Refactored MotionModelPatcher into MotionModelAttachment so that no custom ModelPatcher model is required; no more mismatches between ComfyUI and ADE ModelPatcher features. Also made Scale Ref Image and VAE Encode node use batched vae encoding

This commit is contained in:
Jedrzej Kosinski
2024-11-27 23:31:55 -06:00
parent 32a06f1e15
commit 0a2d1291c9
7 changed files with 137 additions and 134 deletions
+110 -106
View File
@@ -14,7 +14,7 @@ import comfy.lora
import comfy.model_management
import comfy.utils
from comfy.model_patcher import ModelPatcher
from comfy.patcher_extension import WrappersMP, PatcherInjection
from comfy.patcher_extension import CallbacksMP, WrappersMP, PatcherInjection
from comfy.model_base import BaseModel
from comfy.sd import CLIP, VAE
@@ -34,6 +34,12 @@ from .sample_settings import SampleSettings, SeedNoiseGeneration
from .dinklink import get_acn_outer_sample_wrapper
class MotionModelPatcher(ModelPatcher):
'''Class used only for type hints.'''
def __init__(self):
self.model: AnimateDiffModel
class ModelPatcherHelper:
SAMPLE_SETTINGS = "ADE_sample_settings"
PARAMS = "ADE_params"
@@ -55,13 +61,10 @@ class ModelPatcherHelper:
self.remove_motion_models()
self.remove_forward_timestep_embed_patch()
def get_adgs(self):
pass
def get_motion_models(self) -> list['MotionModelPatcher']:
def get_motion_models(self) -> list[MotionModelPatcher]:
return self.model.additional_models.get(self.ADE, [])
def set_motion_models(self, motion_models: list['MotionModelPatcher']):
def set_motion_models(self, motion_models: list[MotionModelPatcher]):
self.model.set_additional_models(self.ADE, motion_models)
self.model.set_injections(self.ADE,
[PatcherInjection(inject=inject_motion_models, eject=eject_motion_models)])
@@ -152,7 +155,7 @@ class ModelPatcherHelper:
def pre_run(self):
# TODO: could implement this as a ModelPatcher ON_PRE_RUN callback
for motion_model in self.get_motion_models():
motion_model.pre_run(self.model)
motion_model.pre_run()
self.get_sample_settings().pre_run(self.model)
@@ -177,11 +180,61 @@ def forward_timestep_embed_patch_ade(layer, x, emb, context, transformer_options
return layer(x, context)
class MotionModelPatcher(ModelPatcher):
# Mostly here so that type hints work in IDEs
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.model: AnimateDiffModel = self.model
def create_MotionModelPatcher(model, load_device, offload_device) -> MotionModelPatcher:
patcher = ModelPatcher(model, load_device=load_device, offload_device=offload_device)
ade = ModelPatcherHelper.ADE
patcher.add_callback_with_key(CallbacksMP.ON_LOAD, ade, _mm_patch_lowvram_extras_callback)
patcher.add_callback_with_key(CallbacksMP.ON_LOAD, ade, _mm_handle_float8_pe_tensors_callback)
patcher.add_callback_with_key(CallbacksMP.ON_PRE_RUN, ade, _mm_pre_run_callback)
patcher.add_callback_with_key(CallbacksMP.ON_CLEANUP, ade, _mm_clean_callback)
patcher.set_attachments(ade, MotionModelAttachment())
return patcher
def _mm_patch_lowvram_extras_callback(self: MotionModelPatcher, device_to, lowvram_model_memory, *args, **kwargs):
if lowvram_model_memory > 0:
# figure out the tensors (likely pe's) that should be cast to device besides just the named_modules
remaining_tensors = list(self.model.state_dict().keys())
named_modules = []
for n, _ in self.model.named_modules():
named_modules.append(n)
named_modules.append(f"{n}.weight")
named_modules.append(f"{n}.bias")
for name in named_modules:
if name in remaining_tensors:
remaining_tensors.remove(name)
for key in remaining_tensors:
self.patch_weight_to_device(key, device_to)
if device_to is not None:
comfy.utils.set_attr(self.model, key, comfy.utils.get_attr(self.model, key).to(device_to))
def _mm_handle_float8_pe_tensors_callback(self: MotionModelPatcher, *args, **kwargs):
remaining_tensors = list(self.model.state_dict().keys())
pe_tensors = [x for x in remaining_tensors if '.pe' in x]
is_first = True
for key in pe_tensors:
if is_first:
is_first = False
if comfy.utils.get_attr(self.model, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]:
break
comfy.utils.set_attr(self.model, key, comfy.utils.get_attr(self.model, key).half())
def _mm_pre_run_callback(self: MotionModelPatcher, *args, **kwargs):
attachment = get_mm_attachment(self)
attachment.pre_run(self)
def _mm_clean_callback(self: MotionModelPatcher, *args, **kwargs):
attachment = get_mm_attachment(self)
attachment.cleanup(self)
def get_mm_attachment(patcher: MotionModelPatcher) -> 'MotionModelAttachment':
return patcher.get_attachment(ModelPatcherHelper.ADE)
class MotionModelAttachment:
def __init__(self):
self.timestep_percent_range = (0.0, 1.0)
self.timestep_range: tuple[float, float] = None
self.keyframes: ADKeyframeGroup = ADKeyframeGroup()
@@ -239,49 +292,14 @@ class MotionModelPatcher(ModelPatcher):
self.prev_sub_idxs = None
self.prev_batched_number = None
def load(self, device_to=None, lowvram_model_memory=0, *args, **kwargs):
to_return = super().load(device_to=device_to, lowvram_model_memory=lowvram_model_memory, *args, **kwargs)
if lowvram_model_memory > 0:
self._patch_lowvram_extras(device_to=device_to)
self._handle_float8_pe_tensors()
return to_return
def _patch_lowvram_extras(self, device_to=None):
# figure out the tensors (likely pe's) that should be cast to device besides just the named_modules
remaining_tensors = list(self.model.state_dict().keys())
named_modules = []
for n, _ in self.model.named_modules():
named_modules.append(n)
named_modules.append(f"{n}.weight")
named_modules.append(f"{n}.bias")
for name in named_modules:
if name in remaining_tensors:
remaining_tensors.remove(name)
for key in remaining_tensors:
self.patch_weight_to_device(key, device_to)
if device_to is not None:
comfy.utils.set_attr(self.model, key, comfy.utils.get_attr(self.model, key).to(device_to))
def _handle_float8_pe_tensors(self):
remaining_tensors = list(self.model.state_dict().keys())
pe_tensors = [x for x in remaining_tensors if '.pe' in x]
is_first = True
for key in pe_tensors:
if is_first:
is_first = False
if comfy.utils.get_attr(self.model, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]:
break
comfy.utils.set_attr(self.model, key, comfy.utils.get_attr(self.model, key).half())
def pre_run(self, model: ModelPatcher):
self.cleanup()
self.model.set_scale(self.scale_multival, self.per_block_list)
self.model.set_effect(self.effect_multival, self.per_block_list)
self.model.set_cameractrl_effect(self.cameractrl_multival)
if self.model.img_encoder is not None:
self.model.img_encoder.set_ref_drift(self.orig_ref_drift)
self.model.img_encoder.set_insertion_weights(self.orig_insertion_weights)
def pre_run(self, patcher: MotionModelPatcher):
self.cleanup(patcher)
patcher.model.set_scale(self.scale_multival, self.per_block_list)
patcher.model.set_effect(self.effect_multival, self.per_block_list)
patcher.model.set_cameractrl_effect(self.cameractrl_multival)
if patcher.model.img_encoder is not None:
patcher.model.img_encoder.set_ref_drift(self.orig_ref_drift)
patcher.model.img_encoder.set_insertion_weights(self.orig_insertion_weights)
def initialize_timesteps(self, model: BaseModel):
self.timestep_range = (model.model_sampling.percent_to_sigma(self.timestep_percent_range[0]),
@@ -290,7 +308,7 @@ class MotionModelPatcher(ModelPatcher):
for keyframe in self.keyframes.keyframes:
keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent)
def prepare_current_keyframe(self, x: Tensor, t: Tensor):
def prepare_current_keyframe(self, patcher: MotionModelPatcher, x: Tensor, t: Tensor):
curr_t: float = t[0]
# if curr_t was previous_t, then do nothing (already accounted for this step)
if curr_t == self.previous_t:
@@ -340,26 +358,26 @@ class MotionModelPatcher(ModelPatcher):
self.combined_pia_mask = get_combined_input(self.pia_input, self.current_pia_input, x)
self.combined_pia_effect = get_combined_input_effect_multival(self.pia_input, self.current_pia_input)
# apply scale and effect
self.model.set_scale(self.combined_scale, self.per_block_list)
self.model.set_effect(self.combined_effect, self.per_block_list) # TODO: set combined_per_block_list
self.model.set_cameractrl_effect(self.combined_cameractrl_effect)
patcher.model.set_scale(self.combined_scale, self.per_block_list)
patcher.model.set_effect(self.combined_effect, self.per_block_list) # TODO: set combined_per_block_list
patcher.model.set_cameractrl_effect(self.combined_cameractrl_effect)
# apply effect - if not within range, set effect to 0, effectively turning model off
if curr_t > self.timestep_range[0] or curr_t < self.timestep_range[1]:
self.model.set_effect(0.0)
patcher.model.set_effect(0.0)
self.was_within_range = False
else:
# if was not in range last step, apply effect to toggle AD status
if not self.was_within_range:
self.model.set_effect(self.combined_effect, self.per_block_list)
patcher.model.set_effect(self.combined_effect, self.per_block_list)
self.was_within_range = True
# update steps current keyframe is used
self.current_used_steps += 1
# update previous_t
self.previous_t = curr_t
def prepare_alcmi2v_features(self, x: Tensor, cond_or_uncond: list[int], ad_params: dict[str], latent_format):
def prepare_alcmi2v_features(self, patcher: MotionModelPatcher, x: Tensor, cond_or_uncond: list[int], ad_params: dict[str], latent_format):
# if no img_encoder, done
if self.model.img_encoder is None:
if patcher.model.img_encoder is None:
return
batched_number = len(cond_or_uncond)
full_length = ad_params["full_length"]
@@ -372,20 +390,20 @@ class MotionModelPatcher(ModelPatcher):
img_latents = comfy.utils.common_upscale(self.orig_img_latents[sub_idxs], x.shape[3], x.shape[2], 'nearest-exact', 'center').to(x.dtype).to(x.device)
else:
img_latents = comfy.utils.common_upscale(self.orig_img_latents, x.shape[3], x.shape[2], 'nearest-exact', 'center').to(x.dtype).to(x.device)
img_latents = latent_format.process_in(img_latents)
img_latents: Tensor = latent_format.process_in(img_latents)
# make sure img_latents matches goal_length
if goal_length != img_latents.shape[0]:
img_latents = ade_broadcast_image_to(img_latents, goal_length, batched_number)
img_features = self.model.img_encoder(img_latents, goal_length, batched_number)
self.model.set_img_features(img_features=img_features, apply_ref_when_disabled=self.orig_apply_ref_when_disabled)
img_features = patcher.model.img_encoder(img_latents, goal_length, batched_number)
patcher.model.set_img_features(img_features=img_features, apply_ref_when_disabled=self.orig_apply_ref_when_disabled)
# cache values for next step
self.img_latents_shape = img_latents.shape
self.prev_sub_idxs = sub_idxs
self.prev_batched_number = batched_number
def prepare_camera_features(self, x: Tensor, cond_or_uncond: list[int], ad_params: dict[str]):
def prepare_camera_features(self, patcher: MotionModelPatcher, x: Tensor, cond_or_uncond: list[int], ad_params: dict[str]):
# if no camera_encoder, done
if self.model.camera_encoder is None:
if patcher.model.camera_encoder is None:
return
batched_number = len(cond_or_uncond)
full_length = ad_params["full_length"]
@@ -410,8 +428,8 @@ class MotionModelPatcher(ModelPatcher):
# create encoded embeddings
b, c, h, w = x.shape
plucker_embedding = prepare_pose_embedding(camera_poses, image_width=w*8, image_height=h*8).to(dtype=x.dtype, device=x.device)
camera_embedding = self.model.camera_encoder(plucker_embedding, video_length=goal_length, batched_number=batched_number)
self.model.set_camera_features(camera_features=camera_embedding)
camera_embedding = patcher.model.camera_encoder(plucker_embedding, video_length=goal_length, batched_number=batched_number)
patcher.model.set_camera_features(camera_features=camera_embedding)
self.camera_features_shape = len(camera_embedding)
self.prev_sub_idxs = sub_idxs
self.prev_batched_number = batched_number
@@ -517,16 +535,15 @@ class MotionModelPatcher(ModelPatcher):
finally:
comfy.model_management.load_models_gpu(cached_loaded_models)
def is_pia(self):
return self.model.mm_info.mm_format == AnimateDiffFormat.PIA and self.orig_pia_images is not None
def is_pia(self, patcher: MotionModelPatcher):
return patcher.model.mm_info.mm_format == AnimateDiffFormat.PIA and self.orig_pia_images is not None
def is_fancyvideo(self):
return self.model.mm_info.mm_format == AnimateDiffFormat.FANCYVIDEO
def is_fancyvideo(self, patcher: MotionModelPatcher):
return patcher.model.mm_info.mm_format == AnimateDiffFormat.FANCYVIDEO
def cleanup(self):
super().cleanup()
if self.model is not None:
self.model.cleanup()
def cleanup(self, patcher: MotionModelPatcher):
if patcher.model is not None:
patcher.model.cleanup()
# AnimateLCM-I2V
del self.img_features
self.img_features = None
@@ -552,24 +569,8 @@ class MotionModelPatcher(ModelPatcher):
self.prev_sub_idxs = None
self.prev_batched_number = None
def clone(self):
# normal ModelPatcher clone actions
n = MotionModelPatcher(self.model, self.load_device, self.offload_device, self.size, weight_inplace_update=self.weight_inplace_update)
n.patches = {}
for k in self.patches:
n.patches[k] = self.patches[k][:]
if hasattr(n, "patches_uuid"):
self.patches_uuid = n.patches_uuid
n.object_patches = self.object_patches.copy()
n.model_options = copy.deepcopy(self.model_options)
if hasattr(n, "model_keys"):
n.model_keys = self.model_keys
if hasattr(n, "backup"):
self.backup = n.backup
if hasattr(n, "object_patches_backup"):
self.object_patches_backup = n.object_patches_backup
n.parent = self
def on_model_patcher_clone(self):
n = MotionModelAttachment()
# extra cloned params
n.timestep_percent_range = self.timestep_percent_range
n.timestep_range = self.timestep_range
@@ -635,11 +636,12 @@ class MotionModelGroup:
def initialize_timesteps(self, model: BaseModel):
for motion_model in self.models:
motion_model.initialize_timesteps(model)
attachment = get_mm_attachment(motion_model)
attachment.initialize_timesteps(model)
def pre_run(self, model: ModelPatcher):
for motion_model in self.models:
motion_model.pre_run(model)
motion_model.pre_run()
def cleanup(self):
for motion_model in self.models:
@@ -647,12 +649,14 @@ class MotionModelGroup:
def prepare_current_keyframe(self, x: Tensor, t: Tensor):
for motion_model in self.models:
motion_model.prepare_current_keyframe(x=x, t=t)
attachment = get_mm_attachment(motion_model)
attachment.prepare_current_keyframe(motion_model, x=x, t=t)
def get_special_models(self):
pia_motion_models: list[MotionModelPatcher] = []
for motion_model in self.models:
if motion_model.is_pia() or motion_model.is_fancyvideo():
attachment = get_mm_attachment(motion_model)
if attachment.is_pia(motion_model) or attachment.is_fancyvideo(motion_model):
pia_motion_models.append(motion_model)
return pia_motion_models
@@ -759,7 +763,7 @@ def load_motion_module_gen1(model_name: str, model: ModelPatcher, motion_lora: M
load_result = ad_wrapper.load_state_dict(mm_state_dict, strict=False)
verify_load_result(load_result=load_result, mm_info=mm_info)
# wrap motion_module into a ModelPatcher, to allow motion lora patches
motion_model = MotionModelPatcher(model=ad_wrapper, load_device=model.load_device, offload_device=model.offload_device)
motion_model = create_MotionModelPatcher(model=ad_wrapper, load_device=model.load_device, offload_device=model.offload_device)
# load motion_lora, if present
if motion_lora is not None:
for lora in motion_lora.loras:
@@ -783,8 +787,8 @@ def load_motion_module_gen2(model_name: str, motion_model_settings: AnimateDiffS
load_result = ad_wrapper.load_state_dict(mm_state_dict, strict=False)
verify_load_result(load_result=load_result, mm_info=mm_info)
# wrap motion_module into a ModelPatcher, to allow motion lora patches
motion_model = MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
offload_device=comfy.model_management.unet_offload_device())
motion_model = create_MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
offload_device=comfy.model_management.unet_offload_device())
return motion_model
@@ -823,7 +827,7 @@ def create_fresh_motion_module(motion_model: MotionModelPatcher) -> MotionModelP
ad_wrapper.to(comfy.model_management.unet_dtype())
ad_wrapper.to(comfy.model_management.unet_offload_device())
ad_wrapper.load_state_dict(motion_model.model.state_dict())
return MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
return create_MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
offload_device=comfy.model_management.unet_offload_device())
@@ -832,8 +836,8 @@ def create_fresh_encoder_only_model(motion_model: MotionModelPatcher) -> MotionM
ad_wrapper.to(comfy.model_management.unet_dtype())
ad_wrapper.to(comfy.model_management.unet_offload_device())
ad_wrapper.load_state_dict(motion_model.model.state_dict(), strict=False)
return MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
offload_device=comfy.model_management.unet_offload_device())
return create_MotionModelPatcher(model=ad_wrapper, load_device=comfy.model_management.get_torch_device(),
offload_device=comfy.model_management.unet_offload_device())
def inject_img_encoder_into_model(motion_model: MotionModelPatcher, w_encoder: MotionModelPatcher):
+7 -11
View File
@@ -7,10 +7,10 @@ from comfy.sd import VAE
from .ad_settings import AnimateDiffSettings
from .logger import logger
from .utils_model import ScaleMethods, CropMethods, get_available_motion_models
from .utils_model import ScaleMethods, CropMethods, get_available_motion_models, vae_encode_raw_batched
from .utils_motion import ADKeyframeGroup
from .motion_lora import MotionLoraList
from .model_injection import (MotionModelGroup, MotionModelPatcher, create_fresh_encoder_only_model,
from .model_injection import (MotionModelGroup, MotionModelPatcher, get_mm_attachment, create_fresh_encoder_only_model,
load_motion_module_gen2, inject_img_encoder_into_model)
from .motion_module_ad import AnimateDiffFormat
from .nodes_gen2 import ApplyAnimateDiffModelNode
@@ -58,9 +58,10 @@ class ApplyAnimateLCMI2VModel:
# confirm that model contains img_encoder
if curr_model.model.img_encoder is None:
raise Exception(f"Motion model '{curr_model.model.mm_info.mm_name}' does not contain an img_encoder; cannot be used with Apply AnimateLCM-I2V Model node.")
curr_model.orig_img_latents = ref_latent["samples"]
curr_model.orig_ref_drift = ref_drift
curr_model.orig_apply_ref_when_disabled = apply_ref_when_disabled
attachment = get_mm_attachment(curr_model)
attachment.orig_img_latents = ref_latent["samples"]
attachment.orig_ref_drift = ref_drift
attachment.orig_apply_ref_when_disabled = apply_ref_when_disabled
return new_m_models
@@ -148,9 +149,4 @@ class UpscaleAndVaeEncode:
image = comfy.utils.common_upscale(samples=image, width=w*8, height=h*8, upscale_method=scale_method, crop=crop)
image = image.movedim(1,-1)
# now that images are the expected size, VAEEncode them
try: # account for old ComfyUI versions (TODO: remove this when other changes require ComfyUI update)
if not hasattr(vae, "vae_encode_crop_pixels"):
image = VAEEncode.vae_encode_crop_pixels(image)
except Exception:
pass
return ({"samples": vae.encode(image[:,:,:,:3])},)
return ({"samples": vae_encode_raw_batched(vae, image)},)
+4 -3
View File
@@ -16,7 +16,7 @@ from .logger import logger
from .utils_model import get_available_motion_models, calculate_file_hash, strip_path, BIGMAX
from .utils_motion import ADKeyframeGroup
from .motion_lora import MotionLoraList
from .model_injection import (MotionModelGroup, MotionModelPatcher, load_motion_module_gen2, inject_camera_encoder_into_model)
from .model_injection import (MotionModelGroup, MotionModelPatcher, get_mm_attachment, load_motion_module_gen2, inject_camera_encoder_into_model)
from .nodes_gen2 import ApplyAnimateDiffModelNode, ADKeyframeNode
@@ -230,8 +230,9 @@ class ApplyAnimateDiffWithCameraCtrl:
if curr_model.model.camera_encoder is None:
raise Exception(f"Motion model '{curr_model.model.mm_info.mm_name}' does not contain a camera_encoder; cannot be used with Apply AnimateDiff-CameraCtrl Model node.")
camera_entries = [CameraEntry(entry) for entry in cameractrl_poses]
curr_model.orig_camera_entries = camera_entries
curr_model.cameractrl_multival = cameractrl_multival
attachment = get_mm_attachment(curr_model)
attachment.orig_camera_entries = camera_entries
attachment.cameractrl_multival = cameractrl_multival
return new_m_models
+4 -3
View File
@@ -10,7 +10,7 @@ from .logger import logger
from .utils_model import BIGMIN, BIGMAX, get_available_motion_models
from .utils_motion import ADKeyframeGroup, InputPIA, InputPIA_Multival, extend_list_to_batch_size, extend_to_batch_size, prepare_mask_batch
from .motion_lora import MotionLoraList
from .model_injection import MotionModelGroup, MotionModelPatcher, load_motion_module_gen2, inject_pia_conv_in_into_model
from .model_injection import MotionModelGroup, MotionModelPatcher, get_mm_attachment, load_motion_module_gen2, inject_pia_conv_in_into_model
from .motion_module_ad import AnimateDiffFormat
from .nodes_gen2 import ApplyAnimateDiffModelNode, ADKeyframeNode
@@ -59,6 +59,7 @@ class ApplyAnimateDiffFancyVideo:
# confirm that model is FancyVideo
if curr_model.model.mm_info.mm_format != AnimateDiffFormat.FANCYVIDEO:
raise Exception(f"Motion model '{curr_model.model.mm_info.mm_name}' is not a FancyVideo model; cannot be used with Apply AD-FancyModel Model node.")
curr_model.orig_fancy_images = image
curr_model.fancy_vae = vae
attachment = get_mm_attachment(curr_model)
attachment.orig_fancy_images = image
attachment.fancy_vae = vae
return new_m_models
+7 -6
View File
@@ -11,7 +11,7 @@ from .utils_motion import ADKeyframeGroup, ADKeyframe, InputPIA
from .motion_lora import MotionLoraList
from .motion_module_ad import AllPerBlocks
from .model_injection import (ModelPatcherHelper,
InjectionParams, MotionModelGroup, MotionModelPatcher, create_fresh_motion_module,
InjectionParams, MotionModelGroup, MotionModelPatcher, get_mm_attachment, create_fresh_motion_module,
load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2, validate_per_block_compatibility)
from .sample_settings import SampleSettings
from .sampling import outer_sample_wrapper, sliding_calc_cond_batch
@@ -128,13 +128,14 @@ class ApplyAnimateDiffModelNode:
if motion_lora is not None:
for lora in motion_lora.loras:
load_motion_lora_as_patches(motion_model, lora)
motion_model.scale_multival = scale_multival
motion_model.effect_multival = effect_multival
attachment = get_mm_attachment(motion_model)
attachment.scale_multival = scale_multival
attachment.effect_multival = effect_multival
if per_block is not None:
validate_per_block_compatibility(motion_model=motion_model, all_per_blocks=per_block)
motion_model.per_block_list = per_block.per_block_list
motion_model.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeyframeGroup()
motion_model.timestep_percent_range = (start_percent, end_percent)
attachment.per_block_list = per_block.per_block_list
attachment.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeyframeGroup()
attachment.timestep_percent_range = (start_percent, end_percent)
# add to beginning, so that after injection, it will be the earliest of prev_m_models to be run
prev_m_models.add_to_start(mm=motion_model)
return (prev_m_models,)
+5 -4
View File
@@ -10,7 +10,7 @@ from .logger import logger
from .utils_model import BIGMIN, BIGMAX, get_available_motion_models
from .utils_motion import ADKeyframeGroup, InputPIA, InputPIA_Multival, extend_list_to_batch_size, extend_to_batch_size, prepare_mask_batch
from .motion_lora import MotionLoraList
from .model_injection import MotionModelGroup, MotionModelPatcher, load_motion_module_gen2, inject_pia_conv_in_into_model
from .model_injection import MotionModelGroup, MotionModelPatcher, get_mm_attachment, load_motion_module_gen2, inject_pia_conv_in_into_model
from .motion_module_ad import AnimateDiffFormat
from .nodes_gen2 import ApplyAnimateDiffModelNode, ADKeyframeNode
@@ -148,11 +148,12 @@ class ApplyAnimateDiffPIAModel:
# confirm that model is PIA
if curr_model.model.mm_info.mm_format != AnimateDiffFormat.PIA:
raise Exception(f"Motion model '{curr_model.model.mm_info.mm_name}' is not a PIA model; cannot be used with Apply AnimateDiff-PIA Model node.")
curr_model.orig_pia_images = image
curr_model.pia_vae = vae
attachment = get_mm_attachment(curr_model)
attachment.orig_pia_images = image
attachment.pia_vae = vae
if pia_input is None:
pia_input = InputPIA_Multival(1.0)
curr_model.pia_input = pia_input
attachment.pia_input = pia_input
#curr_model.pia_multival = ref_multival
return new_m_models
-1
View File
@@ -37,7 +37,6 @@ def vae_encode_raw_dynamic_batched(vae: VAE, pixels: Tensor, max_batch=16, min_b
b, h, w, c = pixels.shape
actual_size = h*w
actual_batch_size = int(max(min_batch, min(max_batch, max_batch // max((actual_size / max_size), 1.0))))
logger.info(f"actual_batch_size: {actual_batch_size}")
return vae_encode_raw_batched(vae=vae, pixels=pixels, per_batch=actual_batch_size, show_pbar=show_pbar)