Skip to content

Commit 44c46f0

Browse files
committed
make it possible to merge inpainting model with non-inpainting one
1 parent 8504db5 commit 44c46f0

1 file changed

Lines changed: 25 additions & 2 deletions

File tree

‎modules/extras.py‎

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -247,6 +247,7 @@ def add_difference(theta0, theta1_2_diff, alpha):
247247
primary_model_info = sd_models.checkpoints_list[primary_model_name]
248248
secondary_model_info = sd_models.checkpoints_list[secondary_model_name]
249249
teritary_model_info = sd_models.checkpoints_list.get(teritary_model_name, None)
250+
result_is_inpainting_model = False
250251

251252
print(f"Loading {primary_model_info.filename}...")
252253
theta_0 = sd_models.read_state_dict(primary_model_info.filename, map_location='cpu')
@@ -280,8 +281,22 @@ def add_difference(theta0, theta1_2_diff, alpha):
280281

281282
for key in tqdm.tqdm(theta_0.keys()):
282283
if 'model' in key and key in theta_1:
284+
a = theta_0[key]
285+
b = theta_1[key]
283286

284-
theta_0[key] = theta_func2(theta_0[key], theta_1[key], multiplier)
287+
# this enables merging an inpainting model (A) with another one (B);
288+
# where normal model would have 4 channels, for latenst space, inpainting model would
289+
# have another 4 channels for unmasked picture's latent space, plus one channel for mask, for a total of 9
290+
if a.shape != b.shape and a.shape[0:1] + a.shape[2:] == b.shape[0:1] + b.shape[2:]:
291+
if a.shape[1] == 4 and b.shape[1] == 9:
292+
raise RuntimeError("When merging inpainting model with a normal one, A must be the inpainting model.")
293+
294+
assert a.shape[1] == 9 and b.shape[1] == 4, f"Bad dimensions for merged layer {key}: A={a.shape}, B={b.shape}"
295+
296+
theta_0[key][:, 0:4, :, :] = theta_func2(a[:, 0:4, :, :], b, multiplier)
297+
result_is_inpainting_model = True
298+
else:
299+
theta_0[key] = theta_func2(a, b, multiplier)
285300

286301
if save_as_half:
287302
theta_0[key] = theta_0[key].half()
@@ -295,8 +310,16 @@ def add_difference(theta0, theta1_2_diff, alpha):
295310

296311
ckpt_dir = shared.cmd_opts.ckpt_dir or sd_models.model_path
297312

298-
filename = primary_model_info.model_name + '_' + str(round(1-multiplier, 2)) + '-' + secondary_model_info.model_name + '_' + str(round(multiplier, 2)) + '-' + interp_method.replace(" ", "_") + '-merged.' + checkpoint_format
313+
filename = \
314+
primary_model_info.model_name + '_' + str(round(1-multiplier, 2)) + '-' + \
315+
secondary_model_info.model_name + '_' + str(round(multiplier, 2)) + '-' + \
316+
interp_method.replace(" ", "_") + \
317+
'-merged.' + \
318+
("inpainting." if result_is_inpainting_model else "") + \
319+
checkpoint_format
320+
299321
filename = filename if custom_name == '' else (custom_name + '.' + checkpoint_format)
322+
300323
output_modelname = os.path.join(ckpt_dir, filename)
301324

302325
print(f"Saving to {output_modelname}...")

0 commit comments

Comments
 (0)