diff --git a/control.py b/control.py index 16a8e8b..2f20782 100644 --- a/control.py +++ b/control.py @@ -12,10 +12,10 @@ from ldm.modules.diffusionmodules.util import timestep_embedding sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy")) from comfy.cldm import cldm -from comfy.t2i_adapter import adapter from comfy.model_patcher import ModelPatcher -from comfy.controlnet import ControlBase, broadcast_image_to, ControlLora +from comfy.controlnet import ControlBase, ControlNet, T2IAdapter, broadcast_image_to, ControlLora +import comfy.t2i_adapter as t2i_adapter import comfy.utils as utils import comfy.model_management as model_management import comfy.model_detection as model_detection @@ -114,109 +114,75 @@ class TimestepKeyframeGroup: return group -# Copied from comfy.sd, weights modified -class ControlNetAdvanced(ControlBase): +# used to inject ControlNetAdvanced and T2IAdapterAdvanced control_merge function +def control_merge_inject(self, control_input, control_output, control_prev, output_dtype): + out = {'input':[], 'middle':[], 'output': []} + + if control_input is not None: + for i in range(len(control_input)): + key = 'input' + x = control_input[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) + + x *= self.strength * self.weights[i] + if x.dtype != output_dtype: + x = x.to(output_dtype) + out[key].insert(0, x) + + if control_output is not None: + for i in range(len(control_output)): + if i == (len(control_output) - 1): + key = 'middle' + index = 0 + else: + key = 'output' + index = i + x = control_output[i] + if x is not None: + self.apply_advanced_strengths_and_masks(x, self.current_timestep_keyframe, self.batched_number) + + if self.global_average_pooling: + x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) + + x *= self.strength * self.weights[i] + if x.dtype != output_dtype: + x = x.to(output_dtype) + + out[key].append(x) + if control_prev is not None: + for x in ['input', 'middle', 'output']: + o = out[x] + for i in range(len(control_prev[x])): + prev_val = control_prev[x][i] + if i >= len(o): + o.append(prev_val) + elif prev_val is not None: + if o[i] is None: + o[i] = prev_val + else: + o[i] += prev_val + return out + + +class ControlNetAdvanced(ControlNet): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None): - super().__init__(device) - self.control_model = control_model - self.control_model_wrapped = ModelPatcher(self.control_model, load_device=model_management.get_torch_device(), offload_device=model_management.unet_offload_device()) + super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device) self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() - + self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0] + # initialize weights self.weights = self.timestep_keyframes.keyframes[0].control_net_weights if self.timestep_keyframes.keyframes[0].control_net_weights else [1.0]*13 - self.global_average_pooling = global_average_pooling # mask for which parts of controlnet output to keep self.cond_hint_mask = None + # override control_merge + self.control_merge = control_merge_inject.__get__(self, type(self)) def get_control(self, x_noisy, t, cond, batched_number): - #print(f"$$$$ t len: f{len(t)}") - control_prev = None - if self.previous_controlnet is not None: - control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) - - # TODO: select based on progress in diffusion - current_timestep_keyframe = self.timestep_keyframes[0] - - if self.timestep_range is not None: - if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]: - if control_prev is not None: - return control_prev - else: - return None - - output_dtype = x_noisy.dtype - if self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]: - if self.cond_hint is not None: - del self.cond_hint - self.cond_hint = None - self.cond_hint = utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device) - if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) - - - context = cond['c_crossattn'] - y = cond.get('c_adm', None) - if y is not None: - y = y.to(self.control_model.dtype) - control = self.control_model(x=x_noisy.to(self.control_model.dtype), hint=self.cond_hint, timesteps=t, context=context.to(self.control_model.dtype), y=y) - return self.control_merge(None, control, control_prev, output_dtype, current_timestep_keyframe, batched_number) - - def control_merge(self, control_input, control_output, control_prev, output_dtype, current_timestep_keyframe: TimestepKeyframe, batched_number: int): - out = {'input':[], 'middle':[], 'output': []} - - #print(f"$$$$ control_input size: {control_input.size() if control_input != None else None}") - #print(f"$$$$ control_output size: {control_output.size() if control_input != None else None}") - - if control_input is not None: - for i in range(len(control_input)): - key = 'input' - x = control_input[i] - - #print(f"$$$$ x size: {x.size()}") - - self.apply_advanced_strengths_and_masks(x, current_timestep_keyframe, batched_number) - - if x is not None: - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - out[key].insert(0, x) - - if control_output is not None: - for i in range(len(control_output)): - if i == (len(control_output) - 1): - key = 'middle' - index = 0 - else: - key = 'output' - index = i - x = control_output[i] - - #print(f"$$$$ x size: {x.size()}") - - self.apply_advanced_strengths_and_masks(x, current_timestep_keyframe, batched_number) - - if x is not None: - if self.global_average_pooling: - x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3]) - - x *= self.strength * self.weights[i] - if x.dtype != output_dtype: - x = x.to(output_dtype) - - out[key].append(x) - if control_prev is not None: - for x in ['input', 'middle', 'output']: - o = out[x] - for i in range(len(control_prev[x])): - prev_val = control_prev[x][i] - if i >= len(o): - o.append(prev_val) - elif prev_val is not None: - if o[i] is None: - o[i] = prev_val - else: - o[i] += prev_val - return out + # need to reference t and batched_number later + self.t = t + self.batched_number = batched_number + # TODO: choose TimestepKeyframe based on t + return super().get_control(x_noisy, t, cond, batched_number) def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): if current_timestep_keyframe.latent_keyframes is not None: @@ -238,7 +204,6 @@ class ControlNetAdvanced(ControlBase): for b in range(batched_number): x[(latent_count*b)+batch_index] *= 0.0 - def copy(self): c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling) self.copy_to(c) @@ -250,6 +215,35 @@ class ControlNetAdvanced(ControlBase): return out +class T2IAdapterAdvanced(T2IAdapter): + def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroup, channels_in, device=None): + super().__init__(t2i_model=t2i_model, channels_in=channels_in, device=device) + self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() + self.current_timestep_keyframe = self.timestep_keyframes.keyframes[0] + first_weight = self.timestep_keyframes.keyframes[0].t2i_adapter_weights if self.timestep_keyframes.get_index(0) else None + self.weights = first_weight if first_weight else [1.0]*12 + # mask for which parts of controlnet output to keep + self.cond_hint_mask = None + # override control_merge + self.control_merge = control_merge_inject.__get__(self, type(self)) + + def get_control(self, x_noisy, t, cond, batched_number): + # need to reference t and batched_number later + self.t = t + self.batched_number = batched_number + # TODO: choose TimestepKeyframe based on t + return super().get_control(x_noisy, t, cond, batched_number) + + def apply_advanced_strengths_and_masks(self, x, current_timestep_keyframe: TimestepKeyframe, batched_number: int): + # For now, do nothing; need to figure out LatentKeyframe control is even possible for T2I Adapters + return + + def copy(self): + c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in) + self.copy_to(c) + return c + + def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, model=None): controlnet_data = utils.load_torch_file(ckpt_path, safe_load=True) if "lora_controlnet" in controlnet_data: @@ -359,81 +353,6 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo return control -# Copied from comfy.sd, weights modified -class T2IAdapterAdvanced(ControlBase): - def __init__(self, t2i_model, timestep_keyframes: TimestepKeyframeGroup, channels_in, device=None): - super().__init__(device) - self.t2i_model = t2i_model - # TODO: make this actually pull values based on timestep instead of first value - self.timestep_keyframes = timestep_keyframes if timestep_keyframes else TimestepKeyframeGroup() - first_weight = self.timestep_keyframes.keyframes[0].t2i_adapter_weights if self.timestep_keyframes.get_index(0) else None - self.weights = first_weight if first_weight else [1.0]*4 - - self.channels_in = channels_in - self.control_input = None - - def get_control(self, x_noisy, t, cond, batched_number): - control_prev = None - if self.previous_controlnet is not None: - control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) - - if self.timestep_range is not None: - if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]: - if control_prev is not None: - return control_prev - else: - return {} - - if self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]: - if self.cond_hint is not None: - del self.cond_hint - self.control_input = None - self.cond_hint = None - self.cond_hint = utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").float().to(self.device) - if self.channels_in == 1 and self.cond_hint.shape[1] > 1: - self.cond_hint = torch.mean(self.cond_hint, 1, keepdim=True) - if x_noisy.shape[0] != self.cond_hint.shape[0]: - self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) - if self.control_input is None: - self.t2i_model.to(self.device) - self.control_input = self.t2i_model(self.cond_hint) - self.t2i_model.cpu() - - output_dtype = x_noisy.dtype - out = {'input':[]} - - autocast_enabled = torch.is_autocast_enabled() - for i in range(len(self.control_input)): - key = 'input' - x = self.control_input[i] * self.strength * self.weights[i] # apply layer weight - if x.dtype != output_dtype and not autocast_enabled: - x = x.to(output_dtype) - - if control_prev is not None and key in control_prev: - index = len(control_prev[key]) - i * 3 - 3 - prev = control_prev[key][index] - if prev is not None: - x += prev - out[key].insert(0, None) - out[key].insert(0, None) - out[key].insert(0, x) - - if control_prev is not None and 'input' in control_prev: - for i in range(len(out['input'])): - if out['input'][i] is None: - out['input'][i] = control_prev['input'][i] - if control_prev is not None and 'middle' in control_prev: - out['middle'] = control_prev['middle'] - if control_prev is not None and 'output' in control_prev: - out['output'] = control_prev['output'] - return out - - def copy(self): - c = T2IAdapterAdvanced(self.t2i_model, self.timestep_keyframes, self.channels_in) - self.copy_to(c) - return c - - def load_t2i_adapter(t2i_data, timestep_keyframes: TimestepKeyframeGroup=None): keys = t2i_data.keys() if 'adapter' in keys: @@ -441,7 +360,7 @@ def load_t2i_adapter(t2i_data, timestep_keyframes: TimestepKeyframeGroup=None): keys = t2i_data.keys() if "body.0.in_conv.weight" in keys: cin = t2i_data['body.0.in_conv.weight'].shape[1] - model_ad = adapter.Adapter_light(cin=cin, channels=[320, 640, 1280, 1280], nums_rb=4) + model_ad = t2i_adapter.adapter.Adapter_light(cin=cin, channels=[320, 640, 1280, 1280], nums_rb=4) elif 'conv_in.weight' in keys: cin = t2i_data['conv_in.weight'].shape[1] channel = t2i_data['conv_in.weight'].shape[0] @@ -450,8 +369,17 @@ def load_t2i_adapter(t2i_data, timestep_keyframes: TimestepKeyframeGroup=None): down_opts = list(filter(lambda a: a.endswith("down_opt.op.weight"), keys)) if len(down_opts) > 0: use_conv = True - model_ad = adapter.Adapter(cin=cin, channels=[channel, channel*2, channel*4, channel*4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv) + xl = False + if cin == 256 or cin == 768: + xl = True + model_ad = t2i_adapter.adapter.Adapter(cin=cin, channels=[channel, channel*2, channel*4, channel*4][:4], nums_rb=2, ksize=ksize, sk=True, use_conv=use_conv, xl=xl) else: return None - model_ad.load_state_dict(t2i_data) - return T2IAdapterAdvanced(model_ad, timestep_keyframes, cin // 64) + missing, unexpected = model_ad.load_state_dict(t2i_data) + if len(missing) > 0: + print("t2i missing", missing) + + if len(unexpected) > 0: + print("t2i unexpected", unexpected) + + return T2IAdapterAdvanced(model_ad, timestep_keyframes, model_ad.input_channels) diff --git a/nodes.py b/nodes.py index 4b1f18d..3e16e5d 100644 --- a/nodes.py +++ b/nodes.py @@ -14,6 +14,15 @@ from .control import load_controlnet, ControlNetWeightsType, T2IAdapterWeightsTy LatentKeyframe, LatentKeyframeGroup, TimestepKeyframe, TimestepKeyframeGroup +def get_properly_arranged_t2i_weights(initial_weights: list[float]): + new_weights = [] + new_weights.extend([initial_weights[0]]*3) + new_weights.extend([initial_weights[1]]*3) + new_weights.extend([initial_weights[2]]*3) + new_weights.extend([initial_weights[3]]*3) + return new_weights + + class ScaledSoftControlNetWeights: @classmethod def INPUT_TYPES(s): @@ -130,6 +139,7 @@ class SoftT2IAdapterWeights: weights = [weight_00, weight_01, weight_02, weight_03] if flip_weights: weights.reverse() + weights = get_properly_arranged_t2i_weights(weights) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights))) @@ -155,7 +165,7 @@ class CustomT2IAdapterWeights: weights = [weight_00, weight_01, weight_02, weight_03] if flip_weights: weights.reverse() - + weights = get_properly_arranged_t2i_weights(weights) return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(t2i_adapter_weights=weights)))