@@ -226,7 +226,8 @@ def __init__(
226226
227227
228228class BucketManager :
229- def __init__ (self , no_upscale , max_reso , min_size , max_size , reso_steps ) -> None :
229+ def __init__ (self , no_upscale , max_reso , min_size , max_size , reso_steps , multires_training = False ) -> None :
230+ self .multires_training = multires_training
230231 if max_size is not None :
231232 if max_reso is not None :
232233 assert max_size >= max_reso [0 ], "the max_size should be larger than the width of max_reso"
@@ -274,7 +275,7 @@ def sort(self):
274275 self .reso_to_id = sorted_reso_to_id
275276
276277 def make_buckets (self ):
277- resos = model_util .make_bucket_resolutions (self .max_reso , self .min_size , self .max_size , self .reso_steps )
278+ resos = model_util .make_bucket_resolutions (self .max_reso , self .min_size , self .max_size , self .reso_steps , self . multires_training )
278279 self .set_predefined_resos (resos )
279280
280281 def set_predefined_resos (self , resos ):
@@ -304,9 +305,19 @@ def select_bucket(self, image_width, image_height):
304305 if reso in self .predefined_resos_set :
305306 pass
306307 else :
307- ar_errors = self .predefined_aspect_ratios - aspect_ratio
308- predefined_bucket_id = np .abs (ar_errors ).argmin () # 当該解像度以外でaspect ratio errorが最も少ないもの
309- reso = self .predefined_resos [predefined_bucket_id ]
308+ ar_errors = np .abs (self .predefined_aspect_ratios - aspect_ratio )
309+ if getattr (self , "multires_training" , False ):
310+ min_ar_error = ar_errors .min ()
311+ # filter out the closest aspect ratios
312+ closest_indices = np .where (ar_errors <= min_ar_error + 1e-4 )[0 ]
313+ # from these, find the one with the closest area
314+ target_area = image_width * image_height
315+ areas = np .array ([self .predefined_resos [i ][0 ] * self .predefined_resos [i ][1 ] for i in closest_indices ])
316+ best_index = closest_indices [np .abs (areas - target_area ).argmin ()]
317+ reso = self .predefined_resos [best_index ]
318+ else :
319+ predefined_bucket_id = ar_errors .argmin () # 当該解像度以外でaspect ratio errorが最も少ないもの
320+ reso = self .predefined_resos [predefined_bucket_id ]
310321
311322 ar_reso = reso [0 ] / reso [1 ]
312323 if aspect_ratio > ar_reso : # 横が長い→縦を合わせる
@@ -1118,6 +1129,7 @@ def make_buckets(self):
11181129 self .min_bucket_reso ,
11191130 self .max_bucket_reso ,
11201131 self .bucket_reso_steps ,
1132+ getattr (self , "multires_training" , False ),
11211133 )
11221134 if not self .bucket_no_upscale :
11231135 self .bucket_manager .make_buckets ()
@@ -2023,6 +2035,7 @@ def __init__(
20232035 max_bucket_reso : int ,
20242036 bucket_reso_steps : int ,
20252037 bucket_no_upscale : bool ,
2038+ multires_training : bool ,
20262039 prior_loss_weight : float ,
20272040 debug_dataset : bool ,
20282041 validation_split : float ,
@@ -2050,11 +2063,13 @@ def __init__(
20502063 self .max_bucket_reso = max_bucket_reso
20512064 self .bucket_reso_steps = bucket_reso_steps
20522065 self .bucket_no_upscale = bucket_no_upscale
2066+ self .multires_training = multires_training
20532067 else :
20542068 self .min_bucket_reso = None
20552069 self .max_bucket_reso = None
20562070 self .bucket_reso_steps = None # この情報は使われない
20572071 self .bucket_no_upscale = False
2072+ self .multires_training = False
20582073
20592074 def read_caption (img_path , caption_extension , enable_wildcard ):
20602075 # captionの候補ファイル名を作る
@@ -2333,6 +2348,7 @@ def __init__(
23332348 max_bucket_reso : int ,
23342349 bucket_reso_steps : int ,
23352350 bucket_no_upscale : bool ,
2351+ multires_training : bool ,
23362352 debug_dataset : bool ,
23372353 validation_seed : int ,
23382354 validation_split : float ,
@@ -2353,11 +2369,13 @@ def __init__(
23532369 self .max_bucket_reso = max_bucket_reso
23542370 self .bucket_reso_steps = bucket_reso_steps
23552371 self .bucket_no_upscale = bucket_no_upscale
2372+ self .multires_training = multires_training
23562373 else :
23572374 self .min_bucket_reso = None
23582375 self .max_bucket_reso = None
23592376 self .bucket_reso_steps = None # この情報は使われない
23602377 self .bucket_no_upscale = False
2378+ self .multires_training = False
23612379
23622380 self .num_train_images = 0
23632381 self .num_reg_images = 0
@@ -2522,6 +2540,7 @@ def __init__(
25222540 max_bucket_reso : int ,
25232541 bucket_reso_steps : int ,
25242542 bucket_no_upscale : bool ,
2543+ multires_training : bool ,
25252544 debug_dataset : bool ,
25262545 validation_split : float ,
25272546 validation_seed : Optional [int ],
@@ -2577,6 +2596,7 @@ def __init__(
25772596 max_bucket_reso ,
25782597 bucket_reso_steps ,
25792598 bucket_no_upscale ,
2599+ multires_training ,
25802600 1.0 ,
25812601 debug_dataset ,
25822602 validation_split ,
@@ -4821,6 +4841,11 @@ def add_dataset_arguments(
48214841 action = "store_true" ,
48224842 help = "make bucket for each image without upscaling / 画像を拡大せずbucketを作成します" ,
48234843 )
4844+ parser .add_argument (
4845+ "--multires_training" ,
4846+ action = "store_true" ,
4847+ help = "make buckets for all resolutions down to minimum resolution for multires training" ,
4848+ )
48244849 parser .add_argument (
48254850 "--resize_interpolation" ,
48264851 type = str ,
0 commit comments