diff --git a/SteerableMotion.py b/SteerableMotion.py index 04d7a2c..ada9468 100644 --- a/SteerableMotion.py +++ b/SteerableMotion.py @@ -7,7 +7,7 @@ import matplotlib.pyplot as plt import folder_paths -from .imports.IPAdapterPlus import IPAdapterApplyImport, prep_image +from .imports.IPAdapterPlus import (IPAdapterApplyImport, prep_image,IPAdapterBatchEmbedsImport, IPAdapterEncoderImport,) from .imports.AdvancedControlNet import ( calculate_weights, LatentKeyframeInterpolationNodeImport, @@ -261,6 +261,9 @@ class BatchCreativeInterpolationNode: cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights = [], [], [], [] last_key_frame_position = (keyframe_positions[-1]) + buffer + batches = [] + current_batch = [] + for i, (start, end) in enumerate(influence_ranges): # set basic values batch_index_from, batch_index_to_excl = influence_ranges[i] @@ -300,6 +303,8 @@ class BatchCreativeInterpolationNode: control_net_loader = ControlNetLoaderAdvancedImport() apply_advanced_control_net = AdvancedControlNetApplyImport() ipadapter_application = IPAdapterApplyImport() + ipadapter_encoder = IPAdapterEncoderImport() + ipadapter_batcher = IPAdapterBatchEmbedsImport() # Load keyframe and append frame numbers and weights weights, frame_numbers, latent_keyframe = latent_keyframe_interpolation_node.load_keyframe( @@ -318,8 +323,8 @@ class BatchCreativeInterpolationNode: # Prepare image prepped_image = prep_image(image=image.unsqueeze(0), interpolation="LANCZOS", crop_position="pad", sharpening=0.0)[0] - # Adjust strength values and influence range - ipa_strength_from, ipa_strength_to = adjust_strength_values(strength_from, strength_to, ipadapter_strength_multiplier) + # Adjust strength values and influence range + ipa_strength_from, ipa_strength_to = adjust_strength_values(strength_from, strength_to, ipadapter_strength_multiplier) ipa_batch_index_from, ipa_batch_index_to_excl = adjust_influence_range(batch_index_from, batch_index_to_excl, last_key_frame_position, ipadapter_influence_multiplier, buffer) # Calculate weights and append frame numbers and weights @@ -328,9 +333,32 @@ class BatchCreativeInterpolationNode: ipadapter_weights.append(ipa_weights) # Create mask batch and apply ipadapter - masks = create_mask_batch(last_key_frame_position, weights, frame_numbers) - model = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, clip_vision=clip_vision, image=prepped_image, weight_type="original", noise=ipadapter_noise, embeds=None, attn_mask=masks, start_at=0.0, end_at=1.0, unfold_batch=True)[0] + + masks = create_mask_batch(last_key_frame_position, weights, frame_numbers) + + # Apply ipadapter + encoded, = ipadapter_encoder.preprocess(clip_vision, prepped_image, True, ipadapter_noise, 1.0, image_2=None, image_3=None, image_4=None, weight_2=1.0, weight_3=1.0, weight_4=1.0) + + current_batch.append(encoded) + + # If (i+1) is divisible by 8, start a new batch + if (i + 1) % 8 == 0: + batches.append(current_batch) + current_batch = [] + # Add the last batch if it's not empty + if current_batch: + batches.append(current_batch) + + # embeds = ipadapter_batcher.batch(self, embed1, embed2) + + for batch in batches: + # Combine all the encoded data in the batch into a single tensor + embeds = torch.cat(batch, dim=1) + + # Apply ipadapter + model, = ipadapter_application.apply_ipadapter(ipadapter=ipadapter, model=model, weight=1.0, image=None, weight_type="original", noise=None, embeds=embeds, attn_mask=masks, start_at=0.0, end_at=1.0, unfold_batch=True) + comparison_diagram, = plot_weight_comparison(cn_frame_numbers, cn_weights, ipadapter_frame_numbers, ipadapter_weights, buffer) return comparison_diagram, positive, negative, model diff --git a/imports/IPAdapterPlus.py b/imports/IPAdapterPlus.py index 58f4eb9..2251521 100644 --- a/imports/IPAdapterPlus.py +++ b/imports/IPAdapterPlus.py @@ -682,3 +682,99 @@ class ResamplerImport(nn.Module): latents = self.proj_out(latents) return self.norm_out(latents) + + +class IPAdapterEncoderImport: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "clip_vision": ("CLIP_VISION",), + "image_1": ("IMAGE",), + "ipadapter_plus": ("BOOLEAN", { "default": False }), + "noise": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01 }), + "weight_1": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), + }, + "optional": { + "image_2": ("IMAGE",), + "image_3": ("IMAGE",), + "image_4": ("IMAGE",), + "weight_2": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), + "weight_3": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), + "weight_4": ("FLOAT", { "default": 1.0, "min": 0, "max": 1.0, "step": 0.01 }), + } + } + + RETURN_TYPES = ("EMBEDS",) + FUNCTION = "preprocess" + CATEGORY = "ipadapter" + + def preprocess(self, clip_vision, image_1, ipadapter_plus, noise, weight_1, image_2=None, image_3=None, image_4=None, weight_2=1.0, weight_3=1.0, weight_4=1.0): + weight_1 *= (0.1 + (weight_1 - 0.1)) + weight_1 = 1.19e-05 if weight_1 <= 1.19e-05 else weight_1 + weight_2 *= (0.1 + (weight_2 - 0.1)) + weight_2 = 1.19e-05 if weight_2 <= 1.19e-05 else weight_2 + weight_3 *= (0.1 + (weight_3 - 0.1)) + weight_3 = 1.19e-05 if weight_3 <= 1.19e-05 else weight_3 + weight_4 *= (0.1 + (weight_4 - 0.1)) + weight_5 = 1.19e-05 if weight_4 <= 1.19e-05 else weight_4 + + image = image_1 + weight = [weight_1]*image_1.shape[0] + + if image_2 is not None: + if image_1.shape[1:] != image_2.shape[1:]: + image_2 = comfy.utils.common_upscale(image_2.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) + image = torch.cat((image, image_2), dim=0) + weight += [weight_2]*image_2.shape[0] + if image_3 is not None: + if image.shape[1:] != image_3.shape[1:]: + image_3 = comfy.utils.common_upscale(image_3.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) + image = torch.cat((image, image_3), dim=0) + weight += [weight_3]*image_3.shape[0] + if image_4 is not None: + if image.shape[1:] != image_4.shape[1:]: + image_4 = comfy.utils.common_upscale(image_4.movedim(-1,1), image.shape[2], image.shape[1], "bilinear", "center").movedim(1,-1) + image = torch.cat((image, image_4), dim=0) + weight += [weight_4]*image_4.shape[0] + + clip_embed = clip_vision.encode_image(image) + neg_image = image_add_noise(image, noise) if noise > 0 else None + + if ipadapter_plus: + clip_embed = clip_embed.penultimate_hidden_states + if noise > 0: + clip_embed_zeroed = clip_vision.encode_image(neg_image).penultimate_hidden_states + else: + clip_embed_zeroed = zeroed_hidden_states(clip_vision, image.shape[0]) + else: + clip_embed = clip_embed.image_embeds + if noise > 0: + clip_embed_zeroed = clip_vision.encode_image(neg_image).image_embeds + else: + clip_embed_zeroed = torch.zeros_like(clip_embed) + + if any(e != 1.0 for e in weight): + weight = torch.tensor(weight).unsqueeze(-1) if not ipadapter_plus else torch.tensor(weight).unsqueeze(-1).unsqueeze(-1) + clip_embed = clip_embed * weight + + output = torch.stack((clip_embed, clip_embed_zeroed)) + + return( output, ) + + + +class IPAdapterBatchEmbedsImport: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embed1": ("EMBEDS",), + "embed2": ("EMBEDS",), + }} + + RETURN_TYPES = ("EMBEDS",) + FUNCTION = "batch" + CATEGORY = "ipadapter" + + def batch(self, embed1, embed2): + output = torch.cat((embed1, embed2), dim=1) + return (output, ) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index e69de29..0000000