@@ -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