66 Commits
Author SHA1 Message Date
kijai b185059ba8 Merge branch 'main' into develop 2025-01-28 16:39:38 +02:00
kijai fed499e971 Update pyproject.toml 2025-01-28 16:39:21 +02:00
Jukka Seppänen f3dda43cdf Update readme.md 2025-01-23 11:18:57 +02:00
kijai 126322139f teacache tweaks 2025-01-20 22:22:45 +02:00
kijai 90e8367f5e Better sampling preview and support VHS live latent preview 2025-01-20 21:53:00 +02:00
kijai 76f7930d07 Add TeaCache 2025-01-20 17:06:06 +02:00
kijai 3d2ee02d83 delete LoRAs if fused to save memory 2025-01-20 17:04:42 +02:00
kijai 51daeef1b7 sageattn fp8/GGUF fix 2025-01-20 11:23:02 +02:00
Jukka Seppänen 5bca0548d9 Update readme.md 2025-01-18 21:27:57 +02:00
kijai 97b7b18f35 Create cut_and_drag_for_noisewarp_01.json 2025-01-18 21:20:24 +02:00
kijai f5454aa806 rename workflows folder, add NoiseWarp example 2025-01-18 21:14:04 +02:00
kijai 3a38d01414 Allow using precalculated noise
only when denoise is 1.0 for now
2025-01-18 17:02:41 +02:00
kijai 8c5e4f812d support other tora model 2025-01-14 14:39:44 +02:00
Jukka Seppänen eaaa0f6e1a Create LICENSE 2024-12-24 02:09:15 +02:00
kijai 25d0ede406 fix 2024-12-22 18:12:07 +02:00
kijai f16d38a5d2 Add Enhance-A-Video
https://github.com/NUS-HPC-AI-Lab/Enhance-A-Video
2024-12-22 01:26:18 +02:00
kijai fcc0f3e65a Update pipeline_cogvideox.py 2024-12-20 17:00:42 +02:00
kijai 0758d2d016 fix 2024-12-17 09:10:03 +02:00
kijai b5eefbf4d4 Initial support for Fun 1.5
I2V works, interpolation doesn't seem to (not sure if it should)
2024-12-17 01:16:11 +02:00
Jukka Seppänen 795f8b0565 Merge pull request #311 from glide-the/fix_pos_embedding
Add cogvideox-2b-img2vid CogVideoXModelLoader support
2024-12-08 11:52:12 +02:00
glide-the d9d30f24bb Add cogvideox-2b-img2vid CogVideoXModelLoader support
fix for transformer model patch_embed.pos_embedding dtype
or at add line ComfyUI-CogVideoXWrapper/embeddings.py:129 code
pos_embedding = pos_embedding.to(embeds.device, dtype=embeds.dtype)
2024-12-06 15:58:52 +08:00
kijai 729a6485ea expose sageattn 2.0.0 functions
_cuda versions seem to be required on RTX 30xx -series GPUs for sageattn + CogVideoX 1.5
2024-12-01 18:45:10 +02:00
kijai 2871f9b9d5 Update nodes.py 2024-11-28 18:20:09 +02:00
kijai 17a7aed013 CogVideoConcatLatent 2024-11-28 18:03:47 +02:00
kijai 411791c748 Update model_loading.py 2024-11-27 20:36:21 +02:00
kijai 7a10e732bb update cogvideox_Fun_180_orbit -example 2024-11-27 12:14:15 +02:00
kijai f9c747eff5 mid_image 2024-11-27 01:52:58 +02:00
kijai 2b211b9d1b sageattn 2.0.0 options 2024-11-27 01:16:22 +02:00
Jukka Seppänen 1ade29084e Merge branch 'main' of https://github.com/kijai/ComfyUI-CogVideoXWrapper 2024-11-24 22:27:27 +02:00
Jukka Seppänen f1b3bc0abf error earlier if sageattention fails to import 2024-11-24 22:27:26 +02:00
kijai c71bca9350 Update pipeline_cogvideox.py 2024-11-24 18:34:00 +02:00
Jukka Seppänen 8d6e53b556 Allow compiling VAE as well 2024-11-23 17:08:57 +02:00
kijai 9baf100366 fix 2b image2vid 2024-11-22 23:13:14 +02:00
Jukka Seppänen 6c7068b5bc Update nodes.py 2024-11-22 22:38:51 +02:00
kijai 895d3b83a4 Update model_loading.py 2024-11-21 03:05:51 +02:00
kijai 276a045a57 use selected load device as LoRA load device too 2024-11-21 02:46:07 +02:00
kijai e52dc36bc5 Update model_loading.py 2024-11-20 21:28:53 +02:00
kijai e5fc7c1bf3 Allow mixing Fun and not fun loras 2024-11-20 21:24:51 +02:00
kijai e187cfe22f Allow loading the "Rewards" LoRAs into 1.5 as well (for what it's worth) 2024-11-20 19:18:40 +02:00
kijai 573150de28 fix Tora when no autocast 2024-11-20 16:41:34 +02:00
kijai b74aa75026 Don't use autocast with fp/bf16 2024-11-20 14:22:10 +02:00
Jukka Seppänen b9f7b6e338 Merge pull request #261 from Dango233/Dango233-patch-1
Fix fused sdpa
2024-11-20 12:37:26 +02:00
Dango233 b31a025673 Fix fused sdpa 2024-11-20 17:40:28 +08:00
kijai ce329e0dce fix T2V 2024-11-20 10:17:11 +02:00
kijai de7e069286 fix noise augment 2024-11-20 02:12:56 +02:00
kijai 5cc570a467 Add start/end percent to image_conds 2024-11-20 02:07:29 +02:00
kijai b9688f3cd2 Add strength parameter for image encode 2024-11-20 01:32:05 +02:00
kijai ecd067260c Add CogVideoX-Fun-V1.1-5b-Control
https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-5b-Control
2024-11-20 01:23:54 +02:00
kijai c9efefe736 Update cogvideox_1_5_5b_I2V_01.json 2024-11-19 20:33:58 +02:00
kijai f7afa7d3be Update cogvideox_1_5_5b_I2V_01.json 2024-11-19 20:33:15 +02:00
kijai ebd6a3a4e8 update workflows to save_output enabled by default 2024-11-19 20:32:30 +02:00
kijai 41a0f33381 Update model_loading.py 2024-11-19 20:27:31 +02:00
kijai 1cfe0835f5 fix GGUF loader 2024-11-19 20:23:47 +02:00
kijai b0eabeba24 fix comfy attention output shape 2024-11-19 20:18:13 +02:00
kijai 822cb4ee1c Update cogvideox_1_5_5b_I2V_01.json 2024-11-19 20:02:22 +02:00
kijai 882faa6dea add comfyui attention mode 2024-11-19 19:55:51 +02:00
kijai cac1f81c51 Update pipeline_cogvideox.py 2024-11-19 19:42:09 +02:00
Jukka Seppänen fc647862b8 Merge pull request #194 from eltociear/patch-1
docs: update readme.md
2024-11-19 19:25:19 +02:00
kijai 516655b689 Update model_loading.py 2024-11-19 19:17:42 +02:00
kijai 67f2f6abb1 Merge branch 'refactor' 2024-11-19 19:16:39 +02:00
Jukka Seppänen 909d7026f3 Update model_loading.py 2024-11-09 20:20:15 +02:00
kijai 806a0fa1d6 Update model_loading.py 2024-11-08 21:31:31 +02:00
kijai f7a88cbd56 Update model_loading.py 2024-11-08 21:23:29 +02:00
kijai 4c2ce52f57 Update model_loading.py 2024-11-08 17:30:51 +02:00
Jukka Seppänen 4a597f1955 Update requirements.txt 2024-11-08 16:43:32 +02:00
Ikko Eltociear Ashimine eb902d9e9c docs: update readme.md
strenght -> strength
2024-10-29 09:21:38 +09:00
27 changed files with 4394 additions and 1724 deletions
+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.
+142 -116
View File
@@ -35,20 +35,56 @@ from diffusers.loaders import PeftAdapterMixin
from diffusers.models.embeddings import apply_rotary_emb
from .embeddings import CogVideoXPatchEmbed
from .enhance_a_video.enhance import get_feta_scores
from .enhance_a_video.globals import is_enhance_enabled, set_num_frames
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
try:
from sageattention import sageattn
SAGEATTN_IS_AVAILABLE = True
except:
SAGEATTN_IS_AVAILABLE = False
@torch.compiler.disable()
def sageattn_func(query, key, value, attn_mask=None, dropout_p=0.0,is_causal=False):
return sageattn(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p,is_causal=is_causal)
from comfy.ldm.modules.attention import optimized_attention
def set_attention_func(attention_mode, heads):
if attention_mode == "sdpa" or attention_mode == "fused_sdpa":
def func(q, k, v, is_causal=False, attn_mask=None):
return F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0, is_causal=is_causal)
return func
elif attention_mode == "comfy":
def func(q, k, v, is_causal=False, attn_mask=None):
return optimized_attention(q, k, v, mask=attn_mask, heads=heads, skip_reshape=True)
return func
elif attention_mode == "sageattn" or attention_mode == "fused_sageattn":
@torch.compiler.disable()
def func(q, k, v, is_causal=False, attn_mask=None):
return sageattn(q.to(v), k.to(v), v, is_causal=is_causal, attn_mask=attn_mask)
return func
elif attention_mode == "sageattn_qk_int8_pv_fp16_cuda":
from sageattention import sageattn_qk_int8_pv_fp16_cuda
@torch.compiler.disable()
def func(q, k, v, is_causal=False, attn_mask=None):
return sageattn_qk_int8_pv_fp16_cuda(q.to(v), k.to(v), v, is_causal=is_causal, attn_mask=attn_mask, pv_accum_dtype="fp32")
return func
elif attention_mode == "sageattn_qk_int8_pv_fp16_triton":
from sageattention import sageattn_qk_int8_pv_fp16_triton
@torch.compiler.disable()
def func(q, k, v, is_causal=False, attn_mask=None):
return sageattn_qk_int8_pv_fp16_triton(q.to(v), k.to(v), v, is_causal=is_causal, attn_mask=attn_mask)
return func
elif attention_mode == "sageattn_qk_int8_pv_fp8_cuda":
from sageattention import sageattn_qk_int8_pv_fp8_cuda
@torch.compiler.disable()
def func(q, k, v, is_causal=False, attn_mask=None):
return sageattn_qk_int8_pv_fp8_cuda(q.to(v), k.to(v), v, is_causal=is_causal, attn_mask=attn_mask, pv_accum_dtype="fp32+fp32")
return func
#for fastercache
def fft(tensor):
tensor_fft = torch.fft.fft2(tensor)
tensor_fft_shifted = torch.fft.fftshift(tensor_fft)
@@ -66,16 +102,25 @@ def fft(tensor):
return low_freq_fft, high_freq_fft
#for teacache
def poly1d(coefficients, x):
result = torch.zeros_like(x)
for i, coeff in enumerate(coefficients):
result += coeff * (x ** (len(coefficients) - 1 - i))
return result.abs()
#region Attention
class CogVideoXAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention for the CogVideoX model. It applies a rotary embedding on
query and key vectors, but does not include spatial normalization.
"""
def __init__(self):
def __init__(self, attn_func, attention_mode: Optional[str] = None):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("CogVideoXAttnProcessor requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
self.attention_mode = attention_mode
self.attn_func = attn_func
def __call__(
self,
attn: Attention,
@@ -83,7 +128,6 @@ class CogVideoXAttnProcessor2_0:
encoder_hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
attention_mode: Optional[str] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
@@ -97,7 +141,10 @@ class CogVideoXAttnProcessor2_0:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attention_mode != "fused_sdpa" or attention_mode != "fused_sageattn":
if attn.to_q.weight.dtype == torch.float16 or attn.to_q.weight.dtype == torch.bfloat16:
hidden_states = hidden_states.to(attn.to_q.weight.dtype)
if not "fused" in self.attention_mode:
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
@@ -123,17 +170,15 @@ class CogVideoXAttnProcessor2_0:
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
if not attn.is_cross_attention:
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
if attention_mode == "sageattn" or attention_mode == "fused_sageattn":
hidden_states = sageattn_func(query, key, value, attn_mask=attention_mask, dropout_p=0.0,is_causal=False)
else:
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
#if torch.isinf(hidden_states).any():
# raise ValueError(f"hidden_states after dot product has inf")
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
#feta
if is_enhance_enabled():
feta_scores = get_feta_scores(attn, query, key, head_dim, text_seq_length)
hidden_states = self.attn_func(query, key, value, attn_mask=attention_mask, is_causal=False)
if self.attention_mode != "comfy":
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
@@ -143,6 +188,10 @@ class CogVideoXAttnProcessor2_0:
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
if is_enhance_enabled():
hidden_states *= feta_scores
return hidden_states, encoder_hidden_states
#region Blocks
@@ -199,13 +248,15 @@ class CogVideoXBlock(nn.Module):
ff_inner_dim: Optional[int] = None,
ff_bias: bool = True,
attention_out_bias: bool = True,
attention_mode: Optional[str] = "sdpa",
):
super().__init__()
# 1. Self Attention
self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
attn_func = set_attention_func(attention_mode, num_attention_heads)
self.attn1 = Attention(
query_dim=dim,
dim_head=attention_head_dim,
@@ -214,7 +265,7 @@ class CogVideoXBlock(nn.Module):
eps=1e-6,
bias=attention_bias,
out_bias=attention_out_bias,
processor=CogVideoXAttnProcessor2_0(),
processor=CogVideoXAttnProcessor2_0(attn_func, attention_mode=attention_mode),
)
# 2. Feed Forward
@@ -243,7 +294,6 @@ class CogVideoXBlock(nn.Module):
fastercache_counter=0,
fastercache_start_step=15,
fastercache_device="cuda:0",
attention_mode="sdpa",
) -> torch.Tensor:
#print("hidden_states in block: ", hidden_states.shape) #1.5: torch.Size([2, 3200, 3072]) 10.: torch.Size([2, 6400, 3072])
text_seq_length = encoder_hidden_states.size(1)
@@ -282,7 +332,6 @@ class CogVideoXBlock(nn.Module):
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
attention_mode=attention_mode,
)
if fastercache_counter == fastercache_start_step:
self.cached_hidden_states = [attn_hidden_states.to(fastercache_device), attn_hidden_states.to(fastercache_device)]
@@ -294,8 +343,7 @@ class CogVideoXBlock(nn.Module):
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
attention_mode=attention_mode,
image_rotary_emb=image_rotary_emb
)
hidden_states = hidden_states + gate_msa * attn_hidden_states
@@ -404,6 +452,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
use_rotary_positional_embeddings: bool = False,
use_learned_positional_embeddings: bool = False,
patch_bias: bool = True,
attention_mode: Optional[str] = "sdpa",
):
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
@@ -457,6 +506,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
dropout=dropout,
activation_fn=activation_fn,
attention_bias=attention_bias,
attention_mode=attention_mode,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
)
@@ -484,7 +534,12 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
self.gradient_checkpointing = False
self.attention_mode = attention_mode
#tora
self.fuser_list = None
#fastercache
self.use_fastercache = False
self.fastercache_counter = 0
self.fastercache_start_step = 15
@@ -492,73 +547,21 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
self.fastercache_hf_step = 30
self.fastercache_device = "cuda"
self.fastercache_num_blocks_to_cache = len(self.transformer_blocks)
self.attention_mode = "sdpa"
#teacache
self.use_teacache = False
self.teacache_rel_l1_thresh = 0.0
if not self.config.use_rotary_positional_embeddings:
#CogVideoX-2B
self.teacache_coefficients = [-3.10658903e+01, 2.54732368e+01, -5.92380459e+00, 1.75769064e+00, -3.61568434e-03]
else:
#CogVideoX-5B
self.teacache_coefficients = [-1.53880483e+03, 8.43202495e+02, -1.34363087e+02, 7.97131516e+00, -5.23162339e-02]
def _set_gradient_checkpointing(self, module, value=False):
self.gradient_checkpointing = value
@property
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
def attn_processors(self) -> Dict[str, AttentionProcessor]:
r"""
Returns:
`dict` of attention processors: A dictionary containing all attention processors used in the model with
indexed by its weight name.
"""
# set recursively
processors = {}
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor()
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
return processors
for name, module in self.named_children():
fn_recursive_add_processors(name, module, processors)
return processors
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
r"""
Sets the attention processor to use to compute attention.
Parameters:
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
The instantiated processor class or a dictionary of processor classes that will be set as the processor
for **all** `Attention` layers.
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
processor. This is strongly recommended when setting trainable attention processors.
"""
count = len(self.attn_processors.keys())
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor)
else:
module.set_processor(processor.pop(f"{name}.processor"))
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
#region forward
def forward(
self,
hidden_states: torch.Tensor,
@@ -573,6 +576,8 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
return_dict: bool = True,
):
batch_size, num_frames, channels, height, width = hidden_states.shape
set_num_frames(num_frames) #enhance a video global
# 1. Time embedding
timesteps = timestep
@@ -620,8 +625,7 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
block_use_fastercache = i <= self.fastercache_num_blocks_to_cache,
fastercache_counter = self.fastercache_counter,
fastercache_start_step = self.fastercache_start_step,
fastercache_device = self.fastercache_device,
attention_mode = self.attention_mode
fastercache_device = self.fastercache_device
)
if (controlnet_states is not None) and (i < len(controlnet_states)):
@@ -680,34 +684,56 @@ class CogVideoXTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
recovered_uncond = rearrange(recovered_uncond.to(output.dtype), "(B T) C H W -> B T C H W", B=bb, C=cc, T=tt, H=hh, W=ww)
output = torch.cat([output, recovered_uncond])
else:
for i, block in enumerate(self.transformer_blocks):
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=emb,
image_rotary_emb=image_rotary_emb,
video_flow_feature=video_flow_features[i] if video_flow_features is not None else None,
fuser = self.fuser_list[i] if self.fuser_list is not None else None,
block_use_fastercache = i <= self.fastercache_num_blocks_to_cache,
fastercache_counter = self.fastercache_counter,
fastercache_start_step = self.fastercache_start_step,
fastercache_device = self.fastercache_device,
attention_mode = self.attention_mode
)
#has_nan = torch.isnan(hidden_states).any()
#if has_nan:
# raise ValueError(f"block output hidden_states has nan: {has_nan}")
if self.use_teacache:
if not hasattr(self, 'accumulated_rel_l1_distance'):
should_calc = True
self.accumulated_rel_l1_distance = 0
else:
self.accumulated_rel_l1_distance += poly1d(self.teacache_coefficients, ((emb-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()))
if self.accumulated_rel_l1_distance < self.teacache_rel_l1_thresh:
should_calc = False
self.teacache_counter += 1
else:
should_calc = True
self.accumulated_rel_l1_distance = 0
#print("self.accumulated_rel_l1_distance ", self.accumulated_rel_l1_distance)
self.previous_modulated_input = emb
if not should_calc:
hidden_states += self.previous_residual
encoder_hidden_states += self.previous_residual_encoder
if not self.use_teacache or (self.use_teacache and should_calc):
if self.use_teacache:
ori_hidden_states = hidden_states.clone()
ori_encoder_hidden_states = encoder_hidden_states.clone()
for i, block in enumerate(self.transformer_blocks):
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=emb,
image_rotary_emb=image_rotary_emb,
video_flow_feature=video_flow_features[i] if video_flow_features is not None else None,
fuser = self.fuser_list[i] if self.fuser_list is not None else None,
block_use_fastercache = i <= self.fastercache_num_blocks_to_cache,
fastercache_counter = self.fastercache_counter,
fastercache_start_step = self.fastercache_start_step,
fastercache_device = self.fastercache_device
)
#controlnet
if (controlnet_states is not None) and (i < len(controlnet_states)):
controlnet_states_block = controlnet_states[i]
controlnet_block_weight = 1.0
if isinstance(controlnet_weights, (list, np.ndarray)) or torch.is_tensor(controlnet_weights):
controlnet_block_weight = controlnet_weights[i]
print(controlnet_block_weight)
elif isinstance(controlnet_weights, (float, int)):
controlnet_block_weight = controlnet_weights
hidden_states = hidden_states + controlnet_states_block * controlnet_block_weight
#controlnet
if (controlnet_states is not None) and (i < len(controlnet_states)):
controlnet_states_block = controlnet_states[i]
controlnet_block_weight = 1.0
if isinstance(controlnet_weights, (list, np.ndarray)) or torch.is_tensor(controlnet_weights):
controlnet_block_weight = controlnet_weights[i]
print(controlnet_block_weight)
elif isinstance(controlnet_weights, (float, int)):
controlnet_block_weight = controlnet_weights
hidden_states = hidden_states + controlnet_states_block * controlnet_block_weight
if self.use_teacache:
self.previous_residual = hidden_states - ori_hidden_states
self.previous_residual_encoder = encoder_hidden_states - ori_encoder_hidden_states
if not self.config.use_rotary_positional_embeddings:
# CogVideoX-2B
View File
+82
View File
@@ -0,0 +1,82 @@
import torch
from einops import rearrange
from diffusers.models.attention import Attention
from .globals import get_enhance_weight, get_num_frames
# def get_feta_scores(query, key):
# img_q, img_k = query, key
# num_frames = get_num_frames()
# B, S, N, C = img_q.shape
# # Calculate spatial dimension
# spatial_dim = S // num_frames
# # Add time dimension between spatial and head dims
# query_image = img_q.reshape(B, spatial_dim, num_frames, N, C)
# key_image = img_k.reshape(B, spatial_dim, num_frames, N, C)
# # Expand time dimension
# query_image = query_image.expand(-1, -1, num_frames, -1, -1) # [B, S, T, N, C]
# key_image = key_image.expand(-1, -1, num_frames, -1, -1) # [B, S, T, N, C]
# # Reshape to match feta_score input format: [(B S) N T C]
# query_image = rearrange(query_image, "b s t n c -> (b s) n t c") #torch.Size([3200, 24, 5, 128])
# key_image = rearrange(key_image, "b s t n c -> (b s) n t c")
# return feta_score(query_image, key_image, C, num_frames)
def get_feta_scores(
attn: Attention,
query: torch.Tensor,
key: torch.Tensor,
head_dim: int,
text_seq_length: int,
) -> torch.Tensor:
num_frames = get_num_frames()
spatial_dim = int((query.shape[2] - text_seq_length) / num_frames)
query_image = rearrange(
query[:, :, text_seq_length:],
"B N (T S) C -> (B S) N T C",
N=attn.heads,
T=num_frames,
S=spatial_dim,
C=head_dim,
)
key_image = rearrange(
key[:, :, text_seq_length:],
"B N (T S) C -> (B S) N T C",
N=attn.heads,
T=num_frames,
S=spatial_dim,
C=head_dim,
)
return feta_score(query_image, key_image, head_dim, num_frames)
def feta_score(query_image, key_image, head_dim, num_frames):
scale = head_dim**-0.5
query_image = query_image * scale
attn_temp = query_image @ key_image.transpose(-2, -1) # translate attn to float32
attn_temp = attn_temp.to(torch.float32)
attn_temp = attn_temp.softmax(dim=-1)
# Reshape to [batch_size * num_tokens, num_frames, num_frames]
attn_temp = attn_temp.reshape(-1, num_frames, num_frames)
# Create a mask for diagonal elements
diag_mask = torch.eye(num_frames, device=attn_temp.device).bool()
diag_mask = diag_mask.unsqueeze(0).expand(attn_temp.shape[0], -1, -1)
# Zero out diagonal elements
attn_wo_diag = attn_temp.masked_fill(diag_mask, 0)
# Calculate mean for each token's attention matrix
# Number of off-diagonal elements per matrix is n*n - n
num_off_diag = num_frames * num_frames - num_frames
mean_scores = attn_wo_diag.sum(dim=(1, 2)) / num_off_diag
enhance_scores = mean_scores.mean() * (num_frames + get_enhance_weight())
enhance_scores = enhance_scores.clamp(min=1)
return enhance_scores
+31
View File
@@ -0,0 +1,31 @@
NUM_FRAMES = None
FETA_WEIGHT = None
ENABLE_FETA = False
def set_num_frames(num_frames: int):
global NUM_FRAMES
NUM_FRAMES = num_frames
def get_num_frames() -> int:
return NUM_FRAMES
def enable_enhance():
global ENABLE_FETA
ENABLE_FETA = True
def disable_enhance():
global ENABLE_FETA
ENABLE_FETA = False
def is_enhance_enabled() -> bool:
return ENABLE_FETA
def set_enhance_weight(feta_weight: float):
global FETA_WEIGHT
FETA_WEIGHT = feta_weight
def get_enhance_weight() -> float:
return FETA_WEIGHT
@@ -877,7 +877,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -818,7 +818,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -559,7 +559,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
File diff suppressed because it is too large Load Diff
@@ -639,7 +639,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -877,7 +877,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -14,7 +14,7 @@
"1": 574
},
"flags": {},
"order": 8,
"order": 7,
"mode": 0,
"inputs": [
{
@@ -103,7 +103,7 @@
"1": 122
},
"flags": {},
"order": 6,
"order": 5,
"mode": 0,
"inputs": [
{
@@ -152,7 +152,7 @@
"1": 168.08047485351562
},
"flags": {},
"order": 5,
"order": 4,
"mode": 0,
"inputs": [
{
@@ -275,7 +275,7 @@
"1": 198
},
"flags": {},
"order": 9,
"order": 8,
"mode": 0,
"inputs": [
{
@@ -310,80 +310,6 @@
true
]
},
{
"id": 44,
"type": "VHS_VideoCombine",
"pos": {
"0": 1884,
"1": -6
},
"size": [
605.3909912109375,
654.5737362132353
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 134
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 8,
"loop_count": 0,
"filename_prefix": "CogVideoX-I2V",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "CogVideoX-I2V_00004.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 8
},
"muted": false
}
}
},
{
"id": 37,
"type": "ImageResizeKJ",
@@ -396,7 +322,7 @@
"1": 266
},
"flags": {},
"order": 4,
"order": 3,
"mode": 0,
"inputs": [
{
@@ -476,7 +402,7 @@
"1": 144
},
"flags": {},
"order": 7,
"order": 6,
"mode": 0,
"inputs": [
{
@@ -575,52 +501,78 @@
]
},
{
"id": 64,
"type": "CogVideoImageEncodeFunInP",
"id": 44,
"type": "VHS_VideoCombine",
"pos": {
"0": 1861.032958984375,
"1": 752.6453247070312
},
"size": {
"0": 380.4000244140625,
"1": 146
"0": 1884,
"1": -6
},
"size": [
605.3909912109375,
310
],
"flags": {},
"order": 3,
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 134
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null
},
{
"name": "start_image",
"type": "IMAGE",
"link": null
},
{
"name": "end_image",
"type": "IMAGE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "image_cond_latents",
"type": "LATENT",
"links": null
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "CogVideoImageEncodeFunInP"
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": [
49,
false,
0
]
"widgets_values": {
"frame_rate": 16,
"loop_count": 0,
"filename_prefix": "CogVideoX_1_5_I2V",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "CogVideoX-I2V_00004.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 8
},
"muted": false
}
}
}
],
"links": [
@@ -725,10 +677,10 @@
"config": {},
"extra": {
"ds": {
"scale": 0.8390545288825803,
"scale": 0.7627768444387097,
"offset": [
351.5513339440394,
161.02862760095286
716.7143770104391,
291.75859557289965
]
}
},
@@ -498,7 +498,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -510,7 +510,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
@@ -168,7 +168,7 @@
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
File diff suppressed because one or more lines are too long
Binary file not shown.
-79
View File
@@ -1,79 +0,0 @@
import io
import torch
from PIL import Image
import struct
import numpy as np
from comfy.cli_args import args, LatentPreviewMethod
from comfy.taesd.taesd import TAESD
import comfy.model_management
import folder_paths
import comfy.utils
import logging
MAX_PREVIEW_RESOLUTION = args.preview_size
def preview_to_image(latent_image):
latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1
.mul(0xFF) # to 0..255
).to(device="cpu", dtype=torch.uint8, non_blocking=comfy.model_management.device_supports_non_blocking(latent_image.device))
return Image.fromarray(latents_ubyte.numpy())
class LatentPreviewer:
def decode_latent_to_preview(self, x0):
pass
def decode_latent_to_preview_image(self, preview_format, x0):
preview_image = self.decode_latent_to_preview(x0)
return ("GIF", preview_image, MAX_PREVIEW_RESOLUTION)
class Latent2RGBPreviewer(LatentPreviewer):
def __init__(self):
latent_rgb_factors = [[0.11945946736445662, 0.09919175788574555, -0.004832707433877734], [-0.0011977028264356232, 0.05496505130267682, 0.021321622433638193], [-0.014088548986590666, -0.008701477861945644, -0.020991313281459367], [0.03063921972519621, 0.12186477097625073, 0.0139593690235148], [0.0927403067854673, 0.030293187650929136, 0.05083134241694003], [0.0379112441305742, 0.04935199882777209, 0.058562766246777774], [0.017749911959153715, 0.008839453404921545, 0.036005638019226294], [0.10610119248526109, 0.02339855688237826, 0.057154257614084596], [0.1273639464837117, -0.010959856130713416, 0.043268631260428896], [-0.01873510946881321, 0.08220930648486932, 0.10613256772247093], [0.008429116376722327, 0.07623856561000408, 0.09295712117576727], [0.12938137079617007, 0.12360403483892413, 0.04478930933220116], [0.04565908794779364, 0.041064156741596365, -0.017695041535528512], [0.00019003240570281826, -0.013965147883381978, 0.05329669529635849], [0.08082391586738358, 0.11548306825496074, -0.021464170006615893], [-0.01517932393230994, -0.0057985555313003236, 0.07216646476618871]]
self.latent_rgb_factors = torch.tensor(latent_rgb_factors, device="cpu").transpose(0, 1)
self.latent_rgb_factors_bias = None
# if latent_rgb_factors_bias is not None:
# self.latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device="cpu")
def decode_latent_to_preview(self, x0):
self.latent_rgb_factors = self.latent_rgb_factors.to(dtype=x0.dtype, device=x0.device)
if self.latent_rgb_factors_bias is not None:
self.latent_rgb_factors_bias = self.latent_rgb_factors_bias.to(dtype=x0.dtype, device=x0.device)
latent_image = torch.nn.functional.linear(x0[0].permute(1, 2, 0), self.latent_rgb_factors,
bias=self.latent_rgb_factors_bias)
return preview_to_image(latent_image)
def get_previewer():
previewer = None
method = args.preview_method
if method != LatentPreviewMethod.NoPreviews:
# TODO previewer method
if method == LatentPreviewMethod.Auto:
method = LatentPreviewMethod.Latent2RGB
if previewer is None:
previewer = Latent2RGBPreviewer()
return previewer
def prepare_callback(model, steps, x0_output_dict=None):
preview_format = "JPEG"
if preview_format not in ["JPEG", "PNG"]:
preview_format = "JPEG"
previewer = get_previewer()
pbar = comfy.utils.ProgressBar(steps)
def callback(step, x0, x, total_steps):
if x0_output_dict is not None:
x0_output_dict["x0"] = x0
preview_bytes = None
if previewer:
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
pbar.update_absolute(step + 1, total_steps, preview_bytes)
return callback
+11 -8
View File
@@ -406,20 +406,23 @@ def merge_lora(transformer, lora_path, multiplier, device='cpu', dtype=torch.flo
else:
temp_name = layer_infos.pop(0)
weight_up = elems['lora_up.weight'].to(dtype)
weight_down = elems['lora_down.weight'].to(dtype)
weight_up = elems['lora_up.weight'].to(dtype).to(device)
weight_down = elems['lora_down.weight'].to(dtype).to(device)
if 'alpha' in elems.keys():
alpha = elems['alpha'].item() / weight_up.shape[1]
else:
alpha = 1.0
curr_layer.weight.data = curr_layer.weight.data.to(device)
if len(weight_up.shape) == 4:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
weight_down.squeeze(3).squeeze(2)).unsqueeze(
2).unsqueeze(3)
else:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
try:
if len(weight_up.shape) == 4:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2),
weight_down.squeeze(3).squeeze(2)).unsqueeze(
2).unsqueeze(3)
else:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
except:
print(f"Could not apply LoRA weight in layer {layer}")
return transformer
+200 -85
View File
@@ -70,6 +70,7 @@ class CogVideoLoraSelect:
RETURN_NAMES = ("lora", )
FUNCTION = "getlorapath"
CATEGORY = "CogVideoWrapper"
DESCRIPTION = "Select a LoRA model from ComfyUI/models/CogVideo/loras"
def getlorapath(self, lora, strength, prev_lora=None, fuse_lora=False):
cog_loras_list = []
@@ -86,6 +87,43 @@ class CogVideoLoraSelect:
cog_loras_list.append(cog_lora)
print(cog_loras_list)
return (cog_loras_list,)
class CogVideoLoraSelectComfy:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora": (folder_paths.get_filename_list("loras"),
{"tooltip": "LORA models are expected to be in ComfyUI/models/loras with .safetensors extension"}),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}),
},
"optional": {
"prev_lora":("COGLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}),
"fuse_lora": ("BOOLEAN", {"default": False, "tooltip": "Fuse the LoRA weights into the transformer"}),
}
}
RETURN_TYPES = ("COGLORA",)
RETURN_NAMES = ("lora", )
FUNCTION = "getlorapath"
CATEGORY = "CogVideoWrapper"
DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras"
def getlorapath(self, lora, strength, prev_lora=None, fuse_lora=False):
cog_loras_list = []
cog_lora = {
"path": folder_paths.get_full_path("loras", lora),
"strength": strength,
"name": lora.split(".")[0],
"fuse_lora": fuse_lora
}
if prev_lora is not None:
cog_loras_list.extend(prev_lora)
cog_loras_list.append(cog_lora)
print(cog_loras_list)
return (cog_loras_list,)
#region DownloadAndLoadCogVideoModel
class DownloadAndLoadCogVideoModel:
@@ -108,6 +146,8 @@ class DownloadAndLoadCogVideoModel:
"alibaba-pai/CogVideoX-Fun-V1.1-5b-InP",
"alibaba-pai/CogVideoX-Fun-V1.1-2b-Pose",
"alibaba-pai/CogVideoX-Fun-V1.1-5b-Pose",
"alibaba-pai/CogVideoX-Fun-V1.1-5b-Control",
"alibaba-pai/CogVideoX-Fun-V1.5-5b-InP",
"feizhengcong/CogvideoX-Interpolation",
"NimVideo/cogvideox-2b-img2vid"
],
@@ -123,7 +163,19 @@ class DownloadAndLoadCogVideoModel:
"block_edit": ("TRANSFORMERBLOCKS", {"default": None}),
"lora": ("COGLORA", {"default": None}),
"compile_args":("COMPILEARGS", ),
"attention_mode": (["sdpa", "sageattn", "fused_sdpa", "fused_sageattn"], {"default": "sdpa"}),
"attention_mode": ([
"sdpa",
"fused_sdpa",
"sageattn",
"fused_sageattn",
"sageattn_qk_int8_pv_fp8_cuda",
"sageattn_qk_int8_pv_fp16_cuda",
"sageattn_qk_int8_pv_fp16_triton",
"fused_sageattn_qk_int8_pv_fp8_cuda",
"fused_sageattn_qk_int8_pv_fp16_cuda",
"fused_sageattn_qk_int8_pv_fp16_triton",
"comfy"
], {"default": "sdpa"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
}
}
@@ -138,6 +190,19 @@ class DownloadAndLoadCogVideoModel:
enable_sequential_cpu_offload=False, block_edit=None, lora=None, compile_args=None,
attention_mode="sdpa", load_device="main_device"):
transformer = None
if "sage" in attention_mode:
try:
from sageattention import sageattn
except Exception as e:
raise ValueError(f"Can't import SageAttention: {str(e)}")
if "qk_int8" in attention_mode:
try:
from sageattention import sageattn_qk_int8_pv_fp16_cuda
except Exception as e:
raise ValueError(f"Can't import SageAttention 2.0.0: {str(e)}")
if precision == "fp16" and "1.5" in model:
raise ValueError("1.5 models do not currently work in fp16")
@@ -151,7 +216,7 @@ class DownloadAndLoadCogVideoModel:
download_path = folder_paths.get_folder_paths("CogVideo")[0]
if "Fun" in model:
if not "1.1" in model:
if "1.1" not in model and "1.5" not in model:
repo_id = "kijai/CogVideoX-Fun-pruned"
if "2b" in model:
base_path = os.path.join(folder_paths.models_dir, "CogVideoX_Fun", "CogVideoX-Fun-2b-InP") # location of the official model
@@ -161,7 +226,7 @@ class DownloadAndLoadCogVideoModel:
base_path = os.path.join(folder_paths.models_dir, "CogVideoX_Fun", "CogVideoX-Fun-5b-InP") # location of the official model
if not os.path.exists(base_path):
base_path = os.path.join(download_path, "CogVideoX-Fun-5b-InP")
elif "1.1" in model:
else:
repo_id = model
base_path = os.path.join(folder_paths.models_dir, "CogVideoX_Fun", (model.split("/")[-1])) # location of the official model
if not os.path.exists(base_path):
@@ -211,10 +276,10 @@ class DownloadAndLoadCogVideoModel:
local_dir_use_symlinks=False,
)
transformer = CogVideoXTransformer3DModel.from_pretrained(base_path, subfolder=subfolder)
transformer = CogVideoXTransformer3DModel.from_pretrained(base_path, subfolder=subfolder, attention_mode=attention_mode)
transformer = transformer.to(dtype).to(transformer_load_device)
if "1.5" in model:
if "1.5" in model and not "fun" in model:
transformer.config.sample_height = 300
transformer.config.sample_width = 300
@@ -233,52 +298,60 @@ class DownloadAndLoadCogVideoModel:
transformer,
scheduler,
dtype=dtype,
is_fun_inpaint=True if "fun" in model.lower() and "pose" not in model.lower() else False
is_fun_inpaint="fun" in model.lower() and not ("pose" in model.lower() or "control" in model.lower())
)
if "cogvideox-2b-img2vid" in model:
pipe.input_with_padding = False
#LoRAs
if lora is not None:
try:
adapter_list = []
adapter_weights = []
for l in lora:
fuse = True if l["fuse_lora"] else False
lora_sd = load_torch_file(l["path"])
for key, val in lora_sd.items():
if "lora_B" in key:
lora_rank = val.shape[1]
break
dimensionx_loras = ["orbit", "dimensionx"] # for now dimensionx loras need scaling
dimensionx_lora = False
adapter_list = []
adapter_weights = []
for l in lora:
if any(item in l["path"].lower() for item in dimensionx_loras):
dimensionx_lora = True
fuse = True if l["fuse_lora"] else False
lora_sd = load_torch_file(l["path"])
lora_rank = None
for key, val in lora_sd.items():
if "lora_B" in key:
lora_rank = val.shape[1]
break
if lora_rank is not None:
log.info(f"Merging rank {lora_rank} LoRA weights from {l['path']} with strength {l['strength']}")
adapter_name = l['path'].split("/")[-1].split(".")[0]
adapter_weight = l['strength']
pipe.load_lora_weights(l['path'], weight_name=l['path'].split("/")[-1], lora_rank=lora_rank, adapter_name=adapter_name)
#transformer = load_lora_into_transformer(lora, transformer)
adapter_list.append(adapter_name)
adapter_weights.append(adapter_weight)
for l in lora:
pipe.set_adapters(adapter_list, adapter_weights=adapter_weights)
else:
try: #Fun trainer LoRAs are loaded differently
from .lora_utils import merge_lora
log.info(f"Merging LoRA weights from {l['path']} with strength {l['strength']}")
pipe.transformer = merge_lora(pipe.transformer, l["path"], l["strength"], device=transformer_load_device, state_dict=lora_sd)
except:
raise ValueError(f"Can't recognize LoRA {l['path']}")
del lora_sd
mm.soft_empty_cache()
if adapter_list:
pipe.set_adapters(adapter_list, adapter_weights=adapter_weights)
if fuse:
lora_scale = 1
dimension_loras = ["orbit", "dimensionx"] # for now dimensionx loras need scaling
if any(item in lora[-1]["path"].lower() for item in dimension_loras):
if dimensionx_lora:
lora_scale = lora_scale / lora_rank
pipe.fuse_lora(lora_scale=lora_scale, components=["transformer"])
except: #Fun trainer LoRAs are loaded differently
from .lora_utils import merge_lora
for l in lora:
log.info(f"Merging LoRA weights from {l['path']} with strength {l['strength']}")
transformer = merge_lora(transformer, l["path"], l["strength"])
pipe.delete_adapters(adapter_list)
if "fused" in attention_mode:
from diffusers.models.attention import Attention
transformer.fuse_qkv_projections = True
for module in transformer.modules():
pipe.transformer.fuse_qkv_projections = True
for module in pipe.transformer.modules():
if isinstance(module, Attention):
module.fuse_projections(fuse=True)
transformer.attention_mode = attention_mode
if compile_args is not None:
pipe.transformer.to(memory_format=torch.channels_last)
@@ -391,6 +464,7 @@ class DownloadAndLoadCogVideoModel:
pipeline = {
"pipe": pipe,
"dtype": dtype,
"quantization": quantization,
"base_path": base_path,
"onediff": True if compile == "onediff" else False,
"cpu_offloading": enable_sequential_cpu_offload,
@@ -425,8 +499,7 @@ class DownloadAndLoadCogVideoGGUFModel:
},
"optional": {
"block_edit": ("TRANSFORMERBLOCKS", {"default": None}),
#"lora": ("COGLORA", {"default": None}),
"compile": (["disabled","torch"], {"tooltip": "compile the model for faster inference, these are advanced options only available on Linux, see readme for more info"}),
#"compile_args":("COMPILEARGS", ),
"attention_mode": (["sdpa", "sageattn"], {"default": "sdpa"}),
}
}
@@ -437,7 +510,13 @@ class DownloadAndLoadCogVideoGGUFModel:
CATEGORY = "CogVideoWrapper"
def loadmodel(self, model, vae_precision, fp8_fastmode, load_device, enable_sequential_cpu_offload,
block_edit=None, compile="disabled", attention_mode="sdpa"):
block_edit=None, compile_args=None, attention_mode="sdpa"):
if "sage" in attention_mode:
try:
from sageattention import sageattn
except Exception as e:
raise ValueError(f"Can't import SageAttention: {str(e)}")
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -496,8 +575,8 @@ class DownloadAndLoadCogVideoGGUFModel:
else:
transformer_config["in_channels"] = 16
transformer = CogVideoXTransformer3DModel.from_config(transformer_config)
transformer = CogVideoXTransformer3DModel.from_config(transformer_config, attention_mode=attention_mode)
cast_dtype = vae_dtype
params_to_keep = {"patch_embed", "pos_embedding", "time_embedding"}
if "2b" in model:
cast_dtype = torch.float16
@@ -524,10 +603,6 @@ class DownloadAndLoadCogVideoGGUFModel:
from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(transformer, vae_dtype, params_to_keep=params_to_keep)
if compile == "torch":
# compilation
for i, block in enumerate(transformer.transformer_blocks):
transformer.transformer_blocks[i] = torch.compile(block, fullgraph=False, dynamic=False, backend="inductor")
with open(scheduler_path) as f:
scheduler_config = json.load(f)
@@ -554,7 +629,12 @@ class DownloadAndLoadCogVideoGGUFModel:
vae = AutoencoderKLCogVideoX.from_config(vae_config).to(vae_dtype).to(offload_device)
vae.load_state_dict(vae_sd)
del vae_sd
pipe = CogVideoXPipeline(transformer, scheduler, dtype=vae_dtype)
pipe = CogVideoXPipeline(
transformer,
scheduler,
dtype=vae_dtype,
is_fun_inpaint="fun" in model.lower() and not ("pose" in model.lower() or "control" in model.lower())
)
if enable_sequential_cpu_offload:
pipe.enable_sequential_cpu_offload()
@@ -571,6 +651,7 @@ class DownloadAndLoadCogVideoGGUFModel:
pipeline = {
"pipe": pipe,
"dtype": vae_dtype,
"quantization": "GGUF",
"base_path": model,
"onediff": False,
"cpu_offloading": enable_sequential_cpu_offload,
@@ -587,7 +668,7 @@ class CogVideoXModelLoader:
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "The name of the checkpoint (model) to load.",}),
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6"], {"default": 'disabled', "tooltip": "optional quantization method"}),
@@ -598,7 +679,19 @@ class CogVideoXModelLoader:
"block_edit": ("TRANSFORMERBLOCKS", {"default": None}),
"lora": ("COGLORA", {"default": None}),
"compile_args":("COMPILEARGS", ),
"attention_mode": (["sdpa", "sageattn", "fused_sdpa", "fused_sageattn"], {"default": "sdpa"}),
"attention_mode": ([
"sdpa",
"fused_sdpa",
"sageattn",
"fused_sageattn",
"sageattn_qk_int8_pv_fp8_cuda",
"sageattn_qk_int8_pv_fp16_cuda",
"sageattn_qk_int8_pv_fp16_triton",
"fused_sageattn_qk_int8_pv_fp8_cuda",
"fused_sageattn_qk_int8_pv_fp16_cuda",
"fused_sageattn_qk_int8_pv_fp16_triton",
"comfy"
], {"default": "sdpa"}),
}
}
@@ -609,6 +702,12 @@ class CogVideoXModelLoader:
def loadmodel(self, model, base_precision, load_device, enable_sequential_cpu_offload,
block_edit=None, compile_args=None, lora=None, attention_mode="sdpa", quantization="disabled"):
transformer = None
if "sage" in attention_mode:
try:
from sageattention import sageattn
except Exception as e:
raise ValueError(f"Can't import SageAttention: {str(e)}")
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -632,6 +731,8 @@ class CogVideoXModelLoader:
model_type = "5b_I2V_1_5"
elif sd["patch_embed.proj.weight"].shape == (1920, 33, 2, 2):
model_type = "fun_2b"
elif sd["patch_embed.proj.weight"].shape == (1920, 32, 2, 2):
model_type = "cogvideox-2b-img2vid"
elif sd["patch_embed.proj.weight"].shape == (1920, 16, 2, 2):
model_type = "2b"
elif sd["patch_embed.proj.weight"].shape == (3072, 32, 2, 2):
@@ -653,7 +754,7 @@ class CogVideoXModelLoader:
with open(transformer_config_path) as f:
transformer_config = json.load(f)
if model_type in ["I2V", "I2V_5b", "fun_5b_pose", "5b_I2V_1_5"]:
if model_type in ["I2V", "I2V_5b", "fun_5b_pose", "5b_I2V_1_5", "cogvideox-2b-img2vid"]:
transformer_config["in_channels"] = 32
if "1_5" in model_type:
transformer_config["ofs_embed_dim"] = 512
@@ -669,7 +770,7 @@ class CogVideoXModelLoader:
transformer_config["sample_width"] = 300
with init_empty_weights():
transformer = CogVideoXTransformer3DModel.from_config(transformer_config)
transformer = CogVideoXTransformer3DModel.from_config(transformer_config, attention_mode=attention_mode)
#load weights
#params_to_keep = {}
@@ -679,7 +780,10 @@ class CogVideoXModelLoader:
#dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
del sd
# TODO fix for transformer model patch_embed.pos_embedding dtype
# or at add line ComfyUI-CogVideoXWrapper/embeddings.py:129 code
# pos_embedding = pos_embedding.to(embeds.device, dtype=embeds.dtype)
transformer = transformer.to(base_dtype).to(transformer_load_device)
#scheduler
with open(scheduler_config_path) as f:
@@ -697,49 +801,53 @@ class CogVideoXModelLoader:
module.fuse_projections(fuse=True)
transformer.attention_mode = attention_mode
if "fun" in model_type:
if not "pose" in model_type:
raise NotImplementedError("Fun models besides pose are not supported with this loader yet")
pipe = CogVideoX_Fun_Pipeline_Inpaint(vae, transformer, scheduler)
else:
pipe = CogVideoXPipeline(transformer, scheduler, dtype=base_dtype)
else:
pipe = CogVideoXPipeline(transformer, scheduler, dtype=base_dtype)
pipe = CogVideoXPipeline(
transformer,
scheduler,
dtype=base_dtype,
is_fun_inpaint="fun" in model.lower() and not ("pose" in model.lower() or "control" in model.lower())
)
if "cogvideox-2b-img2vid" == model_type:
pipe.input_with_padding = False
if enable_sequential_cpu_offload:
pipe.enable_sequential_cpu_offload()
#LoRAs
if lora is not None:
from .lora_utils import merge_lora#, load_lora_into_transformer
if "fun" in model.lower():
for l in lora:
log.info(f"Merging LoRA weights from {l['path']} with strength {l['strength']}")
transformer = merge_lora(transformer, l["path"], l["strength"])
else:
adapter_list = []
adapter_weights = []
for l in lora:
fuse = True if l["fuse_lora"] else False
lora_sd = load_torch_file(l["path"])
for key, val in lora_sd.items():
if "lora_B" in key:
lora_rank = val.shape[1]
break
dimensionx_loras = ["orbit", "dimensionx"] # for now dimensionx loras need scaling
dimensionx_lora = False
adapter_list = []
adapter_weights = []
for l in lora:
if any(item in l["path"].lower() for item in dimensionx_loras):
dimensionx_lora = True
fuse = True if l["fuse_lora"] else False
lora_sd = load_torch_file(l["path"])
lora_rank = None
for key, val in lora_sd.items():
if "lora_B" in key:
lora_rank = val.shape[1]
break
if lora_rank is not None:
log.info(f"Merging rank {lora_rank} LoRA weights from {l['path']} with strength {l['strength']}")
adapter_name = l['path'].split("/")[-1].split(".")[0]
adapter_weight = l['strength']
pipe.load_lora_weights(l['path'], weight_name=l['path'].split("/")[-1], lora_rank=lora_rank, adapter_name=adapter_name)
#transformer = load_lora_into_transformer(lora, transformer)
adapter_list.append(adapter_name)
adapter_weights.append(adapter_weight)
for l in lora:
pipe.set_adapters(adapter_list, adapter_weights=adapter_weights)
else:
try: #Fun trainer LoRAs are loaded differently
from .lora_utils import merge_lora
log.info(f"Merging LoRA weights from {l['path']} with strength {l['strength']}")
pipe.transformer = merge_lora(pipe.transformer, l["path"], l["strength"], device=transformer_load_device, state_dict=lora_sd)
except:
raise ValueError(f"Can't recognize LoRA {l['path']}")
if adapter_list:
pipe.set_adapters(adapter_list, adapter_weights=adapter_weights)
if fuse:
lora_scale = 1
dimension_loras = ["orbit", "dimensionx"] # for now dimensionx loras need scaling
if any(item in lora[-1]["path"].lower() for item in dimension_loras):
if dimensionx_lora:
lora_scale = lora_scale / lora_rank
pipe.fuse_lora(lora_scale=lora_scale, components=["transformer"])
@@ -801,15 +909,11 @@ class CogVideoXModelLoader:
manual_offloading = False # to disable manual .to(device) calls
log.info(f"Quantized transformer blocks to {quantization}")
# if load_device == "offload_device":
# pipe.transformer.to(offload_device)
# else:
# pipe.transformer.to(device)
pipeline = {
"pipe": pipe,
"dtype": base_dtype,
"quantization": quantization,
"base_path": model,
"onediff": False,
"cpu_offloading": enable_sequential_cpu_offload,
@@ -817,7 +921,6 @@ class CogVideoXModelLoader:
"model_name": model,
"manual_offloading": manual_offloading,
}
return (pipeline,)
#region VAE
@@ -827,12 +930,13 @@ class CogVideoXVAELoader:
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "The name of the checkpoint (vae) to load."}),
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
},
"optional": {
"precision": (["fp16", "fp32", "bf16"],
{"default": "bf16"}
),
"compile_args":("COMPILEARGS", ),
}
}
@@ -842,7 +946,7 @@ class CogVideoXVAELoader:
CATEGORY = "CogVideoWrapper"
DESCRIPTION = "Loads CogVideoX VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision):
def loadmodel(self, model_name, precision, compile_args=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -854,6 +958,10 @@ class CogVideoXVAELoader:
vae = AutoencoderKLCogVideoX.from_config(vae_config).to(dtype).to(offload_device)
vae.load_state_dict(vae_sd)
#compile
if compile_args is not None:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
vae = torch.compile(vae, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
return (vae,)
@@ -866,6 +974,7 @@ class DownloadAndLoadToraModel:
"model": (
[
"kijai/CogVideoX-5b-Tora",
"kijai/CogVideoX-5b-Tora-I2V",
],
),
},
@@ -895,14 +1004,17 @@ class DownloadAndLoadToraModel:
pass
download_path = os.path.join(folder_paths.models_dir, 'CogVideo', "CogVideoX-5b-Tora")
fuser_path = os.path.join(download_path, "fuser", "fuser.safetensors")
fuser_model = "fuser.safetensors" if not "I2V" in model else "fuser_I2V.safetensors"
fuser_path = os.path.join(download_path, "fuser", fuser_model)
if not os.path.exists(fuser_path):
log.info(f"Downloading Fuser model to: {fuser_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=model,
allow_patterns=["*fuser.safetensors*"],
allow_patterns=[fuser_model],
local_dir=download_path,
local_dir_use_symlinks=False,
)
@@ -924,14 +1036,15 @@ class DownloadAndLoadToraModel:
param.data = param.data.to(torch.bfloat16).to(device)
del fuser_sd
traj_extractor_path = os.path.join(download_path, "traj_extractor", "traj_extractor.safetensors")
traj_extractor_model = "traj_extractor.safetensors" if not "I2V" in model else "traj_extractor_I2V.safetensors"
traj_extractor_path = os.path.join(download_path, "traj_extractor", traj_extractor_model)
if not os.path.exists(traj_extractor_path):
log.info(f"Downloading trajectory extractor model to: {traj_extractor_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="kijai/CogVideoX-5b-Tora",
allow_patterns=["*traj_extractor.safetensors*"],
allow_patterns=[traj_extractor_model],
local_dir=download_path,
local_dir_use_symlinks=False,
)
@@ -1019,6 +1132,7 @@ NODE_CLASS_MAPPINGS = {
"CogVideoLoraSelect": CogVideoLoraSelect,
"CogVideoXVAELoader": CogVideoXVAELoader,
"CogVideoXModelLoader": CogVideoXModelLoader,
"CogVideoLoraSelectComfy": CogVideoLoraSelectComfy
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadCogVideoModel": "(Down)load CogVideo Model",
@@ -1028,4 +1142,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"CogVideoLoraSelect": "CogVideo LoraSelect",
"CogVideoXVAELoader": "CogVideoX VAE Loader",
"CogVideoXModelLoader": "CogVideoX Model Loader",
"CogVideoLoraSelectComfy": "CogVideo LoraSelect Comfy"
}
+95 -23
View File
@@ -49,6 +49,25 @@ if not "CogVideo" in folder_paths.folder_names_and_paths:
if not "cogvideox_loras" in folder_paths.folder_names_and_paths:
folder_paths.add_model_folder_path("cogvideox_loras", os.path.join(folder_paths.models_dir, "CogVideo", "loras"))
class CogVideoEnhanceAVideo:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"weight": ("FLOAT", {"default": 1.0, "min": 0, "max": 100, "step": 0.01, "tooltip": "The feta Weight of the Enhance-A-Video"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply Enhance-A-Video"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply Enhance-A-Video"}),
},
}
RETURN_TYPES = ("FETAARGS",)
RETURN_NAMES = ("feta_args",)
FUNCTION = "setargs"
CATEGORY = "CogVideoWrapper"
DESCRIPTION = "https://github.com/NUS-HPC-AI-Lab/Enhance-A-Video"
def setargs(self, **kwargs):
return (kwargs, )
class CogVideoContextOptions:
@classmethod
def INPUT_TYPES(s):
@@ -220,6 +239,9 @@ class CogVideoImageEncode:
"end_image": ("IMAGE", ),
"enable_tiling": ("BOOLEAN", {"default": False, "tooltip": "Enable tiling for the VAE to reduce memory usage"}),
"noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Augment image with noise"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
@@ -228,7 +250,7 @@ class CogVideoImageEncode:
FUNCTION = "encode"
CATEGORY = "CogVideoWrapper"
def encode(self, vae, start_image, end_image=None, enable_tiling=False, noise_aug_strength=0.0):
def encode(self, vae, start_image, end_image=None, enable_tiling=False, noise_aug_strength=0.0, strength=1.0, start_percent=0.0, end_percent=1.0):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
generator = torch.Generator(device=device).manual_seed(0)
@@ -251,19 +273,20 @@ class CogVideoImageEncode:
except:
pass
if noise_aug_strength > 0:
start_image = add_noise_to_reference_video(start_image, ratio=noise_aug_strength)
if end_image is not None:
end_image = add_noise_to_reference_video(end_image, ratio=noise_aug_strength)
latents_list = []
start_image = (start_image * 2.0 - 1.0).to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
if noise_aug_strength > 0:
start_image = add_noise_to_reference_video(start_image, ratio=noise_aug_strength)
start_latents = vae.encode(start_image).latent_dist.sample(generator)
start_latents = start_latents.permute(0, 2, 1, 3, 4) # B, T, C, H, W
if end_image is not None:
end_image = (end_image * 2.0 - 1.0).to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3)
if noise_aug_strength > 0:
end_image = add_noise_to_reference_video(end_image, ratio=noise_aug_strength)
end_latents = vae.encode(end_image).latent_dist.sample(generator)
end_latents = end_latents.permute(0, 2, 1, 3, 4) # B, T, C, H, W
latents_list = [start_latents, end_latents]
@@ -271,12 +294,16 @@ class CogVideoImageEncode:
else:
final_latents = start_latents
final_latents = final_latents * vae_scaling_factor
final_latents = final_latents * vae_scaling_factor * strength
log.info(f"Encoded latents shape: {final_latents.shape}")
vae.to(offload_device)
return ({"samples": final_latents}, )
return ({
"samples": final_latents,
"start_percent": start_percent,
"end_percent": end_percent
}, )
class CogVideoImageEncodeFunInP:
@classmethod
@@ -343,20 +370,17 @@ class CogVideoImageEncodeFunInP:
bs = 1
new_mask_pixel_values = []
print("input_image shape: ",input_image.shape)
for i in range(0, input_image.shape[0], bs):
mask_pixel_values_bs = input_image[i : i + bs]
mask_pixel_values_bs = vae.encode(mask_pixel_values_bs)[0]
print("mask_pixel_values_bs: ",mask_pixel_values_bs.parameters.shape)
mask_pixel_values_bs = mask_pixel_values_bs.mode()
print("mask_pixel_values_bs: ",mask_pixel_values_bs.shape, mask_pixel_values_bs.min(), mask_pixel_values_bs.max())
new_mask_pixel_values.append(mask_pixel_values_bs)
masked_image_latents = torch.cat(new_mask_pixel_values, dim = 0)
masked_image_latents = masked_image_latents.permute(0, 2, 1, 3, 4) # B, T, C, H, W
mask = torch.zeros_like(masked_image_latents[:, :, :1, :, :])
if end_image is not None:
mask[:, -1, :, :, :] = vae_scaling_factor
#if end_image is not None:
# mask[:, -1, :, :, :] = 0
mask[:, 0, :, :, :] = vae_scaling_factor
final_latents = masked_image_latents * vae_scaling_factor
@@ -560,6 +584,26 @@ class CogVideoXFasterCache:
"num_blocks_to_cache" : num_blocks_to_cache,
}
return (fastercache,)
class CogVideoXTeaCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Cache threshold, higher values are faster while sacrificing quality"}),
}
}
RETURN_TYPES = ("TEACACHEARGS",)
RETURN_NAMES = ("teacache_args",)
FUNCTION = "args"
CATEGORY = "CogVideoWrapper"
def args(self, rel_l1_thresh):
teacache = {
"rel_l1_thresh": rel_l1_thresh
}
return (teacache,)
#region Sampler
class CogVideoSampler:
@@ -587,6 +631,8 @@ class CogVideoSampler:
"controlnet": ("COGVIDECONTROLNET",),
"tora_trajectory": ("TORAFEATURES", ),
"fastercache": ("FASTERCACHEARGS", ),
"feta_args": ("FETAARGS", ),
"teacache_args": ("TEACACHEARGS", ),
}
}
@@ -596,13 +642,18 @@ class CogVideoSampler:
CATEGORY = "CogVideoWrapper"
def process(self, model, positive, negative, steps, cfg, seed, scheduler, num_frames, samples=None,
denoise_strength=1.0, image_cond_latents=None, context_options=None, controlnet=None, tora_trajectory=None, fastercache=None):
denoise_strength=1.0, image_cond_latents=None, context_options=None, controlnet=None, tora_trajectory=None, fastercache=None, feta_args=None, teacache_args=None):
mm.unload_all_models()
mm.soft_empty_cache()
model_name = model.get("model_name", "")
supports_image_conds = True if "I2V" in model_name or "interpolation" in model_name.lower() or "fun" in model_name.lower() else False
if "fun" in model_name.lower() and "pose" not in model_name.lower() and image_cond_latents is not None:
supports_image_conds = True if (
"I2V" in model_name or
"interpolation" in model_name.lower() or
"fun" in model_name.lower() or
"img2vid" in model_name.lower()
) else False
if "fun" in model_name.lower() and not ("pose" in model_name.lower() or "control" in model_name.lower()) and image_cond_latents is not None:
assert image_cond_latents["mask"] is not None, "For fun inpaint models use CogVideoImageEncodeFunInP"
fun_mask = image_cond_latents["mask"]
else:
@@ -611,7 +662,9 @@ class CogVideoSampler:
if image_cond_latents is not None:
assert supports_image_conds, "Image condition latents only supported for I2V and Interpolation models"
image_conds = image_cond_latents["samples"]
if "1.5" in model_name or "1_5" in model_name:
image_cond_start_percent = image_cond_latents.get("start_percent", 0.0)
image_cond_end_percent = image_cond_latents.get("end_percent", 1.0)
if ("1.5" in model_name or "1_5" in model_name) and not "fun" in model_name.lower():
image_conds = image_conds / 0.7 # needed for 1.5 models
else:
if not "fun" in model_name.lower():
@@ -674,6 +727,13 @@ class CogVideoSampler:
pipe.transformer.use_fastercache = False
pipe.transformer.fastercache_counter = 0
if teacache_args is not None:
pipe.transformer.use_teacache = True
pipe.transformer.teacache_rel_l1_thresh = teacache_args["rel_l1_thresh"]
log.info(f"TeaCache enabled with rel_l1_thresh: {pipe.transformer.teacache_rel_l1_thresh}")
else:
pipe.transformer.use_teacache = False
if not isinstance(cfg, list):
cfg = [cfg for _ in range(steps)]
else:
@@ -683,8 +743,9 @@ class CogVideoSampler:
except:
pass
autocastcondition = not model["onediff"] or not dtype == torch.float32
autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocastcondition else nullcontext()
autocast_context = torch.autocast(
mm.get_autocast_device(device), dtype=dtype
) if any(q in model["quantization"] for q in ("e4m3fn", "GGUF")) else nullcontext()
with autocast_context:
latents = model["pipe"](
num_inference_steps=steps,
@@ -707,6 +768,9 @@ class CogVideoSampler:
freenoise=context_options["freenoise"] if context_options is not None else None,
controlnet=controlnet,
tora=tora_trajectory if tora_trajectory is not None else None,
image_cond_start_percent=image_cond_start_percent if image_cond_latents is not None else 0.0,
image_cond_end_percent=image_cond_end_percent if image_cond_latents is not None else 1.0,
feta_args=feta_args,
)
if not model["cpu_offloading"] and model["manual_offloading"]:
pipe.transformer.to(offload_device)
@@ -718,6 +782,9 @@ class CogVideoSampler:
block.cached_encoder_hidden_states = None
print_memory(device)
if teacache_args is not None:
log.info(f"TeaCache skipped steps: {pipe.transformer.teacache_counter}")
mm.soft_empty_cache()
try:
torch.cuda.reset_peak_memory_stats(device)
@@ -855,8 +922,8 @@ class CogVideoXFunResizeToClosestBucket:
from .cogvideox_fun.utils import ASPECT_RATIO_512, get_closest_ratio
B, H, W, C = images.shape
# Count most suitable height and width
aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
# Find most suitable height and width
aspect_ratio_sample_size = {key : [x / 512 * base_resolution for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()}
closest_size, closest_ratio = get_closest_ratio(H, W, ratios=aspect_ratio_sample_size)
height, width = [int(x / 16) * 16 for x in closest_size]
@@ -896,7 +963,8 @@ class CogVideoLatentPreview:
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
#[[0.0658900170023352, 0.04687556512203313, -0.056971557475649186], [-0.01265770449940036, -0.02814809569100843, -0.0768912512529372], [0.061456544746314665, 0.0005511617552452358, -0.0652574975291287], [-0.09020669168815276, -0.004755440180558637, -0.023763970904494294], [0.031766964513999865, -0.030959599938418375, 0.08654669098083616], [-0.005981764690055846, -0.08809119252349802, -0.06439852368217663], [-0.0212114426433989, 0.08894281999597677, 0.05155629477559985], [-0.013947446911030725, -0.08987475069900677, -0.08923124751217484], [-0.08235967967978511, 0.07268025379974379, 0.08830486164536037], [-0.08052049179735378, -0.050116143175332195, 0.02023752569687405], [-0.07607527759162447, 0.06827156419895981, 0.08678111754261035], [-0.04689089232553825, 0.017294986041038893, -0.10280492336438908], [-0.06105783150270304, 0.07311850680875913, 0.019995735372550075], [-0.09232589996527711, -0.012869815059053047, -0.04355587834255975], [-0.06679931010802251, 0.018399815879067458, 0.06802404982033876], [-0.013062632927118165, -0.04292991477896661, 0.07476243356192845]]
latent_rgb_factors =[[0.11945946736445662, 0.09919175788574555, -0.004832707433877734], [-0.0011977028264356232, 0.05496505130267682, 0.021321622433638193], [-0.014088548986590666, -0.008701477861945644, -0.020991313281459367], [0.03063921972519621, 0.12186477097625073, 0.0139593690235148], [0.0927403067854673, 0.030293187650929136, 0.05083134241694003], [0.0379112441305742, 0.04935199882777209, 0.058562766246777774], [0.017749911959153715, 0.008839453404921545, 0.036005638019226294], [0.10610119248526109, 0.02339855688237826, 0.057154257614084596], [0.1273639464837117, -0.010959856130713416, 0.043268631260428896], [-0.01873510946881321, 0.08220930648486932, 0.10613256772247093], [0.008429116376722327, 0.07623856561000408, 0.09295712117576727], [0.12938137079617007, 0.12360403483892413, 0.04478930933220116], [0.04565908794779364, 0.041064156741596365, -0.017695041535528512], [0.00019003240570281826, -0.013965147883381978, 0.05329669529635849], [0.08082391586738358, 0.11548306825496074, -0.021464170006615893], [-0.01517932393230994, -0.0057985555313003236, 0.07216646476618871]]
#latent_rgb_factors =[[0.11945946736445662, 0.09919175788574555, -0.004832707433877734], [-0.0011977028264356232, 0.05496505130267682, 0.021321622433638193], [-0.014088548986590666, -0.008701477861945644, -0.020991313281459367], [0.03063921972519621, 0.12186477097625073, 0.0139593690235148], [0.0927403067854673, 0.030293187650929136, 0.05083134241694003], [0.0379112441305742, 0.04935199882777209, 0.058562766246777774], [0.017749911959153715, 0.008839453404921545, 0.036005638019226294], [0.10610119248526109, 0.02339855688237826, 0.057154257614084596], [0.1273639464837117, -0.010959856130713416, 0.043268631260428896], [-0.01873510946881321, 0.08220930648486932, 0.10613256772247093], [0.008429116376722327, 0.07623856561000408, 0.09295712117576727], [0.12938137079617007, 0.12360403483892413, 0.04478930933220116], [0.04565908794779364, 0.041064156741596365, -0.017695041535528512], [0.00019003240570281826, -0.013965147883381978, 0.05329669529635849], [0.08082391586738358, 0.11548306825496074, -0.021464170006615893], [-0.01517932393230994, -0.0057985555313003236, 0.07216646476618871]]
latent_rgb_factors = [[0.03197404301362048, 0.04091260743347359, 0.0015679806301828524], [0.005517101026578029, 0.0052348639043457755, -0.005613441650464035], [0.0012485338264583965, -0.016096744206117782, 0.025023940031635054], [0.01760126794276171, 0.0036818415416642893, -0.0006019202528157255], [0.000444954842288864, 0.006102128982092191, 0.0008457999272962447], [-0.010531904354560697, -0.0032275501924977175, -0.00886595780267917], [-0.0001454543946122991, 0.010199210750845965, -0.00012702234832386188], [0.02078497279904325, -0.001669617778939972, 0.006712703698951264], [0.005529571599763264, 0.009733929789086743, 0.001887302765339838], [0.012138415094654218, 0.024684961927224837, 0.037211249767461915], [0.0010364484570000384, 0.01983636315929172, 0.009864602025627755], [0.006802862648143341, -0.0010509255113510681, -0.007026003345126021], [0.0003532208468418043, 0.005351971582801936, -0.01845912126717106], [-0.009045079994694397, -0.01127941143183089, 0.0042294057970470806], [0.002548289972720752, 0.025224244654428216, -0.0006086130121693347], [-0.011135669222532816, 0.0018181308593668505, 0.02794541485349922]]
import random
random.seed(seed)
latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)]
@@ -945,6 +1013,8 @@ NODE_CLASS_MAPPINGS = {
"CogVideoLatentPreview": CogVideoLatentPreview,
"CogVideoXTorchCompileSettings": CogVideoXTorchCompileSettings,
"CogVideoImageEncodeFunInP": CogVideoImageEncodeFunInP,
"CogVideoEnhanceAVideo": CogVideoEnhanceAVideo,
"CogVideoXTeaCache": CogVideoXTeaCache,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"CogVideoSampler": "CogVideo Sampler",
@@ -961,4 +1031,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"CogVideoLatentPreview": "CogVideo LatentPreview",
"CogVideoXTorchCompileSettings": "CogVideo TorchCompileSettings",
"CogVideoImageEncodeFunInP": "CogVideo ImageEncode FunInP",
}
"CogVideoEnhanceAVideo": "CogVideo Enhance-A-Video",
"CogVideoXTeaCache": "CogVideoX TeaCache",
}
+92 -23
View File
@@ -29,6 +29,7 @@ from diffusers.loaders import CogVideoXLoraLoaderMixin
from .embeddings import get_3d_rotary_pos_embed
from .custom_cogvideox_transformer_3d import CogVideoXTransformer3DModel
from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
from comfy.utils import ProgressBar
@@ -110,6 +111,34 @@ def retrieve_timesteps(
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
class CogVideoXLatentFormat():
latent_channels = 16
latent_dimensions = 3
scale_factor = 0.7
taesd_decoder_name = None
latent_rgb_factors = [[0.03197404301362048, 0.04091260743347359, 0.0015679806301828524],
[0.005517101026578029, 0.0052348639043457755, -0.005613441650464035],
[0.0012485338264583965, -0.016096744206117782, 0.025023940031635054],
[0.01760126794276171, 0.0036818415416642893, -0.0006019202528157255],
[0.000444954842288864, 0.006102128982092191, 0.0008457999272962447],
[-0.010531904354560697, -0.0032275501924977175, -0.00886595780267917],
[-0.0001454543946122991, 0.010199210750845965, -0.00012702234832386188],
[0.02078497279904325, -0.001669617778939972, 0.006712703698951264],
[0.005529571599763264, 0.009733929789086743, 0.001887302765339838],
[0.012138415094654218, 0.024684961927224837, 0.037211249767461915],
[0.0010364484570000384, 0.01983636315929172, 0.009864602025627755],
[0.006802862648143341, -0.0010509255113510681, -0.007026003345126021],
[0.0003532208468418043, 0.005351971582801936, -0.01845912126717106],
[-0.009045079994694397, -0.01127941143183089, 0.0042294057970470806],
[0.002548289972720752, 0.025224244654428216, -0.0006086130121693347],
[-0.011135669222532816, 0.0018181308593668505, 0.02794541485349922]]
latent_rgb_factors_bias = [ -0.023, 0.0, -0.017]
class CogVideoXModelPlaceholder():
def __init__(self):
self.latent_format = CogVideoXLatentFormat
class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
r"""
Pipeline for text-to-video generation using CogVideoX.
@@ -195,7 +224,7 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
noise[:, place_idx:place_idx + delta, :, :, :] = noise[:, list_idx, :, :, :]
if latents is None:
latents = noise.to(device)
else:
elif denoise_strength < 1.0:
latents = latents.to(device)
timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
latent_timestep = timesteps[:1]
@@ -212,6 +241,8 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
latents = latents[:, :frames_needed, :, :, :]
latents = self.scheduler.add_noise(latents, noise.to(device), latent_timestep)
else:
latents = latents.to(device)
latents = latents * self.scheduler.init_noise_sigma # scale the initial noise by the standard deviation required by the scheduler
return latents, timesteps
@@ -349,6 +380,9 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
freenoise: Optional[bool] = True,
controlnet: Optional[dict] = None,
tora: Optional[dict] = None,
image_cond_start_percent: float = 0.0,
image_cond_end_percent: float = 1.0,
feta_args: Optional[dict] = None,
):
"""
@@ -395,7 +429,6 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
height = height or self.transformer.config.sample_size * self.vae_scale_factor_spatial
width = width or self.transformer.config.sample_size * self.vae_scale_factor_spatial
num_videos_per_prompt = 1
self.num_frames = num_frames
@@ -470,7 +503,9 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
# 5.5.
if image_cond_latents is not None:
if image_cond_latents.shape[1] == 2:
image_cond_frame_count = image_cond_latents.size(1)
patch_size_t = self.transformer.config.patch_size_t
if image_cond_frame_count == 2:
logger.info("More than one image conditioning frame received, interpolating")
padding_shape = (
batch_size,
@@ -481,12 +516,12 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
)
latent_padding = torch.zeros(padding_shape, device=device, dtype=self.vae_dtype)
image_cond_latents = torch.cat([image_cond_latents[:, 0, :, :, :].unsqueeze(1), latent_padding, image_cond_latents[:, -1, :, :, :].unsqueeze(1)], dim=1)
if self.transformer.config.patch_size_t is not None:
first_frame = image_cond_latents[:, : image_cond_latents.size(1) % self.transformer.config.patch_size_t, ...]
if patch_size_t:
first_frame = image_cond_latents[:, : image_cond_latents.size(1) % patch_size_t, ...]
image_cond_latents = torch.cat([first_frame, image_cond_latents], dim=1)
logger.info(f"image cond latents shape: {image_cond_latents.shape}")
elif image_cond_latents.shape[1] == 1:
elif image_cond_frame_count == 1:
logger.info("Only one image conditioning frame received, img2vid")
if self.input_with_padding:
padding_shape = (
@@ -499,13 +534,21 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
latent_padding = torch.zeros(padding_shape, device=device, dtype=self.vae_dtype)
image_cond_latents = torch.cat([image_cond_latents, latent_padding], dim=1)
# Select the first frame along the second dimension
if self.transformer.config.patch_size_t is not None:
first_frame = image_cond_latents[:, : image_cond_latents.size(1) % self.transformer.config.patch_size_t, ...]
if patch_size_t:
first_frame = image_cond_latents[:, : image_cond_latents.size(1) % patch_size_t, ...]
image_cond_latents = torch.cat([first_frame, image_cond_latents], dim=1)
else:
image_cond_latents = image_cond_latents.repeat(1, latents.shape[1], 1, 1, 1)
else:
logger.info(f"Received {image_cond_latents.shape[1]} image conditioning frames")
if fun_mask is not None and patch_size_t:
logger.info(f"1.5 model received {fun_mask.shape[1]} masks")
first_frame = image_cond_latents[:, : image_cond_frame_count % patch_size_t, ...]
image_cond_latents = torch.cat([first_frame, image_cond_latents], dim=1)
fun_mask_first_frame = fun_mask[:, : image_cond_frame_count % patch_size_t, ...]
fun_mask = torch.cat([fun_mask_first_frame, fun_mask], dim=1)
fun_mask[:, 1:, ...] = 0
image_cond_latents = image_cond_latents.to(self.vae_dtype)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
@@ -514,8 +557,8 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
# 7. context schedule
if context_schedule is not None:
if image_cond_latents is not None:
raise NotImplementedError("Context schedule not currently supported with image conditioning")
# if image_cond_latents is not None:
# raise NotImplementedError("Context schedule not currently supported with image conditioning")
logger.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap")
use_context_schedule = True
from .context import get_context_scheduler
@@ -562,7 +605,7 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
else:
controlnet_states = None
control_weights= None
# 9. Tora
if tora is not None:
trajectory_length = tora["video_flow_features"].shape[1]
logger.info(f"Tora trajectory length: {trajectory_length}")
@@ -570,20 +613,45 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
# raise ValueError(f"Tora trajectory length {trajectory_length} does not match inpaint_latents count {latents.shape[2]}")
for module in self.transformer.fuser_list:
for param in module.parameters():
param.data = param.data.to(device)
param.data = param.data.to(self.vae_dtype).to(device)
logger.info(f"Sampling {num_frames} frames in {latent_frames} latent frames at {width}x{height} with {num_inference_steps} inference steps")
from .latent_preview import prepare_callback
callback = prepare_callback(self.transformer, num_inference_steps)
if feta_args is not None:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
feta_end_percent = feta_args["end_percent"]
enable_enhance()
else:
disable_enhance()
# reset TeaCache
if hasattr(self.transformer, 'accumulated_rel_l1_distance'):
delattr(self.transformer, 'accumulated_rel_l1_distance')
self.transformer.teacache_counter = 0
# 11. Denoising loop
#from .latent_preview import prepare_callback
#callback = prepare_callback(self.transformer, num_inference_steps)
from latent_preview import prepare_callback
self.model = CogVideoXModelPlaceholder()
self.load_device = device
callback = prepare_callback(self, num_inference_steps)
# 9. Denoising loop
comfy_pbar = ProgressBar(len(timesteps))
with self.progress_bar(total=len(timesteps)) as progress_bar:
old_pred_original_sample = None # for DPM-solver++
for i, t in enumerate(timesteps):
if self.interrupt:
continue
current_step_percentage = i / num_inference_steps
if feta_args is not None:
if feta_start_percent <= current_step_percentage <= feta_end_percent:
enable_enhance()
else:
disable_enhance()
# region context schedule sampling
if use_context_schedule:
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
@@ -598,8 +666,6 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
current_step_percentage = i / num_inference_steps
# use same rotary embeddings for all context windows
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, context_frames, device)
@@ -710,7 +776,10 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
if image_cond_latents is not None:
latent_image_input = torch.cat([image_cond_latents] * 2) if do_classifier_free_guidance else image_cond_latents
if not image_cond_start_percent <= current_step_percentage <= image_cond_end_percent:
latent_image_input = torch.zeros_like(latent_model_input)
else:
latent_image_input = torch.cat([image_cond_latents] * 2) if do_classifier_free_guidance else image_cond_latents
if fun_mask is not None: #for fun img2vid and interpolation
fun_inpaint_mask = torch.cat([fun_mask] * 2) if do_classifier_free_guidance else fun_mask
masks_input = torch.cat([fun_inpaint_mask, latent_image_input], dim=2)
@@ -727,8 +796,6 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
current_step_percentage = i / num_inference_steps
if controlnet is not None:
controlnet_states = None
if (control_start <= current_step_percentage <= control_end):
@@ -746,7 +813,6 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
else:
controlnet_states = controlnet_states.to(dtype=self.vae_dtype)
# predict noise model_output
noise_pred = self.transformer(
hidden_states=latent_model_input,
@@ -786,8 +852,11 @@ class CogVideoXPipeline(DiffusionPipeline, CogVideoXLoraLoaderMixin):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if callback is not None:
callback(i, latents.detach()[-1], None, num_inference_steps)
if callback is not None:
alpha_prod_t = self.scheduler.alphas_cumprod[t]
beta_prod_t = 1 - alpha_prod_t
callback_tensor = (alpha_prod_t**0.5) * latent_model_input[0][:, :16, :, :] - (beta_prod_t**0.5) * noise_pred.detach()[0]
callback(i, callback_tensor * 5, None, num_inference_steps)
else:
comfy_pbar.update(1)
+2 -2
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-cogvideoxwrapper"
description = "Diffusers wrapper for CogVideoX -models: [a/https://github.com/THUDM/CogVideo](https://github.com/THUDM/CogVideo)"
version = "1.5.0"
description = "Diffusers wrapper for CogVideoX -models: https://github.com/THUDM/CogVideo"
version = "1.5.1"
license = {file = "LICENSE"}
dependencies = ["huggingface_hub", "diffusers>=0.31.0", "accelerate>=0.33.0"]
+15
View File
@@ -1,5 +1,20 @@
# WORK IN PROGRESS
Spreadsheet (WIP) of supported models and their supported features: https://docs.google.com/spreadsheets/d/16eA6mSL8XkTcu9fSWkPSHfRIqyAKJbR1O99xnuGdCKY/edit?usp=sharing
## Update 9
Added preliminary support for [Go-with-the-Flow](https://github.com/VGenAI-Netflix-Eyeline-Research/Go-with-the-Flow)
This uses LoRA weights available here: https://huggingface.co/Eyeline-Research/Go-with-the-Flow/tree/main
To create the input videos for the NoiseWarp process, I've added a node to KJNodes that works alongside my SplineEditor, and either [comfyui-inpaint-nodes](https://github.com/Acly/comfyui-inpaint-nodes) or just cv2 inpainting to create the cut and drag input videos.
The workflows are in the example_workflows -folder.
Quick video to showcase: First mask the subject, then use the cut and drag -workflow to create a video as seen here, then that video is used as input to the NoiseWarp node in the main workflow.
https://github.com/user-attachments/assets/112706b0-a38b-4c3c-b779-deba0827af4f
## BREAKING Update8
This is big one, and unfortunately to do the necessary cleanup and refactoring this will break every old workflow as they are.