VideoAsPrompt CogVideoX

This commit is contained in:
Your Name
2025-10-29 12:17:07 +00:00
commit 9ecd90b037
7 changed files with 2642 additions and 0 deletions
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS
NODE_DISPLAY_NAME_MAPPINGS = {k:k for k,v in NODE_CLASS_MAPPINGS.items()}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+117
View File
@@ -0,0 +1,117 @@
# Copyright (c) 2025 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
# SPDX-License-Identifier: Apache-2.0
#
# MOT (Motion Transfer) attention processor for Video-As-Prompt
# Extracted from Video-As-Prompt modified diffusers
from typing import Optional
import torch
import torch.nn.functional as F
class CogVideoXAttnMOTProcessor2_0:
r"""
Processor for implementing scaled dot-product attention for the CogVideoX model with MOT support.
It applies a rotary embedding on query and key vectors, but does not include spatial normalization.
This processor handles motion transfer by processing reference video attention separately.
"""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("CogVideoXAttnMOTProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn, # Attention module from diffusers.models.attention
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
# MOT specific parameters
is_before_attn: bool = False,
is_ref_video: Optional[bool] = False,
text_seq_length: Optional[int] = None,
) -> torch.Tensor:
"""
Apply attention with MOT support.
Args:
attn: The Attention module
hidden_states: Input hidden states
encoder_hidden_states: Encoder hidden states (text embeddings)
attention_mask: Attention mask
image_rotary_emb: Rotary position embeddings for images
is_before_attn: If True, only compute Q, K, V projections (before attention)
is_ref_video: Whether this is processing reference video
text_seq_length: Length of text sequence for splitting
Returns:
If is_before_attn=True: (query, key, value, attention_mask)
If is_before_attn=False: (hidden_states, encoder_hidden_states)
"""
if is_before_attn:
# Phase 1: Compute Q, K, V projections
text_seq_length = encoder_hidden_states.size(1)
# Concatenate text and video sequences
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
batch_size, sequence_length, _ = hidden_states.shape
if attention_mask is not None:
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])
# Linear projections
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
# Reshape for multi-head attention
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# Apply normalization if available
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# Apply RoPE (Rotary Position Embedding) if needed
if image_rotary_emb is not None:
# Import here to avoid circular dependency
from diffusers.models.embeddings import apply_rotary_emb
# Apply RoPE only to video tokens (skip text tokens)
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)
return query, key, value, attention_mask
else:
# Phase 2: Post-attention processing
batch_size, _, sequence_length, head_dim = hidden_states.shape
# Reshape back from multi-head format
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, sequence_length, attn.heads * head_dim)
# Linear projection
hidden_states = attn.to_out[0](hidden_states)
# Dropout
hidden_states = attn.to_out[1](hidden_states)
# Split back into text and video sequences
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
return hidden_states, encoder_hidden_states
File diff suppressed because it is too large Load Diff
+148
View File
@@ -0,0 +1,148 @@
# Copyright (c) 2025 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
# SPDX-License-Identifier: Apache-2.0
#
# MOT-specific embedding functions for Video-As-Prompt
# Extracted from Video-As-Prompt modified diffusers
from typing import Optional, Tuple, Union
import torch
# Import get_1d_rotary_pos_embed from official diffusers
from diffusers.models.embeddings import get_1d_rotary_pos_embed
def get_3d_rotary_pos_embed(
embed_dim,
crops_coords,
grid_size,
temporal_size,
theta: int = 10000,
use_real: bool = True,
grid_type: str = "linspace",
max_size: Optional[Tuple[int, int]] = None,
device: Optional[torch.device] = None,
mot_num: int = 0, # MOT-specific parameter
ref_type: str = "continous_negative", # MOT-specific parameter
start_point: int = 50,
gap: int = 30,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
RoPE for video tokens with 3D structure, with MOT support.
This is the MOT-modified version that supports motion transfer by handling
reference video position embeddings differently.
Args:
embed_dim: (`int`):
The embedding dimension size, corresponding to hidden_size_head.
crops_coords (`Tuple[int]`):
The top-left and bottom-right coordinates of the crop.
grid_size (`Tuple[int]`):
The grid size of the spatial positional embedding (height, width).
temporal_size (`int`):
The size of the temporal dimension.
theta (`float`):
Scaling factor for frequency computation.
grid_type (`str`):
Whether to use "linspace" or "slice" to compute grids.
mot_num (`int`):
Number of motion reference videos (MOT-specific).
ref_type (`str`):
Type of reference video position encoding (MOT-specific).
Returns:
`Tuple[torch.Tensor, torch.Tensor]`: cos and sin positional embeddings.
"""
if use_real is not True:
raise ValueError("`use_real = False` is not currently supported for get_3d_rotary_pos_embed")
if grid_type == "linspace":
start, stop = crops_coords
grid_size_h, grid_size_w = grid_size
grid_h = torch.linspace(
start[0], stop[0] * (grid_size_h - 1) / grid_size_h, grid_size_h, device=device, dtype=torch.float32
)
grid_w = torch.linspace(
start[1], stop[1] * (grid_size_w - 1) / grid_size_w, grid_size_w, device=device, dtype=torch.float32
)
grid_t = torch.arange(temporal_size, device=device, dtype=torch.float32)
grid_t = torch.linspace(
0, temporal_size * (temporal_size - 1) / temporal_size, temporal_size, device=device, dtype=torch.float32
)
# MOT-specific: Handle reference video position embeddings
if mot_num > 0:
if ref_type == "continous_negative":
orig_t_start = 0
orig_t_stop = temporal_size * (temporal_size - 1) / temporal_size
t_range = orig_t_stop - orig_t_start + 1
temporal_size = temporal_size * mot_num
grid_t = torch.linspace(-mot_num * t_range, -1, temporal_size, device=device, dtype=torch.float32)
elif ref_type == "discrete_long_reference":
start_offsets = start_point + torch.arange(mot_num, device=device, dtype=torch.float32) * gap
base_range = torch.arange(temporal_size, device=device, dtype=torch.float32)
grid_t = start_offsets.unsqueeze(1) + base_range
grid_t = grid_t.flatten().to(device=device, dtype=torch.float32)
else:
raise ValueError(f"Invalid {ref_type} passed for `ref_type`.")
elif grid_type == "slice":
max_h, max_w = max_size
grid_size_h, grid_size_w = grid_size
grid_h = torch.arange(max_h, device=device, dtype=torch.float32)
grid_w = torch.arange(max_w, device=device, dtype=torch.float32)
grid_t = torch.arange(temporal_size, device=device, dtype=torch.float32)
if mot_num > 0:
grid_t = torch.arange(-mot_num * temporal_size, 0, device=device, dtype=torch.float32)
else:
raise ValueError("Invalid value passed for `grid_type`.")
# Compute dimensions for each axis
dim_t = embed_dim // 4
dim_h = embed_dim // 8 * 3
dim_w = embed_dim // 8 * 3
# Temporal frequencies
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, theta=theta, use_real=True)
# Spatial frequencies for height and width
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, theta=theta, use_real=True)
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, theta=theta, use_real=True)
# BroadCast and concatenate temporal and spatial frequencies (height and width) into a 3d tensor
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
freqs_t = freqs_t[:, None, None, :].expand(
-1, grid_size_h, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_w, dim_t
freqs_h = freqs_h[None, :, None, :].expand(
temporal_size, -1, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_2, dim_h
freqs_w = freqs_w[None, None, :, :].expand(
temporal_size, grid_size_h, -1, -1
) # temporal_size, grid_size_h, grid_size_2, dim_w
freqs = torch.cat(
[freqs_t, freqs_h, freqs_w], dim=-1
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
freqs = freqs.view(
temporal_size * grid_size_h * grid_size_w, -1
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
return freqs
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
if grid_type == "slice":
t_cos, t_sin = t_cos[:temporal_size], t_sin[:temporal_size]
h_cos, h_sin = h_cos[:grid_size_h], h_sin[:grid_size_h]
w_cos, w_sin = w_cos[:grid_size_w], w_sin[:grid_size_w]
cos = combine_time_height_width(t_cos, h_cos, w_cos)
sin = combine_time_height_width(t_sin, h_sin, w_sin)
return cos, sin
+136
View File
@@ -0,0 +1,136 @@
from weakref import ref
import torch
import os
from diffusers import (
AutoencoderKLCogVideoX,
# CogVideoXImageToVideoMOTPipeline,
# CogVideoXTransformer3DMOTModel,
)
from diffusers.utils import export_to_video, load_video
from .pipeline_cogvideox_image2video_mot import CogVideoXImageToVideoMOTPipeline
from PIL import Image
from optimum.quanto import freeze, qint8, quantize
import folder_paths
from .cogvideox_transformer_3d_mot import CogVideoXTransformer3DMOTModel
import numpy as np
import comfy.utils
def pil_2_tensor(pil_image):
image = np.array(pil_image).astype(np.float32) / 255.0
image = torch.from_numpy(image)
return image
def tensor_2_pil(img_tensor):
i = 255. * img_tensor.squeeze().cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
return img
class RunningHub_VideoAsPrompt_Loader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"type": (["CogVideoX"], ),
}
}
RETURN_TYPES = ('RH_VideoAsPrompt_Pipeline', )
FUNCTION = "load"
CATEGORY = "RunningHub/VideoAsPrompt"
def load(self, type):
if type == "CogVideoX":
return (self.load_cogvideox(), )
return (None, )
def load_cogvideox(self):
model_base = os.path.join(folder_paths.models_dir, "Video-As-Prompt", "CogVideoX-5B")
vae = AutoencoderKLCogVideoX.from_pretrained(model_base, subfolder="vae", torch_dtype=torch.bfloat16)
transformer = CogVideoXTransformer3DMOTModel.from_pretrained(model_base, subfolder="transformer", torch_dtype=torch.bfloat16)
pipe = CogVideoXImageToVideoMOTPipeline.from_pretrained(
model_base, vae=vae, transformer=transformer, torch_dtype=torch.bfloat16,
)
if hasattr(pipe.vae, 'enable_slicing'):
pipe.vae.enable_slicing()
if hasattr(pipe.vae, 'enable_tiling'):
pipe.vae.enable_tiling()
quantize(pipe.transformer, qint8)
freeze(pipe.transformer)
pipe.enable_model_cpu_offload()
return pipe
class RunningHub_VideoAsPrompt_Sampler_CogVideoX:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"pipeline": ("RH_VideoAsPrompt_Pipeline", ),
"image": ("IMAGE", ),
"ref_video": ("IMAGE", ),
"prompt": ("STRING", {"default": "", "multiline": True}),
"prompt_mot_ref": ("STRING", {"default": "", "multiline": True}),
"height": ("INT", {"default": 480, "min": 16, "max": 1024}),
"width": ("INT", {"default": 720, "min": 16, "max": 1024}),
"num_frames": ("INT", {"default": 49, "min": 1, "max": 1024}),
# "frames_selection": ("STRING", {"default": "evenly", "choices": ["first", "evenly", "random"]}),
# "use_dynamic_cfg": ("BOOLEAN", {"default": False}),
"num_inference_steps": ("INT", {"default": 50, "min": 1, "max": 1000}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff,
"tooltip": "The random seed used for creating the noise."}),
}
}
RETURN_TYPES = ('IMAGE', )
FUNCTION = "sample"
TITLE = 'RunningHub VideoAsPrompt Sampler(CogVideoX)'
CATEGORY = "RunningHub/VideoAsPrompt"
def sample(self, **kwargs):
pipeline = kwargs["pipeline"]
image = kwargs["image"]
ref_video = kwargs["ref_video"]
prompt = kwargs["prompt"]
prompt_mot_ref = kwargs["prompt_mot_ref"]
height = kwargs["height"]
width = kwargs["width"]
num_frames = kwargs["num_frames"]
num_inference_steps = kwargs["num_inference_steps"]
self.pbar = comfy.utils.ProgressBar(num_inference_steps + 2)
# seed = kwargs["seed"]
ref_video = [tensor_2_pil(ref_frame) for ref_frame in ref_video]
image = tensor_2_pil(image)
idx = torch.linspace(0, len(ref_video) - 1, num_frames).long().tolist()
ref_frames = [ref_video[i] for i in idx]
output_frames = pipeline(
image=image,
ref_videos=[ref_frames],
prompt=prompt,
prompt_mot_ref=[prompt_mot_ref],
height=height,
width=width,
num_frames=num_frames,
frames_selection="evenly",
use_dynamic_cfg=True,
num_inference_steps = num_inference_steps,
update_func=self.update,
).frames[0]
export_to_video(output_frames, "output.mp4")
output_frames = [pil_2_tensor(output_frame) for output_frame in output_frames]
return (output_frames, )
def update(self):
self.pbar.update(1)
NODE_CLASS_MAPPINGS = {
"RunningHub VideoAsPrompt Sampler(CogVideoX)": RunningHub_VideoAsPrompt_Sampler_CogVideoX,
"RunningHub VideoAsPrompt Loader": RunningHub_VideoAsPrompt_Loader,
}
File diff suppressed because it is too large Load Diff
+1
View File
@@ -0,0 +1 @@
{"enable": true, "untracked_paths": []}