diff --git a/animatediff/context.py b/animatediff/context.py index 8ba6dcf..df116d1 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -604,9 +604,14 @@ def draw_view(window: list[int], gd: GridDisplay): draw_subidxs(window=window, gd=gd, y_grid_offset=2, color=gd.vs.view_color) -def generate_context_visualization(context_opts: ContextOptionsGroup, model: ModelPatcher, sampler_name: str=None, scheduler: str=None, +def generate_context_visualization(model: ModelPatcher, context_opts: ContextOptionsGroup=None, sampler_name: str=None, scheduler: str=None, width=1440, height=200, video_length=32, steps=None, start_step=None, end_step=None, sigmas=None, force_full_denoise=False, denoise=None): + if context_opts is None: + context_opts = ContextOptionsGroup.default() + params = model.get_attachment("ADE_params") + if params is not None: + context_opts = params.context_options context_opts = context_opts.clone() vs = VisualizeSettings(width, video_length) all_imgs = [] @@ -642,7 +647,9 @@ def generate_context_visualization(context_opts: ContextOptionsGroup, model: Mod # check if context should even be active in this case context_active = True - if video_length < context_opts.context_length: + if context_opts.context_length is None: + context_active = False + elif video_length < context_opts.context_length: context_active = False elif video_length == context_opts.context_length and not context_opts.use_on_equal_length: context_active = False diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 44eff90..6332471 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -110,14 +110,14 @@ class ModelPatcherHelper: def get_sample_settings(self) -> SampleSettings: - return self.model.attachments.get(self.SAMPLE_SETTINGS, None) + return self.model.get_attachment(self.SAMPLE_SETTINGS) def set_sample_settings(self, sample_settings: SampleSettings): self.model.set_attachments(self.SAMPLE_SETTINGS, sample_settings) def get_params(self) -> 'InjectionParams': - return self.model.attachments.get(self.PARAMS) + return self.model.get_attachment(self.PARAMS) def set_params(self, params: 'InjectionParams'): self.model.set_attachments(self.PARAMS, params) diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index f7bc6d8..92babc4 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -362,11 +362,11 @@ class VisualizeContextOptionsKAdv: return { "required": { "model": ("MODEL",), - "context_opts": ("CONTEXT_OPTIONS",), "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), }, "optional": { + "context_opts": ("CONTEXT_OPTIONS",), "visual_width": ("INT", {"min": 32, "max": MAX_RESOLUTION, "default": 1440}), "latents_length": ("INT", {"min": 1, "max": BIGMAX, "default": 32}), "steps": ("INT", {"min": 0, "max": BIGMAX, "default": 20}), @@ -379,9 +379,9 @@ class VisualizeContextOptionsKAdv: CATEGORY = "Animate Diff 🎭🅐🅓/context opts/visualize" FUNCTION = "visualize" - def visualize(self, model: ModelPatcher, context_opts: ContextOptionsGroup, sampler_name: str, scheduler: str, - visual_width: 1280, latents_length=32, steps=20, start_step=0, end_step=20): - images = generate_context_visualization(context_opts=context_opts, model=model, width=visual_width, video_length=latents_length, + def visualize(self, model: ModelPatcher, sampler_name: str, scheduler: str, context_opts: ContextOptionsGroup=None, + visual_width=1440, latents_length=32, steps=20, start_step=0, end_step=20): + images = generate_context_visualization(model=model, context_opts=context_opts, width=visual_width, video_length=latents_length, sampler_name=sampler_name, scheduler=scheduler, steps=steps, start_step=start_step, end_step=end_step) return (images,) @@ -393,11 +393,11 @@ class VisualizeContextOptionsK: return { "required": { "model": ("MODEL",), - "context_opts": ("CONTEXT_OPTIONS",), "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), }, "optional": { + "context_opts": ("CONTEXT_OPTIONS",), "visual_width": ("INT", {"min": 32, "max": MAX_RESOLUTION, "default": 1440}), "latents_length": ("INT", {"min": 1, "max": BIGMAX, "default": 32}), "steps": ("INT", {"min": 0, "max": BIGMAX, "default": 20}), @@ -409,9 +409,9 @@ class VisualizeContextOptionsK: CATEGORY = "Animate Diff 🎭🅐🅓/context opts/visualize" FUNCTION = "visualize" - def visualize(self, model: ModelPatcher, context_opts: ContextOptionsGroup, sampler_name: str, scheduler: str, - visual_width: 1280, latents_length=32, steps=20, denoise=1.0): - images = generate_context_visualization(context_opts=context_opts, model=model, width=visual_width, video_length=latents_length, + def visualize(self, model: ModelPatcher, sampler_name: str, scheduler: str, context_opts: ContextOptionsGroup=None, + visual_width=1440, latents_length=32, steps=20, denoise=1.0): + images = generate_context_visualization(model=model, context_opts=context_opts, width=visual_width, video_length=latents_length, sampler_name=sampler_name, scheduler=scheduler, steps=steps, denoise=denoise) return (images,) @@ -423,10 +423,10 @@ class VisualizeContextOptionsSCustom: return { "required": { "model": ("MODEL",), - "context_opts": ("CONTEXT_OPTIONS",), "sigmas": ("SIGMAS", ), }, "optional": { + "context_opts": ("CONTEXT_OPTIONS",), "visual_width": ("INT", {"min": 32, "max": MAX_RESOLUTION, "default": 1440}), "latents_length": ("INT", {"min": 1, "max": BIGMAX, "default": 32}), } @@ -436,8 +436,8 @@ class VisualizeContextOptionsSCustom: CATEGORY = "Animate Diff 🎭🅐🅓/context opts/visualize" FUNCTION = "visualize" - def visualize(self, model: ModelPatcher, context_opts: ContextOptionsGroup, sigmas, - visual_width: 1280, latents_length=32): - images = generate_context_visualization(context_opts=context_opts, model=model, width=visual_width, video_length=latents_length, + def visualize(self, model: ModelPatcher, sigmas, context_opts: ContextOptionsGroup=None, + visual_width=1440, latents_length=32): + images = generate_context_visualization(model=model, context_opts=context_opts, width=visual_width, video_length=latents_length, sigmas=sigmas) return (images,)