22 Commits
Author SHA1 Message Date
kongwykon 08aad38986 fix: compatible with 3.9 python syntax 2023-12-18 02:48:24 +07:00
Tung Nguyen e132c43556 improve preview 2023-12-06 14:50:57 +07:00
Tung Nguyen e6909ae6b2 update forward_timestep_embed following changes from ComfyUI 2023-12-02 10:34:58 +07:00
Ken Simpson 4dc7fa8f39 Fix for https://github.com/comfyanonymous/ComfyUI/issues/2038 2023-12-02 09:43:44 +07:00
Tung Nguyen 77dd7dbb91 improve preview code 2023-11-30 10:35:23 +07:00
Tung Nguyen 86401148f7 fix sampling issue 2023-11-17 16:48:06 +07:00
Tung Nguyen bce8450e07 only run animatediff replated code on animatediff node 2023-11-14 13:18:54 +07:00
Nuked aed11f1196 Fixed code to avoid conflict with (#64)
Fixed code to avoid conflict with other nodes. This code will be executed only inside the AnimateDiffCombine node now.
2023-11-14 12:52:57 +07:00
Tung Nguyen 97404944f5 remove maximum_batch_area reference 2023-11-14 09:14:55 +07:00
Tung Nguyen 3619aee188 update sampling function to match new ComfyUI changes 2023-11-14 04:50:40 +07:00
ArtVenture f1e326ad4a Add LICENSE 2023-11-10 09:21:24 +07:00
Tung Nguyen a62ca4222e update to fix new issue with recent changes from ComfyUI 2023-11-07 15:55:40 +07:00
Tung Nguyen 45858eebc4 fix LoadVideo node not resize when load video 2023-10-25 22:39:39 +07:00
Tung Nguyen 32aad09a9a update get_resized_cond to match new cond format 2023-10-25 22:12:20 +07:00
Tung Nguyen f40207480f Merge branch 'main' of https://github.com/ArtVentureX/comfyui-animatediff 2023-10-25 22:11:22 +07:00
Tung Nguyen (Blockchain) 6416ffafa2 update sampling function to match new ComfyUI update (#58) 2023-10-25 16:40:39 +07:00
Tung Nguyen 49916eae86 update sampling function to match new ComfyUI update 2023-10-25 16:30:06 +07:00
Tung Nguyen dccfb2326e fix: new comfyui break sampling 2023-10-25 15:24:50 +07:00
Tung Nguyen 62569f0184 fix VanillaTemporalModule.forward() missing 1 required positional argument: 'encoder_hidden_states' 2023-10-12 21:46:57 +07:00
Tung Nguyen dc6268aea5 update to match new changes from ComfyUI 2023-10-12 21:44:21 +07:00
Tung Nguyen 9d32153349 increase max alpha for motion lora 2023-09-25 22:59:21 +07:00
Tung Nguyen (Blockchain) 63c068c59f Support motion LoRA (#38)
Support new motion LoRa from AnimateDiff
2023-09-25 22:53:44 +07:00
13 changed files with 1333 additions and 510 deletions
Vendored
BIN
View File
Binary file not shown.
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+42
View File
@@ -11,6 +11,48 @@
- Community modules: [manshoety/AD_Stabilized_Motion](https://huggingface.co/manshoety/AD_Stabilized_Motion) | [CiaraRowles/TemporalDiff](https://huggingface.co/CiaraRowles/TemporalDiff) - Community modules: [manshoety/AD_Stabilized_Motion](https://huggingface.co/manshoety/AD_Stabilized_Motion) | [CiaraRowles/TemporalDiff](https://huggingface.co/CiaraRowles/TemporalDiff)
- AnimateDiff v2 [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt) - AnimateDiff v2 [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt)
## Update 2023/09/25
#### **Motion LoRA** is now supported!
Download [motion LoRAs](https://huggingface.co/guoyww/animatediff/tree/main) and put them under `comfyui-animatediff/loras/` folder.
Note: LoRAs only work with **AnimateDiff v2** [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt) module.
#### New node: `AnimateDiffLoraLoader`
<img width="370" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/7a9f62f7-702e-48a4-934c-bbfe1e23aff2">
Example workflow:
<img width="1280" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/93e7550f-4648-4482-9961-6cece5132dc9">
Workflow: [lora.json](https://github.com/ArtVentureX/comfyui-animatediff/blob/main/workflows/lora.json)
Samples:
<table>
<tr>
<td>
<img width="512" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/2c5aa25e-0682-481f-8842-066c5b988864">
</td>
</tr>
<tr>
<td>
<img width="512" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/adfbad45-3ba5-42e3-9bee-d2b83f43989c">
</td>
</tr>
<tr>
<td>
<img width="512" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/8e484c74-c691-4d1c-9514-719dbfe3a0b5">
</td>
</tr>
<tr>
<td>
<img width="512" alt="image" src="https://github.com/ArtVentureX/comfyui-animatediff/assets/133728487/4921a335-9207-4a7b-9d66-61a5d76e3179">
</td>
</tr>
</table>
## Update 2023/09/21 ## Update 2023/09/21
#### **Sliding Window** is now available! #### **Sliding Window** is now available!
+45 -1
View File
@@ -1,5 +1,6 @@
import os import os
import hashlib import hashlib
import torch
from typing import Dict from typing import Dict
import folder_paths import folder_paths
@@ -11,6 +12,7 @@ from .motion_module import MotionWrapper
motion_modules: Dict[str, MotionWrapper] = {} motion_modules: Dict[str, MotionWrapper] = {}
motion_loras: Dict[str, Dict[str, torch.Tensor]] = {}
folder_paths.folder_names_and_paths["AnimateDiff"] = ( folder_paths.folder_names_and_paths["AnimateDiff"] = (
@@ -20,19 +22,34 @@ folder_paths.folder_names_and_paths["AnimateDiff"] = (
], ],
folder_paths.supported_pt_extensions, folder_paths.supported_pt_extensions,
) )
folder_paths.folder_names_and_paths["AnimateDiffLora"] = (
[
os.path.join(folder_paths.models_dir, "AnimateDiffLora"),
os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "loras"),
],
folder_paths.supported_pt_extensions,
)
def get_available_models(): def get_available_models():
return folder_paths.get_filename_list("AnimateDiff") return folder_paths.get_filename_list("AnimateDiff")
def get_available_loras():
return folder_paths.get_filename_list("AnimateDiffLora")
def get_model_path(model_name): def get_model_path(model_name):
return folder_paths.get_full_path("AnimateDiff", model_name) return folder_paths.get_full_path("AnimateDiff", model_name)
def get_lora_path(lora_name):
return folder_paths.get_full_path("AnimateDiffLora", lora_name)
def get_model_hash(file_path): def get_model_hash(file_path):
with open(file_path, "rb") as f: with open(file_path, "rb") as f:
bytes = f.read() # read entire file as bytes bytes = f.read(1024 * 1024) # read entire file as bytes
return hashlib.sha256(bytes).hexdigest() return hashlib.sha256(bytes).hexdigest()
@@ -54,3 +71,30 @@ def load_motion_module(model_name: str):
motion_modules[model_hash] = motion_module motion_modules[model_hash] = motion_module
return motion_modules[model_hash] return motion_modules[model_hash]
def load_lora(lora_name: str):
lora_path = get_lora_path(lora_name)
lora_hash = get_model_hash(lora_path)
if lora_hash not in motion_modules:
logger.info(f"Loading lora {lora_name}")
state_dict = load_torch_file(lora_path)
updated_state_dict: Dict[str, torch.Tensor] = {}
for key in state_dict:
# only process lora down key
if "up." in key:
continue
up_key = key.replace(".down.", ".up.")
model_key = key.replace("processor.", "").replace("_lora", "").replace("down.", "").replace("up.", "")
model_key = model_key.replace("to_out.", "to_out.0.")
combined_key = ".".join(model_key.split(".")[:-1])
weight_down = state_dict[key]
weight_up = state_dict[up_key]
updated_state_dict[combined_key] = torch.mm(weight_up, weight_down).to("cpu")
motion_loras[lora_hash] = updated_state_dict
return motion_loras[lora_hash]
+31 -8
View File
@@ -6,25 +6,30 @@ from einops import rearrange, repeat
import comfy.model_management as model_management import comfy.model_management as model_management
from comfy.ldm.modules.attention import ( from comfy.ldm.modules.attention import (
default,
FeedForward, FeedForward,
CrossAttention as ComfyCrossAttention, CrossAttention as ComfyCrossAttention,
CrossAttentionDoggettx, attention_basic,
CrossAttentionBirchSan, attention_pytorch,
attention_split,
attention_sub_quad,
) )
from comfy.cli_args import args from comfy.cli_args import args
from .logger import logger from .logger import logger
CrossAttention = ComfyCrossAttention attention = attention_basic
if model_management.xformers_enabled(): if model_management.xformers_enabled():
logger.warn("xformers is enabled but it has a bug that can cause issue while using with AnimateDiff.") logger.warn("xformers is enabled but it has a bug that can cause issue while using with AnimateDiff.")
if model_management.pytorch_attention_enabled():
attention = attention_pytorch
else:
if args.use_split_cross_attention: if args.use_split_cross_attention:
logger.warn("Using split optimization for AnimateDiff cross attention instead.") attention = attention_split
CrossAttention = CrossAttentionDoggettx
else: else:
logger.warn("Using sub quadratic optimization for AnimateDiff cross attention instead.") attention = attention_sub_quad
CrossAttention = CrossAttentionBirchSan
def zero_module(module): def zero_module(module):
@@ -51,6 +56,24 @@ def has_mid_block(mm_state_dict: dict[str, Tensor]):
return False return False
class CrossAttention(ComfyCrossAttention):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def forward(self, x, context=None, value=None, mask=None):
q = self.to_q(x)
context = default(context, x)
k = self.to_k(context)
if value is not None:
v = self.to_v(value)
del value
else:
v = self.to_v(context)
out = attention(q, k, v, self.heads, mask)
return self.to_out(out)
class MotionWrapper(nn.Module): class MotionWrapper(nn.Module):
def __init__(self, mm_type: str, encoding_max_len: int = 24, is_v2=False): def __init__(self, mm_type: str, encoding_max_len: int = 24, is_v2=False):
super().__init__() super().__init__()
@@ -156,7 +179,7 @@ class VanillaTemporalModule(nn.Module):
def set_video_length(self, video_length: int): def set_video_length(self, video_length: int):
self.temporal_transformer.set_video_length(video_length) self.temporal_transformer.set_video_length(video_length)
def forward(self, input_tensor, encoder_hidden_states, attention_mask=None): def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None):
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)
+84 -2
View File
@@ -3,16 +3,18 @@ import json
import torch import torch
import numpy as np import numpy as np
import hashlib import hashlib
from typing import List from typing import List, Dict, Tuple
from torch import Tensor from torch import Tensor
from PIL import Image, ImageSequence from PIL import Image, ImageSequence
from PIL.PngImagePlugin import PngInfo from PIL.PngImagePlugin import PngInfo
import folder_paths import folder_paths
from .model_utils import get_available_models, load_motion_module from .motion_module import MotionWrapper
from .model_utils import get_available_models, load_motion_module, get_available_loras, load_lora
from .utils import pil2tensor, ensure_opencv from .utils import pil2tensor, ensure_opencv
from .sampler import AnimateDiffSampler, AnimateDiffSlidingWindowOptions from .sampler import AnimateDiffSampler, AnimateDiffSlidingWindowOptions
from .logger import logger
SLIDING_CONTEXT_LENGTH = 16 SLIDING_CONTEXT_LENGTH = 16
@@ -28,21 +30,99 @@ class AnimateDiffModuleLoader:
"required": { "required": {
"model_name": (get_available_models(),), "model_name": (get_available_models(),),
}, },
"optional": {
"lora_stack": ("MOTION_LORA_STACK",),
},
} }
RETURN_TYPES = ("MOTION_MODULE",) RETURN_TYPES = ("MOTION_MODULE",)
CATEGORY = "Animate Diff" CATEGORY = "Animate Diff"
FUNCTION = "load_motion_module" FUNCTION = "load_motion_module"
def inject_loras(self, motion_module: MotionWrapper, lora_stack: List[Tuple[Dict[str, Tensor], float]]):
for lora in lora_stack:
(state_dict, alpha) = lora
for key in state_dict:
layer_infos = key.split(".")
curr_layer = motion_module
while len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
curr_layer = curr_layer.__getattr__(temp_name)
curr_layer.weight.data += alpha * state_dict[key].to(curr_layer.weight.data.device)
def eject_loras(self, motion_module: MotionWrapper, lora_stack: List[Tuple[float, Dict[str, Tensor]]]):
lora_stack.reverse() # should not matter but just in case
for lora in lora_stack:
(state_dict, alpha) = lora
for key in state_dict:
layer_infos = key.split(".")
curr_layer = motion_module
while len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
curr_layer = curr_layer.__getattr__(temp_name)
curr_layer.weight.data -= alpha * state_dict[key].to(curr_layer.weight.data.device)
def load_motion_module( def load_motion_module(
self, self,
model_name: str, model_name: str,
lora_stack: List = None,
): ):
motion_module = load_motion_module(model_name) motion_module = load_motion_module(model_name)
# inject loras
if motion_module.is_v2:
if hasattr(motion_module, "lora_stack") and isinstance(motion_module.lora_stack, list):
self.eject_loras(motion_module, motion_module.lora_stack)
delattr(motion_module, "lora_stack")
if isinstance(lora_stack, list):
self.inject_loras(motion_module, lora_stack)
setattr(motion_module, "lora_stack", lora_stack)
elif isinstance(lora_stack, list):
logger.warning("LoRA is provided but only motion module v2 is supported.")
return (motion_module,) return (motion_module,)
class AnimateDiffLoraLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora_name": (get_available_loras(),),
"alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
},
"optional": {
"lora_stack": ("MOTION_LORA_STACK",),
},
}
RETURN_TYPES = ("MOTION_LORA_STACK",)
CATEGORY = "Animate Diff"
FUNCTION = "load_lora"
def load_lora(
self,
lora_name: str,
alpha: float,
lora_stack: List = None,
):
if not lora_stack:
lora_stack = []
lora = load_lora(lora_name)
lora_stack.append((lora, alpha))
return (lora_stack,)
class AnimateDiffCombine: class AnimateDiffCombine:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -324,6 +404,7 @@ class ImageChunking:
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"AnimateDiffModuleLoader": AnimateDiffModuleLoader, "AnimateDiffModuleLoader": AnimateDiffModuleLoader,
"AnimateDiffLoraLoader": AnimateDiffLoraLoader,
"AnimateDiffCombine": AnimateDiffCombine, "AnimateDiffCombine": AnimateDiffCombine,
"AnimateDiffSampler": AnimateDiffSampler, "AnimateDiffSampler": AnimateDiffSampler,
"AnimateDiffSlidingWindowOptions": AnimateDiffSlidingWindowOptions, "AnimateDiffSlidingWindowOptions": AnimateDiffSlidingWindowOptions,
@@ -332,6 +413,7 @@ NODE_CLASS_MAPPINGS = {
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"AnimateDiffModuleLoader": "Animate Diff Module Loader", "AnimateDiffModuleLoader": "Animate Diff Module Loader",
"AnimateDiffLoraLoader": "Animate Diff Lora Loader",
"AnimateDiffSampler": "Animate Diff Sampler", "AnimateDiffSampler": "Animate Diff Sampler",
"AnimateDiffSlidingWindowOptions": "Sliding Window Options", "AnimateDiffSlidingWindowOptions": "Sliding Window Options",
"AnimateDiffCombine": "Animate Diff Combine", "AnimateDiffCombine": "Animate Diff Combine",
+17 -35
View File
@@ -4,9 +4,7 @@ from torch.nn.functional import group_norm
from einops import rearrange from einops import rearrange
import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel
import comfy.model_management as model_management from comfy.model_base import BaseModel, model_sampling
from comfy.model_base import BaseModel
from comfy.ldm.modules.attention import SpatialTransformer
from nodes import KSampler from nodes import KSampler
from .logger import logger from .logger import logger
@@ -18,19 +16,19 @@ from .sliding_context_sampling import SlidingContext, inject_sampling_function,
SLIDING_CONTEXT_LENGTH = 16 SLIDING_CONTEXT_LENGTH = 16
def forward_timestep_embed(ts, x, emb, context=None, transformer_options={}, output_shape=None): class ModelSamplingConfig:
def __init__(self, beta_schedule: str):
self.sampling_settings = {}
self.sampling_settings["beta_schedule"] = beta_schedule
def forward_timestep_embed(ts, x, emb, context=None, *args, **kwargs):
for layer in ts: for layer in ts:
if isinstance(layer, openaimodel.TimestepBlock): if isinstance(layer, VanillaTemporalModule):
x = layer(x, emb)
elif isinstance(layer, VanillaTemporalModule):
x = layer(x, context) x = layer(x, context)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context, transformer_options)
transformer_options["current_index"] += 1
elif isinstance(layer, openaimodel.Upsample):
x = layer(x, output_shape=output_shape)
else: else:
x = layer(x) x = orig_forward_timestep_embed([layer], x, emb, context, *args, **kwargs)
return x return x
@@ -49,7 +47,6 @@ def groupnorm_mm_factory(video_length: int):
orig_forward_timestep_embed = openaimodel.forward_timestep_embed orig_forward_timestep_embed = openaimodel.forward_timestep_embed
orig_maximum_batch_area = model_management.maximum_batch_area
orig_groupnorm_forward = torch.nn.GroupNorm.forward orig_groupnorm_forward = torch.nn.GroupNorm.forward
@@ -181,32 +178,17 @@ class AnimateDiffSampler(KSampler):
def __init__(self) -> None: def __init__(self) -> None:
super().__init__() super().__init__()
self.prev_beta = None self.model_sampling = None
self.prev_linear_start = None
self.prev_linear_end = None
def override_beta_schedule(self, model: BaseModel): def override_beta_schedule(self, model: BaseModel):
self.prev_beta = model.get_buffer("betas").cpu().clone().detach() self.model_sampling = model.model_sampling
self.prev_linear_start = model.linear_start model.model_sampling = model_sampling(
self.prev_linear_end = model.linear_end ModelSamplingConfig(beta_schedule="sqrt_linear"), model_type=model.model_type
model.register_schedule(
given_betas=None,
beta_schedule="sqrt_linear",
timesteps=1000,
linear_start=0.00085,
linear_end=0.012,
cosine_s=8e-3,
) )
def restore_beta_schedule(self, model: BaseModel): def restore_beta_schedule(self, model: BaseModel):
model.register_schedule( model.model_sampling = self.model_sampling
given_betas=self.prev_beta, self.model_sampling = None
linear_start=self.prev_linear_start,
linear_end=self.prev_linear_end,
)
self.prev_beta = None
self.prev_linear_start = None
self.prev_linear_end = None
def inject_motion_module(self, model, motion_module: MotionWrapper, inject_method: str, frame_number: int): def inject_motion_module(self, model, motion_module: MotionWrapper, inject_method: str, frame_number: int):
model = model.clone() model = model.clone()
+83 -132
View File
@@ -1,6 +1,7 @@
import math
import torch import torch
from torch import Tensor from torch import Tensor
import math from typing import List, Dict
import comfy.utils import comfy.utils
import comfy.sample import comfy.sample
@@ -17,6 +18,10 @@ orig_comfy_sample = comfy.sample.sample
orig_sampling_function = comfy_samplers.sampling_function orig_sampling_function = comfy_samplers.sampling_function
def lcm(a, b):
return abs(a * b) // math.gcd(a, b)
class SlidingContext: class SlidingContext:
def __init__( def __init__(
self, self,
@@ -70,47 +75,34 @@ def __sliding_sample_factory(ctx: SlidingContext):
ctx.current_step = start_step + step + 1 ctx.current_step = start_step + step + 1
try: return orig_comfy_sample(model, *args, **kwargs, callback=callback)
return orig_comfy_sample(model, *args, **kwargs, callback=callback)
except RuntimeError as e:
if str(e).startswith("CUDA error: invalid configuration argument"):
raise RuntimeError(
f"An xformers bug was encountered in AnimateDiff - to run your workflow, \
disable xformers in ComfyUI using '--disable-xformers' startup argument."
)
raise
def sampling_function( def sampling_function(model, x, timestep, uncond, cond, cond_scale, model_options={}, seed=None):
model_function, x, timestep, uncond, cond, cond_scale, cond_concat=None, model_options={}, seed=None def get_area_and_mult(conds, x_in, timestep_in):
):
def get_area_and_mult(cond, x_in, cond_concat_in, timestep_in):
area = (x_in.shape[2], x_in.shape[3], 0, 0) area = (x_in.shape[2], x_in.shape[3], 0, 0)
strength = 1.0 strength = 1.0
if "timestep_start" in cond[1]:
timestep_start = cond[1]["timestep_start"] if "timestep_start" in conds:
timestep_start = conds["timestep_start"]
if timestep_in[0] > timestep_start: if timestep_in[0] > timestep_start:
return None return None
if "timestep_end" in cond[1]: if "timestep_end" in conds:
timestep_end = cond[1]["timestep_end"] timestep_end = conds["timestep_end"]
if timestep_in[0] < timestep_end: if timestep_in[0] < timestep_end:
return None return None
if "area" in cond[1]: if "area" in conds:
area = cond[1]["area"] area = conds["area"]
if "strength" in cond[1]: if "strength" in conds:
strength = cond[1]["strength"] strength = conds["strength"]
adm_cond = None
if "adm_encoded" in cond[1]:
adm_cond = cond[1]["adm_encoded"]
input_x = x_in[:, :, area[2] : area[0] + area[2], area[3] : area[1] + area[3]] input_x = x_in[:, :, area[2] : area[0] + area[2], area[3] : area[1] + area[3]]
if "mask" in cond[1]: if "mask" in conds:
# Scale the mask to the size of the input # Scale the mask to the size of the input
# The mask should have been resized as we began the sampling process # The mask should have been resized as we began the sampling process
mask_strength = 1.0 mask_strength = 1.0
if "mask_strength" in cond[1]: if "mask_strength" in conds:
mask_strength = cond[1]["mask_strength"] mask_strength = conds["mask_strength"]
mask = cond[1]["mask"] mask = conds["mask"]
assert mask.shape[1] == x_in.shape[2] assert mask.shape[1] == x_in.shape[2]
assert mask.shape[2] == x_in.shape[3] assert mask.shape[2] == x_in.shape[3]
mask = mask[:, area[2] : area[0] + area[2], area[3] : area[1] + area[3]] * mask_strength mask = mask[:, area[2] : area[0] + area[2], area[3] : area[1] + area[3]] * mask_strength
@@ -119,7 +111,7 @@ def __sliding_sample_factory(ctx: SlidingContext):
mask = torch.ones_like(input_x) mask = torch.ones_like(input_x)
mult = mask * strength mult = mask * strength
if "mask" not in cond[1]: if "mask" not in conds:
rr = 8 rr = 8
if area[2] != 0: if area[2] != 0:
for t in range(rr): for t in range(rr):
@@ -135,24 +127,17 @@ def __sliding_sample_factory(ctx: SlidingContext):
mult[:, :, :, area[1] - 1 - t : area[1] - t] *= (1.0 / rr) * (t + 1) mult[:, :, :, area[1] - 1 - t : area[1] - t] *= (1.0 / rr) * (t + 1)
conditionning = {} conditionning = {}
conditionning["c_crossattn"] = cond[0] model_conds = conds["model_conds"]
if cond_concat_in is not None and len(cond_concat_in) > 0: for c in model_conds:
cropped = [] conditionning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area)
for x in cond_concat_in:
cr = x[:, :, area[2] : area[0] + area[2], area[3] : area[1] + area[3]]
cropped.append(cr)
conditionning["c_concat"] = torch.cat(cropped, dim=1)
if adm_cond is not None:
conditionning["c_adm"] = adm_cond
control = None control = None
if "control" in cond[1]: if "control" in conds:
control = cond[1]["control"] control = conds["control"]
patches = None patches = None
if "gligen" in cond[1]: if "gligen" in conds:
gligen = cond[1]["gligen"] gligen = conds["gligen"]
patches = {} patches = {}
gligen_type = gligen[0] gligen_type = gligen[0]
gligen_model = gligen[1] gligen_model = gligen[1]
@@ -170,24 +155,8 @@ def __sliding_sample_factory(ctx: SlidingContext):
return True return True
if c1.keys() != c2.keys(): if c1.keys() != c2.keys():
return False return False
if "c_crossattn" in c1: for k in c1:
s1 = c1["c_crossattn"].shape if not c1[k].can_concat(c2[k]):
s2 = c2["c_crossattn"].shape
if s1 != s2:
if s1[0] != s2[0] or s1[2] != s2[2]: # these 2 cases should not happen
return False
mult_min = comfy_samplers.lcm(s1[1], s2[1])
diff = mult_min // min(s1[1], s2[1])
if (
diff > 4
): # arbitrary limit on the padding because it's probably going to impact performance negatively if it's too much
return False
if "c_concat" in c1:
if c1["c_concat"].shape != c2["c_concat"].shape:
return False
if "c_adm" in c1:
if c1["c_adm"].shape != c2["c_adm"].shape:
return False return False
return True return True
@@ -216,55 +185,41 @@ def __sliding_sample_factory(ctx: SlidingContext):
c_concat = [] c_concat = []
c_adm = [] c_adm = []
crossattn_max_len = 0 crossattn_max_len = 0
for x in c_list:
if "c_crossattn" in x:
c = x["c_crossattn"]
if crossattn_max_len == 0:
crossattn_max_len = c.shape[1]
else:
crossattn_max_len = comfy_samplers.lcm(crossattn_max_len, c.shape[1])
c_crossattn.append(c)
if "c_concat" in x:
c_concat.append(x["c_concat"])
if "c_adm" in x:
c_adm.append(x["c_adm"])
out = {}
c_crossattn_out = []
for c in c_crossattn:
if c.shape[1] < crossattn_max_len:
c = c.repeat(1, crossattn_max_len // c.shape[1], 1) # padding with repeat doesn't change result
c_crossattn_out.append(c)
if len(c_crossattn_out) > 0: temp = {}
out["c_crossattn"] = torch.cat(c_crossattn_out) for x in c_list:
if len(c_concat) > 0: for k in x:
out["c_concat"] = torch.cat(c_concat) cur = temp.get(k, [])
if len(c_adm) > 0: cur.append(x[k])
out["c_adm"] = torch.cat(c_adm) temp[k] = cur
out = {}
for k in temp:
conds = temp[k]
out[k] = conds[0].concat(conds[1:])
return out return out
def calc_cond_uncond_batch( def calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options):
model_function, cond, uncond, x_in, timestep, max_total_area, cond_concat_in, model_options
):
out_cond = torch.zeros_like(x_in) out_cond = torch.zeros_like(x_in)
out_count = torch.ones_like(x_in) / 100000.0 out_count = torch.ones_like(x_in) * 1e-37
out_uncond = torch.zeros_like(x_in) out_uncond = torch.zeros_like(x_in)
out_uncond_count = torch.ones_like(x_in) / 100000.0 out_uncond_count = torch.ones_like(x_in) * 1e-37
COND = 0 COND = 0
UNCOND = 1 UNCOND = 1
to_run = [] to_run = []
for x in cond: for x in cond:
p = get_area_and_mult(x, x_in, cond_concat_in, timestep) p = get_area_and_mult(x, x_in, timestep)
if p is None: if p is None:
continue continue
to_run += [(p, COND)] to_run += [(p, COND)]
if uncond is not None: if uncond is not None:
for x in uncond: for x in uncond:
p = get_area_and_mult(x, x_in, cond_concat_in, timestep) p = get_area_and_mult(x, x_in, timestep)
if p is None: if p is None:
continue continue
@@ -281,9 +236,11 @@ def __sliding_sample_factory(ctx: SlidingContext):
to_batch_temp.reverse() to_batch_temp.reverse()
to_batch = to_batch_temp[:1] to_batch = to_batch_temp[:1]
free_memory = model_management.get_free_memory(x_in.device)
for i in range(1, len(to_batch_temp) + 1): for i in range(1, len(to_batch_temp) + 1):
batch_amount = to_batch_temp[: len(to_batch_temp) // i] batch_amount = to_batch_temp[: len(to_batch_temp) // i]
if len(batch_amount) * first_shape[0] * first_shape[2] * first_shape[3] < max_total_area: input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:]
if model.memory_required(input_shape) < free_memory:
to_batch = batch_amount to_batch = batch_amount
break break
@@ -333,11 +290,11 @@ def __sliding_sample_factory(ctx: SlidingContext):
if "model_function_wrapper" in model_options: if "model_function_wrapper" in model_options:
output = model_options["model_function_wrapper"]( output = model_options["model_function_wrapper"](
model_function, model.apply_model,
{"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond},
).chunk(batch_chunks) ).chunk(batch_chunks)
else: else:
output = model_function(input_x, timestep_, **c).chunk(batch_chunks) output = model.apply_model(input_x, timestep_, **c).chunk(batch_chunks)
del input_x del input_x
for o in range(batch_chunks): for o in range(batch_chunks):
@@ -361,14 +318,11 @@ def __sliding_sample_factory(ctx: SlidingContext):
del out_count del out_count
out_uncond /= out_uncond_count out_uncond /= out_uncond_count
del out_uncond_count del out_uncond_count
return out_cond, out_uncond return out_cond, out_uncond
# sliding_calc_cond_uncond_batch inspired by ashen's initial hack for 16-frame sliding context: # sliding_calc_cond_uncond_batch inspired by ashen's initial hack for 16-frame sliding context:
# https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master # https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master
def sliding_calc_cond_uncond_batch( def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options):
model_function, cond, uncond, x_in, timestep, max_total_area, cond_concat_in, model_options
):
# figure out how input is split # figure out how input is split
axes_factor = x.size(0) // ctx.video_length axes_factor = x.size(0) // ctx.video_length
@@ -384,37 +338,29 @@ def __sliding_sample_factory(ctx: SlidingContext):
control.full_latent_length = ctx.video_length control.full_latent_length = ctx.video_length
control.context_length = ctx.context_length control.context_length = ctx.context_length
def get_resized_cond(cond_in, full_idxs) -> list: def get_resized_cond(cond_in: List[Dict], full_idxs) -> list:
# reuse or resize cond items to match context requirements # reuse or resize cond items to match context requirements
resized_cond = [] resized_cond = []
# cond object is a list containing a list - outer list is irrelevant, so just loop through it # cond object is a list containing a list - outer list is irrelevant, so just loop through it
for actual_cond in cond_in: for actual_cond in cond_in:
resized_actual_cond = [] new_cond_item = actual_cond.copy()
# now we are in the inner list - index 0 is tensor, index 1 is dictionary for key, cond_item in new_cond_item.items():
for cond_idx, cond_item in enumerate(actual_cond):
if isinstance(cond_item, Tensor): if isinstance(cond_item, Tensor):
# check that tensor is the expected length - x.size(0) # check that tensor is the expected length - x.size(0)
if cond_item.size(0) == x.size(0): if cond_item.size(0) == x.size(0):
pass
# if so, it's subsetting time - tell controls the expected indeces so they can handle them # if so, it's subsetting time - tell controls the expected indeces so they can handle them
actual_cond_item = cond_item[full_idxs] actual_cond_item = cond_item[full_idxs]
resized_actual_cond.append(actual_cond_item) new_cond_item[key] = actual_cond_item
elif key == "control":
control_item = cond_item
if hasattr(control_item, "sub_idxs"):
prepare_control_objects(control_item, full_idxs)
else: else:
resized_actual_cond.append(cond_item) raise ValueError(
elif isinstance(cond_item, dict): f"Control type {type(control_item).__name__} may not support required features for sliding context window; use Control objects from Kosinkadink/Advanced-ControlNet nodes."
# when in dictionary, look for control )
if "control" in cond_item: new_cond_item[key] = cond_item
control_item = cond_item["control"] resized_cond.append(new_cond_item)
if hasattr(control_item, "sub_idxs"):
prepare_control_objects(control_item, full_idxs)
else:
raise ValueError(
f"Control type {type(control_item).__name__} may not support required features for sliding context window; use Control objects from Kosinkadink/Advanced-ControlNet nodes."
)
resized_actual_cond.append(cond_item)
else:
resized_actual_cond.append(cond_item)
resized_cond.append(resized_actual_cond)
return resized_cond return resized_cond
# perform calc_cond_uncond_batch per context window # perform calc_cond_uncond_batch per context window
@@ -437,16 +383,13 @@ def __sliding_sample_factory(ctx: SlidingContext):
sub_timestep = timestep[full_idxs] sub_timestep = timestep[full_idxs]
sub_cond = get_resized_cond(cond, full_idxs) if cond is not None else None sub_cond = get_resized_cond(cond, full_idxs) if cond is not None else None
sub_uncond = get_resized_cond(uncond, full_idxs) if uncond is not None else None sub_uncond = get_resized_cond(uncond, full_idxs) if uncond is not None else None
sub_cond_concat = get_resized_cond(cond_concat, full_idxs) if cond_concat is not None else None
sub_cond_out, sub_uncond_out = calc_cond_uncond_batch( sub_cond_out, sub_uncond_out = calc_cond_uncond_batch(
model_function, model,
sub_cond, sub_cond,
sub_uncond, sub_uncond,
sub_x, sub_x,
sub_timestep, sub_timestep,
max_total_area,
sub_cond_concat,
model_options, model_options,
) )
@@ -459,17 +402,21 @@ def __sliding_sample_factory(ctx: SlidingContext):
uncond_final /= out_count_final uncond_final /= out_count_final
return cond_final, uncond_final return cond_final, uncond_final
max_total_area = model_management.maximum_batch_area()
if math.isclose(cond_scale, 1.0): if math.isclose(cond_scale, 1.0):
uncond = None uncond = None
cond, uncond = sliding_calc_cond_uncond_batch( cond, uncond = sliding_calc_cond_uncond_batch(model, cond, uncond, x, timestep, model_options)
model_function, cond, uncond, x, timestep, max_total_area, cond_concat, model_options
)
if "sampler_cfg_function" in model_options: if "sampler_cfg_function" in model_options:
args = {"cond": cond, "uncond": uncond, "cond_scale": cond_scale, "timestep": timestep} args = {
return model_options["sampler_cfg_function"](args) "cond": x - cond,
"uncond": x - uncond,
"cond_scale": cond_scale,
"timestep": timestep,
"input": x,
"sigma": timestep,
}
return x - model_options["sampler_cfg_function"](args)
else: else:
return uncond + (cond - uncond) * cond_scale return uncond + (cond - uncond) * cond_scale
@@ -477,6 +424,10 @@ def __sliding_sample_factory(ctx: SlidingContext):
def inject_sampling_function(ctx: SlidingContext): def inject_sampling_function(ctx: SlidingContext):
global orig_comfy_sample, orig_sampling_function
orig_comfy_sample = comfy.sample.sample
orig_sampling_function = comfy_samplers.sampling_function
(sample, sampling_function) = __sliding_sample_factory(ctx) (sample, sampling_function) = __sliding_sample_factory(ctx)
comfy.sample.sample = sample comfy.sample.sample = sample
comfy_samplers.sampling_function = sampling_function comfy_samplers.sampling_function = sampling_function
+8 -10
View File
@@ -142,14 +142,12 @@ def uniform_constant(
# yield if not skipped # yield if not skipped
yield to_yield yield to_yield
def get_context_scheduler(name: str) -> Callable: def get_context_scheduler(name: str) -> Callable:
match name: if name == ContextSchedules.UNIFORM:
case ContextSchedules.UNIFORM: return uniform
return uniform elif name == ContextSchedules.UNIFORM_CONSTANT:
case ContextSchedules.UNIFORM_CONSTANT: return uniform_constant
return uniform_constant elif name == ContextSchedules.UNIFORM_V2:
case ContextSchedules.UNIFORM_V2: return uniform_v2
return uniform_v2 else:
case _: raise ValueError(f"Unknown context_overlap policy {name}")
raise ValueError(f"Unknown context_overlap policy {name}")
View File
+236 -153
View File
@@ -1,162 +1,245 @@
import { app } from "../../../scripts/app.js"; import { app, ANIM_PREVIEW_WIDGET } from '../../../scripts/app.js';
import { api } from "../../../scripts/api.js"; import { api } from "../../../scripts/api.js";
import { $el } from '../../../scripts/ui.js';
import { createImageHost } from "../../../scripts/ui/imagePreview.js"
function offsetDOMWidget(widget, ctx, node, widgetWidth, widgetY, height) { const URL_REGEX = /^(https?:\/\/|\/view\?|data:image\/)/;
const margin = 10;
const elRect = ctx.canvas.getBoundingClientRect();
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(0, widgetY + margin);
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d); const style = `
Object.assign(widget.inputEl.style, { .comfy-img-preview video {
transformOrigin: "0 0", object-fit: contain;
transform: scale, width: var(--comfy-img-preview-width);
left: `${transform.e}px`, height: var(--comfy-img-preview-height);
top: `${transform.d + transform.f}px`, }
width: `${widgetWidth}px`, `;
height: `${(height || widget.parent?.inputHeight || 32) - margin}px`,
position: "absolute", export function chainCallback(object, property, callback) {
background: !node.color ? "" : node.color, if (object == undefined) {
color: !node.color ? "" : "white", //This should not happen.
zIndex: 5, //app.graph._nodes.indexOf(node), console.error("Tried to add callback to non-existant object");
return;
}
if (property in object) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
callback.apply(this, arguments);
return r;
};
} else {
object[property] = callback;
}
};
export function formatUploadedUrl(params) {
if (params.url) {
return params.url;
}
params = { ...params };
if (!params.filename && params.name) {
params.filename = params.name;
delete params.name;
}
return api.apiURL("/view?" + new URLSearchParams(params));
};
export function addVideoPreview(nodeType, options = {}) {
const createVideoNode = (url) => {
return new Promise((cb) => {
const videoEl = document.createElement('video');
Object.defineProperty(videoEl, 'naturalWidth', {
get: () => {
return videoEl.videoWidth;
},
});
Object.defineProperty(videoEl, 'naturalHeight', {
get: () => {
return videoEl.videoHeight;
},
});
videoEl.addEventListener('loadedmetadata', () => {
videoEl.controls = false;
videoEl.loop = true;
videoEl.muted = true;
cb(videoEl);
});
videoEl.addEventListener('error', () => {
cb();
});
videoEl.src = url;
});
};
const createImageNode = (url) => {
return new Promise((cb) => {
const imgEl = document.createElement('img');
imgEl.onload = () => {
cb(imgEl);
};
imgEl.addEventListener('error', () => {
cb();
});
imgEl.src = url;
});
};
nodeType.prototype.onDrawBackground = function (ctx) {
if (this.flags.collapsed) return;
let imageURLs = (this.images ?? []).map((i) =>
typeof i === 'string' ? i : formatUploadedUrl(i),
);
let imagesChanged = false;
if (JSON.stringify(this.displayingImages) !== JSON.stringify(imageURLs)) {
this.displayingImages = imageURLs;
imagesChanged = true;
}
if (!imagesChanged) return;
if (!imageURLs.length) {
this.imgs = null;
this.animatedImages = false;
return;
}
const promises = imageURLs.map((url) => {
if (url.startsWith('/view')) {
url = window.location.origin + url;
}
const u = new URL(url);
const filename =
u.searchParams.get('filename') || u.searchParams.get('name') || u.pathname.split('/').pop();
const ext = filename.split('.').pop();
const format = ['gif', 'webp', 'avif'].includes(ext) ? 'image' : 'video';
if (format === 'video') {
return createVideoNode(url);
} else {
return createImageNode(url);
}
});
Promise.all(promises)
.then((imgs) => {
this.imgs = imgs.filter(Boolean);
})
.then(() => {
if (!this.imgs.length) return;
this.animatedImages = true;
const widgetIdx = this.widgets?.findIndex((w) => w.name === ANIM_PREVIEW_WIDGET);
// Instead of using the canvas we'll use a IMG
if (widgetIdx > -1) {
// Replace content
const widget = this.widgets[widgetIdx];
widget.options.host.updateImages(this.imgs);
} else {
const host = createImageHost(this);
this.setSizeForImage(true);
const widget = this.addDOMWidget(ANIM_PREVIEW_WIDGET, 'img', host.el, {
host,
getHeight: host.getHeight,
onDraw: host.onDraw,
hideOnZoom: false,
});
widget.serializeValue = () => ({
height: host.el.clientHeight,
});
// widget.computeSize = (w) => ([w, 220]);
widget.options.host.updateImages(this.imgs);
}
this.imgs.forEach((img) => {
if (img instanceof HTMLVideoElement) {
img.muted = true;
img.autoplay = true;
img.play();
}
});
});
};
const { textWidget, comboWidget } = options;
if (textWidget) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const pathWidget = this.widgets.find((w) => w.name === textWidget);
pathWidget._value = pathWidget.value;
Object.defineProperty(pathWidget, 'value', {
set: (value) => {
pathWidget._value = value;
pathWidget.inputEl.value = value;
this.images = (value ?? '').split('\n').filter((url) => URL_REGEX.test(url));
},
get: () => {
return pathWidget._value;
},
});
pathWidget.inputEl.addEventListener('change', (e) => {
const value = e.target.value;
pathWidget._value = value;
this.images = (value ?? '').split('\n').filter((url) => URL_REGEX.test(url));
});
// Set value to ensure preview displays on initial add.
pathWidget.value = pathWidget._value;
});
}
if (comboWidget) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const pathWidget = this.widgets.find((w) => w.name === comboWidget);
pathWidget._value = pathWidget.value;
Object.defineProperty(pathWidget, 'value', {
set: (value) => {
pathWidget._value = value;
if (!value) {
return this.images = []
}
const parts = value.split("/")
const filename = parts.pop()
const subfolder = parts.join("/")
const extension = filename.split(".").pop();
const format = (["gif", "webp", "avif"].includes(extension)) ? 'image' : 'video'
this.images = [formatUploadedUrl({ filename, subfolder, type: "input", format: format })]
},
get: () => {
return pathWidget._value;
},
});
// Set value to ensure preview displays on initial add.
pathWidget.value = pathWidget._value;
});
}
chainCallback(nodeType.prototype, "onExecuted", function (message) {
if (message?.videos) {
this.images = message?.videos.map(formatUploadedUrl);
}
}); });
} }
export const hasWidgets = (node) => { app.registerExtension({
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
return false;
}
return true;
};
export const cleanupNode = (node) => {
if (!hasWidgets(node)) {
return;
}
for (const w of node.widgets) {
if (w.canvas) {
w.canvas.remove();
}
if (w.inputEl) {
w.inputEl.remove();
}
// calls the widget remove callback
w.onRemoved?.();
}
};
export const CreatePreviewElement = (name, val, format, callback) => {
const [type] = format.split("/");
const w = {
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth);
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch);
},
computeSize: function (_) {
const ratio = this.inputRatio || 1;
const width = Math.max(220, this.parent.size[0]);
return [width, width / ratio + 10];
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove();
}
},
};
w.inputEl = document.createElement(type === "video" ? "video" : "img");
w.inputEl.src = w.value;
if (type === "video") {
w.inputEl.setAttribute("type", "video/webm");
w.inputEl.autoplay = true;
w.inputEl.loop = true;
w.inputEl.controls = false;
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight;
callback?.();
};
document.body.appendChild(w.inputEl);
return w;
};
const videoPreview = {
name: "AnimateDiff.VideoPreview", name: "AnimateDiff.VideoPreview",
async beforeRegisterNodeDef(nodeType, nodeData, app) { init() {
const onExecuted = nodeType.prototype.onExecuted; $el('style', {
nodeType.prototype.onExecuted = function (message) { textContent: style,
const r = onExecuted ? onExecuted.apply(this, message) : undefined; parent: document.head,
});
if (message?.videos) {
this.videos = message.videos;
}
return r;
};
const onDrawBackground = nodeType.prototype.onDrawBackground;
nodeType.prototype.onDrawBackground = function (ctx) {
const r = onDrawBackground ? onDrawBackground.apply(this, arguments) : undefined;
const node = this;
const prefix = "ad_video_preview_";
if (node.videos_rendered === node.videos) {
return r;
}
if (node.widgets) {
const pos = node.widgets.findIndex((w) => w.name === `${prefix}_0`);
if (pos !== -1) {
for (let i = pos; i < node.widgets.length; i++) {
node.widgets[i].onRemoved?.();
}
node.widgets.length = pos;
}
}
if (node.videos) {
node.videos.forEach((params, i) => {
const previewUrl = api.apiURL(
"/view?" + new URLSearchParams(params).toString()
);
const w = node.addCustomWidget(
CreatePreviewElement(
`${prefix}_${i}`,
previewUrl,
params.format || "image/gif",
node.computeSizeKeepWidth.bind(node)
)
);
w.parent = node;
});
node.videos_rendered = node.videos;
}
return r;
};
const onRemoved = nodeType.prototype.onRemoved;
nodeType.prototype.onRemoved = function () {
cleanupNode(this);
return onRemoved ? onRemoved.apply(this, arguments) : undefined;
};
nodeType.prototype.computeSizeKeepWidth = function () {
this.setSize([
this.size[0],
this.computeSize([this.size[0], this.size[1]])[1],
]);
};
}, },
}; async beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData.name !== "AnimateDiffCombine") {
return;
}
app.registerExtension(videoPreview); addVideoPreview(nodeType);
},
});
+71 -169
View File
@@ -1,188 +1,90 @@
import { app } from "../../../scripts/app.js"; import { app } from "../../../scripts/app.js";
import { api } from "../../../scripts/api.js"; import { api } from "../../../scripts/api.js";
import { ComfyWidgets } from "../../../scripts/widgets.js";
const supportedVideoTypes = [ import {
"image/gif", chainCallback,
"video/webm", addVideoPreview,
"video/mp4", } from "./vid_preview.js";
"video/mov",
];
const VIDEOUPLOAD = (node, inputName, inputData, app) => { async function uploadFile(file) {
const previewWidget = "ad_video_preview"; try {
const videoWidget = node.widgets.find((w) => w.name === "video"); // Wrap file in formdata so it includes filename
let uploadWidget; const body = new FormData();
const new_file = new File([file], file.name, {
type: file.type,
lastModified: file.lastModified,
});
body.append("image", new_file);
body.append("subfolder", "video");
const resp = await api.fetchApi("/upload/image", {
method: "POST",
body,
});
const showVideo = (name) => { if (resp.status === 200 || resp.status === 201) {
let folder_separator = name.lastIndexOf("/"); return resp.json();
let subfolder = ""; } else {
if (folder_separator > -1) { alert(`Upload failed: ${resp.statusText}`);
subfolder = name.substring(0, folder_separator);
name = name.substring(folder_separator + 1);
}
const ext = name.substring(name.lastIndexOf(".") + 1);
const format = supportedVideoTypes.find((t) => t.endsWith(ext));
node.videos = [
{
filename: name,
type: "input",
subfolder: subfolder,
format,
},
];
};
var default_value = videoWidget.value;
Object.defineProperty(videoWidget, "value", {
set: function (value) {
this._real_value = value;
},
get: function () {
let value = "";
if (this._real_value) {
value = this._real_value;
} else {
return default_value;
}
if (value.filename) {
let real_value = value;
value = "";
if (real_value.subfolder) {
value = real_value.subfolder + "/";
}
value += real_value.filename;
if (real_value.type && real_value.type !== "input")
value += ` [${real_value.type}]`;
}
return value;
},
});
// Add our own callback to the combo widget to render an image when it changes
const cb = node.callback;
videoWidget.callback = function () {
showVideo(videoWidget.value);
if (cb) {
return cb.apply(this, arguments);
}
};
// On load if we have a value then render the image
// The value isnt set immediately so we need to wait a moment
// No change callbacks seem to be fired on initial setting of the value
requestAnimationFrame(() => {
if (videoWidget.value) {
showVideo(videoWidget.value);
}
});
async function uploadFile(file, updateNode, pasted = false) {
try {
// Wrap file in formdata so it includes filename
const body = new FormData();
body.append("image", file);
body.append("subfolder", "video");
const resp = await api.fetchApi("/upload/image", {
method: "POST",
body,
});
if (resp.status === 200) {
const data = await resp.json();
// Add the file to the dropdown list and update the widget value
let path = data.name;
if (data.subfolder) path = data.subfolder + "/" + path;
if (!videoWidget.options.values.includes(path)) {
videoWidget.options.values.push(path);
}
if (updateNode) {
showVideo(path);
videoWidget.value = path;
}
} else {
alert(resp.status + " - " + resp.statusText);
}
} catch (error) {
alert(error);
} }
} catch (error) {
alert(`Upload failed: ${error}`);
} }
}
const fileInput = document.createElement("input"); function addUploadWidget(nodeType, widgetName) {
Object.assign(fileInput, { chainCallback(nodeType.prototype, "onNodeCreated", function () {
type: "file", const pathWidget = this.widgets.find((w) => w.name === widgetName);
accept: supportedVideoTypes.join(","), if (pathWidget.element) {
style: "display: none", pathWidget.options.getMinHeight = () => 50;
onchange: async () => { pathWidget.options.getMaxHeight = () => 150;
if (fileInput.files.length) { }
await uploadFile(fileInput.files[0], true);
const fileInput = document.createElement("input");
chainCallback(this, "onRemoved", () => {
fileInput?.remove();
});
Object.assign(fileInput, {
type: "file",
accept: "video/webm,video/mp4,video/mkv,image/gif,image/webp",
style: "display: none",
onchange: async () => {
if (fileInput.files.length) {
const params = await uploadFile(fileInput.files[0]);
if (!params) {
// upload failed and file can not be added to options
return;
}
fileInput.value = "";
const filename = [params.subfolder, params.name || params.filename].filter(Boolean).join('/')
pathWidget.value = filename;
pathWidget.options.values.push(filename);
}
},
});
document.body.append(fileInput);
let uploadWidget = this.addWidget(
"button",
"choose video to upload",
"image",
() => {
app.canvas.node_widget = null;
fileInput.click();
} }
}, );
uploadWidget.options.serialize = false;
}); });
document.body.append(fileInput); }
// Create the button widget for selecting the files
uploadWidget = node.addWidget(
"button",
"choose file to upload",
"image",
() => {
fileInput.click();
}
);
uploadWidget.serialize = false;
// Add handler to check if an image is being dragged over our node
node.onDragOver = function (e) {
if (e.dataTransfer && e.dataTransfer.items) {
const image = [...e.dataTransfer.items].find((f) => f.kind === "file");
return !!image;
}
return false;
};
// On drop upload files
node.onDragDrop = function (e) {
console.log("onDragDrop called");
let handled = false;
for (const file of e.dataTransfer.files) {
if (file.type.startsWith("image/")) {
uploadFile(file, !handled); // Dont await these, any order is fine, only update on first one
handled = true;
}
}
return handled;
};
node.pasteFile = function (file) {
if (supportedVideoTypes.indexOf(file.type) > -1) {
const is_pasted =
file.name === "image.png" && file.lastModified - Date.now() < 2000;
uploadFile(file, true, is_pasted);
return true;
}
return false;
};
return { widget: uploadWidget };
};
ComfyWidgets["VIDEOUPLOAD"] = VIDEOUPLOAD;
// Adds an upload button to the nodes // Adds an upload button to the nodes
app.registerExtension({ app.registerExtension({
name: "AnimateDiff.UploadVideo", name: "AnimateDiff.UploadVideo",
async beforeRegisterNodeDef(nodeType, nodeData, app) { async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData?.input?.required?.video?.[1]?.video_upload === true) { if (nodeData?.input?.required?.video?.[1]?.video_upload === true) {
nodeData.input.required.upload = ["VIDEOUPLOAD"]; addUploadWidget(nodeType, 'video');
addVideoPreview(nodeType, { comboWidget: 'video' });
} }
}, },
}); });
+515
View File
@@ -0,0 +1,515 @@
{
"last_node_id": 21,
"last_link_id": 38,
"nodes": [
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
415,
186
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 3
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
29
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"photo of coastline, rocks, storm weather, wind, waves, lightning, 8k uhd, dslr, soft lighting, high quality, film grain, Fujifilm XT3"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1253,
191
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 28
},
{
"name": "vae",
"type": "VAE",
"link": 20
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
19
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 12,
"type": "AnimateDiffCombine",
"pos": [
1254,
290
],
"size": [
315,
507
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 19
}
],
"properties": {
"Node name for S&R": "AnimateDiffCombine"
},
"widgets_values": [
8,
0,
false,
"AnimateDiff",
"image/gif",
false,
"/view?filename=AnimateDiff_00003_.gif&subfolder=&type=temp&format=image%2Fgif"
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 5
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
30
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"blur, haze, deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers, deformed, distorted, disfigured, poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation"
]
},
{
"id": 20,
"type": "EmptyLatentImage",
"pos": [
522,
621
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
35
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
512,
512,
1
]
},
{
"id": 13,
"type": "VAELoader",
"pos": [
28,
223
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "VAE",
"type": "VAE",
"links": [
20
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAELoader"
},
"widgets_values": [
"vae-ft-mse-840000-ema-pruned.safetensors"
]
},
{
"id": 15,
"type": "AnimateDiffSampler",
"pos": [
882,
192
],
"size": {
"0": 315,
"1": 350
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "motion_module",
"type": "MOTION_MODULE",
"link": 24,
"slot_index": 0
},
{
"name": "model",
"type": "MODEL",
"link": 25,
"slot_index": 1
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 29
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 30
},
{
"name": "latent_image",
"type": "LATENT",
"link": 35
},
{
"name": "sliding_window_opts",
"type": "SLIDING_WINDOW_OPTS",
"link": null
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
28
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "AnimateDiffSampler"
},
"widgets_values": [
"default",
14,
45987230,
"fixed",
25,
7.5,
"ddim",
"ddim_uniform",
1
]
},
{
"id": 16,
"type": "AnimateDiffModuleLoader",
"pos": [
27,
345
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "lora_stack",
"type": "MOTION_LORA_STACK",
"link": 38,
"slot_index": 0
}
],
"outputs": [
{
"name": "MOTION_MODULE",
"type": "MOTION_MODULE",
"links": [
24
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "AnimateDiffModuleLoader"
},
"widgets_values": [
"mm_sd_v15_v2.ckpt"
]
},
{
"id": 21,
"type": "AnimateDiffLoraLoader",
"pos": [
-317,
350
],
"size": [
310,
80
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "lora_stack",
"type": "MOTION_LORA_STACK",
"link": null
}
],
"outputs": [
{
"name": "MOTION_LORA_STACK",
"type": "MOTION_LORA_STACK",
"links": [
38
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "AnimateDiffLoraLoader"
},
"widgets_values": [
"v2_lora_ZoomIn.ckpt",
1
]
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
28,
457
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
25
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
3,
5
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"RealisticVision_v20.safetensors"
]
}
],
"links": [
[
3,
4,
1,
6,
0,
"CLIP"
],
[
5,
4,
1,
7,
0,
"CLIP"
],
[
19,
8,
0,
12,
0,
"IMAGE"
],
[
20,
13,
0,
8,
1,
"VAE"
],
[
24,
16,
0,
15,
0,
"MOTION_MODULE"
],
[
25,
4,
0,
15,
1,
"MODEL"
],
[
28,
15,
0,
8,
0,
"LATENT"
],
[
29,
6,
0,
15,
2,
"CONDITIONING"
],
[
30,
7,
0,
15,
3,
"CONDITIONING"
],
[
35,
20,
0,
15,
4,
"LATENT"
],
[
38,
21,
0,
16,
0,
"MOTION_LORA_STACK"
]
],
"groups": [],
"config": {},
"extra": {},
"version": 0.4
}