Added multigpu support to a bunch of core AnimateDiff features via refactor
This commit is contained in:
@@ -69,7 +69,14 @@ class ModelPatcherHelper:
|
||||
self.remove_motion_models()
|
||||
self.remove_forward_timestep_embed_patch()
|
||||
|
||||
def get_motion_models(self) -> list[MotionModelPatcher]:
|
||||
def get_motion_models(self, all_devices=False) -> list[MotionModelPatcher]:
|
||||
if all_devices:
|
||||
patchers = [self.model]
|
||||
patchers.extend(self.model.get_additional_models_with_key("multigpu"))
|
||||
all_motion_models = []
|
||||
for patcher in patchers:
|
||||
all_motion_models.extend(patcher.get_additional_models_with_key(self.ADE))
|
||||
return all_motion_models
|
||||
return self.model.additional_models.get(self.ADE, [])
|
||||
|
||||
def set_motion_models(self, motion_models: list[MotionModelPatcher]):
|
||||
@@ -82,7 +89,7 @@ class ModelPatcherHelper:
|
||||
self.model.remove_injections(self.ADE)
|
||||
|
||||
def cleanup_motion_models(self):
|
||||
for motion_model in self.get_motion_models():
|
||||
for motion_model in self.get_motion_models(all_devices=True):
|
||||
motion_model.cleanup()
|
||||
|
||||
|
||||
@@ -108,7 +115,7 @@ class ModelPatcherHelper:
|
||||
##########################
|
||||
# motion models helpers
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
for motion_model in self.get_motion_models():
|
||||
for motion_model in self.get_motion_models(all_devices=True):
|
||||
motion_model.model.set_video_length(video_length=video_length, full_length=full_length)
|
||||
|
||||
def get_name_string(self, show_version=False):
|
||||
@@ -164,7 +171,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():
|
||||
for motion_model in self.get_motion_models(all_devices=True):
|
||||
motion_model.pre_run()
|
||||
self.get_sample_settings().pre_run(self.model)
|
||||
|
||||
@@ -670,7 +677,7 @@ class MotionModelAttachment:
|
||||
|
||||
|
||||
class MotionModelGroup:
|
||||
def __init__(self, init_motion_model: MotionModelPatcher=None):
|
||||
def __init__(self, init_motion_model: Union[MotionModelPatcher, list[MotionModelPatcher]]=None):
|
||||
self.models: list[MotionModelPatcher] = []
|
||||
if init_motion_model is not None:
|
||||
if isinstance(init_motion_model, list):
|
||||
|
||||
+68
-38
@@ -40,57 +40,82 @@ class AnimateDiffGlobalState:
|
||||
def __init__(self):
|
||||
self.model_patcher: ModelPatcher = None
|
||||
self.motion_models: MotionModelGroup = None
|
||||
|
||||
self.model_patcher_devices: dict[torch.device, ModelPatcher] = {}
|
||||
self.motion_models_devices: dict[torch.device, MotionModelGroup] = {}
|
||||
|
||||
self.params: InjectionParams = None
|
||||
self.sample_settings: SampleSettings = None
|
||||
self.callback_output_dict: dict[str] = {}
|
||||
self.function_injections: FunctionInjectionHolder = None
|
||||
self.reset()
|
||||
|
||||
def initialize(self, model: BaseModel):
|
||||
def initialize(self, model: BaseModel, model_options: dict[str]):
|
||||
# this function is to be run in sampling func
|
||||
if not self.initialized:
|
||||
self.initialized = True
|
||||
if self.motion_models is not None:
|
||||
self.motion_models.initialize_timesteps(model)
|
||||
# initialize multigpu stuff
|
||||
if "multigpu_clones" in model_options:
|
||||
self.model_patcher_devices = model_options["multigpu_clones"].copy()
|
||||
for device, patcher in self.model_patcher_devices.items():
|
||||
mm_list = ModelPatcherHelper(patcher).get_motion_models()
|
||||
if len(mm_list) > 0:
|
||||
self.motion_models_devices[device] = MotionModelGroup(mm_list)
|
||||
else:
|
||||
self.model_patcher_devices[self.model_patcher.load_device] = self.model_patcher
|
||||
if self.motion_models is not None:
|
||||
self.motion_models_devices[self.model_patcher.load_device] = self.motion_models
|
||||
# initialize timesteps
|
||||
for device, motion_models in self.motion_models_devices.items():
|
||||
base_model = self.model_patcher_devices[device].model
|
||||
motion_models.initialize_timesteps(base_model)
|
||||
if self.params.context_options is not None:
|
||||
self.params.context_options.initialize_timesteps(model)
|
||||
if self.sample_settings.custom_cfg is not None:
|
||||
self.sample_settings.custom_cfg.initialize_timesteps(model)
|
||||
|
||||
def prepare_current_keyframes(self, x: Tensor, timestep: Tensor, transformer_options: dict[str, Tensor]):
|
||||
if self.motion_models is not None:
|
||||
self.motion_models.prepare_current_keyframe(x=x, t=timestep, transformer_options=transformer_options)
|
||||
for motion_models in self.motion_models_devices.values():
|
||||
motion_models.prepare_current_keyframe(x=x, t=timestep, transformer_options=transformer_options)
|
||||
if self.params.context_options is not None:
|
||||
self.params.context_options.prepare_current(t=timestep, transformer_options=transformer_options)
|
||||
if self.sample_settings.custom_cfg is not None:
|
||||
self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep, transformer_options=transformer_options)
|
||||
|
||||
def perform_special_model_features(self, model: BaseModel, conds: list, x_in: Tensor, model_options: dict[str]):
|
||||
if self.motion_models is not None:
|
||||
special_models = self.motion_models.get_special_models()
|
||||
def perform_special_model_features(self, conds: list, x_in: Tensor, model_options: dict[str]):
|
||||
# despite there be
|
||||
clone_uuids = set()
|
||||
for device, motion_models in self.motion_models_devices.items():
|
||||
model: BaseModel = self.model_patcher_devices[device].model
|
||||
special_models = motion_models.get_special_models()
|
||||
if len(special_models) > 0:
|
||||
for special_model in special_models:
|
||||
if special_model.model.is_in_effect():
|
||||
attachment = get_mm_attachment(special_model)
|
||||
if attachment.is_pia(special_model):
|
||||
special_model.model.inject_unet_conv_in_pia_fancyvideo(model)
|
||||
conds = get_conds_with_c_concat(conds,
|
||||
attachment.get_pia_c_concat(model, x_in))
|
||||
if special_model.clone_base_uuid not in clone_uuids:
|
||||
clone_uuids.add(special_model.clone_base_uuid)
|
||||
conds = get_conds_with_c_concat(conds,
|
||||
attachment.get_pia_c_concat(model, x_in))
|
||||
elif attachment.is_fancyvideo(special_model):
|
||||
# TODO: handle other weights
|
||||
special_model.model.inject_unet_conv_in_pia_fancyvideo(model)
|
||||
conds = get_conds_with_c_concat(conds,
|
||||
attachment.get_fancy_c_concat(model, x_in))
|
||||
# add fps_embedding/motion_embedding patches
|
||||
emb_patches = special_model.model.get_fancyvideo_emb_patches(dtype=x_in.dtype, device=x_in.device)
|
||||
transformer_patches = model_options["transformer_options"].get("patches", {})
|
||||
transformer_patches["emb_patch"] = emb_patches
|
||||
model_options["transformer_options"]["patches"] = transformer_patches
|
||||
if special_model.clone_base_uuid not in clone_uuids:
|
||||
clone_uuids.add(special_model.clone_base_uuid)
|
||||
conds = get_conds_with_c_concat(conds,
|
||||
attachment.get_fancy_c_concat(model, x_in))
|
||||
# add fps_embedding/motion_embedding patches
|
||||
emb_patches = special_model.model.get_fancyvideo_emb_patches(dtype=x_in.dtype, device=x_in.device)
|
||||
transformer_patches = model_options["transformer_options"].get("patches", {})
|
||||
transformer_patches["emb_patch"] = emb_patches
|
||||
model_options["transformer_options"]["patches"] = transformer_patches
|
||||
return conds
|
||||
|
||||
def restore_special_model_features(self, model: BaseModel):
|
||||
if self.motion_models is not None:
|
||||
special_models = self.motion_models.get_special_models()
|
||||
def restore_special_model_features(self):
|
||||
for device, motion_models in self.motion_models_devices.items():
|
||||
model: BaseModel = self.model_patcher_devices[device].model
|
||||
special_models = motion_models.get_special_models()
|
||||
if len(special_models) > 0:
|
||||
for special_model in reversed(special_models):
|
||||
attachment = get_mm_attachment(special_model)
|
||||
@@ -109,13 +134,17 @@ class AnimateDiffGlobalState:
|
||||
self.total_steps: int = 0
|
||||
self.callback_output_dict.clear()
|
||||
self.callback_output_dict = {}
|
||||
# model patchers
|
||||
self.model_patcher_devices.clear()
|
||||
if self.model_patcher is not None:
|
||||
self.model_patcher.clean_hooks()
|
||||
del self.model_patcher
|
||||
self.model_patcher = None
|
||||
# motion models
|
||||
self.motion_models_devices.clear()
|
||||
if self.motion_models is not None:
|
||||
del self.motion_models
|
||||
self.motion_models = None
|
||||
# other
|
||||
if self.params is not None:
|
||||
self.params.context_options.reset()
|
||||
del self.params
|
||||
@@ -191,12 +220,13 @@ def _apply_model_wrapper(executor, *args, **kwargs):
|
||||
cond_or_uncond = transformer_options["cond_or_uncond"]
|
||||
ad_params = transformer_options["ad_params"]
|
||||
ADGS: AnimateDiffGlobalState = transformer_options["ADGS"]
|
||||
if ADGS.motion_models is not None:
|
||||
for motion_model in ADGS.motion_models.models:
|
||||
attachment = get_mm_attachment(motion_model)
|
||||
attachment.prepare_alcmi2v_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params, latent_format=executor.class_obj.latent_format)
|
||||
attachment.prepare_camera_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params)
|
||||
attachment.prepare_motionctrl_camera(motion_model, x=x, transformer_options=transformer_options)
|
||||
motion_models = ADGS.motion_models_devices.get(x.device, None)
|
||||
if motion_models is not None:
|
||||
for motion_model in motion_models.models:
|
||||
attachment = get_mm_attachment(motion_model)
|
||||
attachment.prepare_alcmi2v_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params, latent_format=executor.class_obj.latent_format)
|
||||
attachment.prepare_camera_features(motion_model, x=x, cond_or_uncond=cond_or_uncond, ad_params=ad_params)
|
||||
attachment.prepare_motionctrl_camera(motion_model, x=x, transformer_options=transformer_options)
|
||||
del x
|
||||
return executor(*args, **kwargs)
|
||||
|
||||
@@ -545,7 +575,7 @@ def outer_sample_wrapper(executor: WrapperExecutor, *args, **kwargs):
|
||||
|
||||
def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, cond_scale, model_options: dict={}, seed=None):
|
||||
ADGS: AnimateDiffGlobalState = model_options["transformer_options"]["ADGS"]
|
||||
ADGS.initialize(model)
|
||||
ADGS.initialize(model, model_options)
|
||||
ADGS.prepare_current_keyframes(x=x, timestep=timestep, transformer_options=model_options["transformer_options"])
|
||||
try:
|
||||
# add AD/evolved-sampling params to model_options (transformer_options)
|
||||
@@ -556,7 +586,7 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond,
|
||||
model_options["transformer_options"] = model_options["transformer_options"].copy()
|
||||
model_options["transformer_options"]["ad_params"] = ADGS.create_exposed_params()
|
||||
|
||||
cond, uncond = ADGS.perform_special_model_features(model, [cond, uncond], x, model_options)
|
||||
cond, uncond = ADGS.perform_special_model_features([cond, uncond], x, model_options)
|
||||
|
||||
# only use cfg1_optimization if not using custom_cfg or explicitly set to 1.0
|
||||
uncond_ = uncond
|
||||
@@ -576,7 +606,7 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond,
|
||||
|
||||
return comfy.samplers.cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_options, cond, uncond)
|
||||
finally:
|
||||
ADGS.restore_special_model_features(model)
|
||||
ADGS.restore_special_model_features()
|
||||
|
||||
|
||||
def perform_image_injection(ADGS: AnimateDiffGlobalState, model: BaseModel, latents: Tensor, to_inject: NoisedImageToInject) -> Tensor:
|
||||
@@ -634,9 +664,10 @@ def prepare_control_objects(control: ControlBase, full_idxs: list[int], ADGS: An
|
||||
if not hasattr(control, "sub_idxs"):
|
||||
raise ValueError(f"Control type {type(control).__name__} may not support required features for sliding context window; \
|
||||
use ControlNet nodes from Kosinkadink/ComfyUI-Advanced-ControlNet, or make sure ComfyUI-Advanced-ControlNet is updated.")
|
||||
control.sub_idxs = full_idxs
|
||||
control.full_latent_length = ADGS.params.full_length
|
||||
control.context_length = ADGS.params.context_options.context_length
|
||||
if not hasattr(control, "ACN_VERSION"):
|
||||
control.sub_idxs = full_idxs
|
||||
control.full_latent_length = ADGS.params.full_length
|
||||
control.context_length = ADGS.params.context_options.context_length
|
||||
|
||||
# initial sliding_calc_conds_batch inspired by ashen's initial hack for 16-frame sliding context:
|
||||
# https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master
|
||||
@@ -771,16 +802,16 @@ def sliding_calc_cond_batch(executor: Callable, model, conds: list[list[dict]],
|
||||
return conds_final
|
||||
|
||||
ContextResults = collections.namedtuple("ContextResults", ['window_idx', 'sub_conds_out', 'sub_conds', 'ctx_idxs'])
|
||||
def evaluate_context_windows(executor, model, x_in, conds, timestep, enumerated_context_windows: list[tuple[int, list[int]]],
|
||||
def evaluate_context_windows(executor, model: BaseModel, x_in: Tensor, conds, timestep: Tensor, enumerated_context_windows: list[tuple[int, list[int]]],
|
||||
model_options, CREF: ContextRefHandler, ADGS: AnimateDiffGlobalState):
|
||||
results: list[ContextResults] = []
|
||||
for window_idx, ctx_idxs in enumerated_context_windows:
|
||||
# allow processing to end between context window executions for faster Cancel
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
ADGS.params.sub_idxs = ctx_idxs
|
||||
if ADGS.motion_models is not None:
|
||||
ADGS.motion_models.set_sub_idxs(ctx_idxs)
|
||||
ADGS.motion_models.set_video_length(len(ctx_idxs), ADGS.params.full_length)
|
||||
for motion_models in ADGS.motion_models_devices.values():
|
||||
motion_models.set_sub_idxs(ctx_idxs)
|
||||
motion_models.set_video_length(len(ctx_idxs), ADGS.params.full_length)
|
||||
# update exposed params
|
||||
model_options["transformer_options"]["ad_params"]["sub_idxs"] = ctx_idxs
|
||||
model_options["transformer_options"]["ad_params"]["context_length"] = len(ctx_idxs)
|
||||
@@ -800,7 +831,6 @@ def combine_context_window_results(x_in: Tensor, sub_conds_out, sub_conds, ctx_i
|
||||
ADGS: AnimateDiffGlobalState, NAIVE: NaiveReuseHandler, CREF: ContextRefHandler,
|
||||
conds_final: list[Tensor], counts_final: list[Tensor], biases_final: list[Tensor]):
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
full_length = ADGS.params.full_length
|
||||
for pos, idx in enumerate(ctx_idxs):
|
||||
# bias is the influence of a specific index in relation to the whole context window
|
||||
bias = 1 - abs(idx - (ctx_idxs[0] + ctx_idxs[-1]) / 2) / ((ctx_idxs[-1] - ctx_idxs[0] + 1e-2) / 2)
|
||||
|
||||
Reference in New Issue
Block a user