Skip to content

Commit 77b3e99

Browse files
committed
Add multires_training arg for networks
- multires_training allows creation of buckets down to the set minimum bucket resolution defined.
1 parent 640b169 commit 77b3e99

5 files changed

Lines changed: 54 additions & 19 deletions

File tree

finetune/prepare_buckets_latents.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -97,7 +97,7 @@ def main(args):
9797
), f"illegal resolution (not 'width,height') / 画像サイズに誤りがあります。'幅,高さ'で指定してください: {args.max_resolution}"
9898

9999
bucket_manager = train_util.BucketManager(
100-
args.bucket_no_upscale, max_reso, args.min_bucket_reso, args.max_bucket_reso, args.bucket_reso_steps
100+
args.bucket_no_upscale, max_reso, args.min_bucket_reso, args.max_bucket_reso, args.bucket_reso_steps, args.multires_training
101101
)
102102
if not args.bucket_no_upscale:
103103
bucket_manager.make_buckets()
@@ -242,6 +242,11 @@ def setup_parser() -> argparse.ArgumentParser:
242242
action="store_true",
243243
help="make bucket for each image without upscaling / 画像を拡大せずbucketを作成します",
244244
)
245+
parser.add_argument(
246+
"--multires_training",
247+
action="store_true",
248+
help="make buckets for all resolutions down to minimum resolution for multires training",
249+
)
245250
parser.add_argument(
246251
"--mixed_precision",
247252
type=str,

library/config_util.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -112,6 +112,7 @@ class BaseDatasetParams:
112112
validation_seed: Optional[int] = None
113113
validation_split: float = 0.0
114114
resize_interpolation: Optional[str] = None
115+
multires_training: bool = False
115116

116117
@dataclass
117118
class DreamBoothDatasetParams(BaseDatasetParams):
@@ -251,6 +252,7 @@ def __validate_and_convert_scalar_or_twodim(klass, value: Union[float, Sequence]
251252
"resolution": functools.partial(__validate_and_convert_scalar_or_twodim.__func__, int),
252253
"network_multiplier": float,
253254
"resize_interpolation": str,
255+
"multires_training": bool,
254256
}
255257

256258
# options handled by argparse but not handled by user config
@@ -560,6 +562,7 @@ def print_info(_datasets, dataset_type: str):
560562
max_bucket_reso: {dataset.max_bucket_reso}
561563
bucket_reso_steps: {dataset.bucket_reso_steps}
562564
bucket_no_upscale: {dataset.bucket_no_upscale}
565+
multires_training: {dataset.multires_training}
563566
\n"""), " ")
564567
else:
565568
info += "\n"

library/model_util.py

Lines changed: 14 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1333,7 +1333,7 @@ def use_reflection_padding(vae):
13331333
# endregion
13341334

13351335

1336-
def make_bucket_resolutions(max_reso, min_size=256, max_size=1024, divisible=64):
1336+
def make_bucket_resolutions(max_reso, min_size=256, max_size=1024, divisible=64, multires_training=False):
13371337
max_width, max_height = max_reso
13381338
max_area = max_width * max_height
13391339

@@ -1344,18 +1344,19 @@ def make_bucket_resolutions(max_reso, min_size=256, max_size=1024, divisible=64)
13441344

13451345
width = min_size
13461346
while width <= max_size:
1347-
height = min(max_size, int((max_area // width) // divisible) * divisible)
1348-
if height >= min_size:
1349-
resos.add((width, height))
1350-
resos.add((height, width))
1351-
1352-
# # make additional resos
1353-
# if width >= height and width - divisible >= min_size:
1354-
# resos.add((width - divisible, height))
1355-
# resos.add((height, width - divisible))
1356-
# if height >= width and height - divisible >= min_size:
1357-
# resos.add((width, height - divisible))
1358-
# resos.add((height - divisible, width))
1347+
max_h = min(max_size, int((max_area // width) // divisible) * divisible)
1348+
1349+
if not multires_training:
1350+
height = max_h
1351+
if height >= min_size:
1352+
resos.add((width, height))
1353+
resos.add((height, width))
1354+
else:
1355+
height = min_size
1356+
while height <= max_h:
1357+
resos.add((width, height))
1358+
resos.add((height, width))
1359+
height += divisible
13591360

13601361
width += divisible
13611362

library/train_util.py

Lines changed: 30 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -226,7 +226,8 @@ def __init__(
226226

227227

228228
class 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,

train_network.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1623,6 +1623,7 @@ def load_model_hook(models, input_dir):
16231623
"ss_shuffle_caption": bool(args.shuffle_caption),
16241624
"ss_enable_bucket": bool(dataset.enable_bucket),
16251625
"ss_bucket_no_upscale": bool(dataset.bucket_no_upscale),
1626+
"ss_multires_training": bool(getattr(dataset, "multires_training", False)),
16261627
"ss_min_bucket_reso": dataset.min_bucket_reso,
16271628
"ss_max_bucket_reso": dataset.max_bucket_reso,
16281629
"ss_keep_tokens": args.keep_tokens,

0 commit comments

Comments
 (0)