VideoAsPrompt CogVideoX
This commit is contained in:
@@ -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']
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -0,0 +1 @@
|
||||
{"enable": true, "untracked_paths": []}
|
||||
Reference in New Issue
Block a user