Compare commits

...
24 Commits
Author SHA1 Message Date
“BrianChen1129” d7b07adc08 pipeline retrace; bug for black video 2024-12-09 09:47:06 +00:00
“BrianChen1129” 8a41e61cec add orginal mochi but performance bad 2024-12-09 08:44:24 +00:00
BrianChen1129â9 60ce6a62a6 delete sage in mochi-genmo 2024-12-08 03:56:26 +00:00
BrianChen1129â9 7b2ca9abec inference success 2024-12-08 03:54:07 +00:00
BrianChen1129â9 b537e01a88 inference success 2024-12-08 03:53:00 +00:00
Yongqi Chen a233b58a6c syn 2024-12-07 22:38:31 -05:00
Yongqi Chen d24b25a3e1 syn 2024-12-07 22:34:01 -05:00
BrianChen1129â9 2a7c147c4e syn with main 2024-12-08 02:38:43 +00:00
BrianChen1129â9 0c1c939d59 genmo mochi inference ready: 2024-12-08 02:35:40 +00:00
BrianChen1129â9 881e1f130a genmo mochi inference ready: 2024-12-08 02:35:08 +00:00
BrianChen1129â9 00e899cd90 syn 2024-12-07 09:21:02 +00:00
BrianChen1129 9f4151526b add sageattn 2024-12-07 00:33:12 +00:00
Yongqi Chen 1ce3983d68 add cpu offload 2024-12-06 12:45:08 -05:00
Yongqi Chen ee241cfa4d syn 2024-12-06 01:32:55 -05:00
Yongqi Chen 8106ee3f3d syn 2024-12-06 01:30:43 -05:00
Yongqi Chen 4dde52be9c remove conflict in train.py 2024-12-06 01:27:39 -05:00
BrianChenn1129 58abda5c09 syn with main: 2024-12-06 06:19:49 +00:00
BrianChenn1129 3e6019415a add demo 2024-12-06 06:12:43 +00:00
Yongqi Chen ccaf43c195 syn 2024-12-05 21:31:25 -05:00
Yongqi Chen 351538db29 syn 2024-12-05 21:27:03 -05:00
Yongqi Chen 361b24612d syn 2024-12-04 15:28:18 -05:00
Yongqi Chen 6ad03bea79 syn 2024-12-04 15:27:38 -05:00
Yongqi Chen 86e1f88877 syn 2024-12-04 00:11:22 -05:00
Yongqi Chen e212a9c6b9 syn with yongqi-dev2 and main 2024-12-03 23:51:46 -05:00
27 changed files with 4387 additions and 43 deletions
+18 -5
View File
@@ -1,13 +1,15 @@
import gradio as gr
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
import tempfile
import os
import argparse
from safetensors.torch import load_file
def init_args():
parser = argparse.ArgumentParser()
@@ -36,10 +38,21 @@ def load_model(args):
else:
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, False, args.linear_threshold, args.linear_range)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
mochi_genmo = True
if mochi_genmo:
model_path = "/root/fastmochi_genmo/dit.safetensors"
state_dcit = load_file(model_path)
transformer = AsymmDiTJoint()
transformer.load_state_dict(state_dcit)
# from IPython import embed
# embed()
transformer.config.in_channels = 12
print("load gennmo mochi successfully")
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder='transformer/')
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder='transformer/')
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
+2 -2
View File
@@ -18,7 +18,7 @@ from torch.distributed.fsdp import (
)
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmetricJointBlock
from functools import partial
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
@@ -88,7 +88,7 @@ def get_dit_fsdp_kwargs(
auto_wrap_policy = functools.partial(
transformer_auto_wrap_policy,
transformer_layer_cls={
MochiTransformerBlock,
MochiTransformerBlock, # AsymmetricJointBlock
},
)
@@ -0,0 +1,29 @@
from contextlib import contextmanager
import torch
try:
from flash_attn import flash_attn_varlen_func as flash_varlen_attn
except ImportError:
flash_varlen_attn = None
try:
from sageattention import sageattn as sage_attn
except ImportError:
sage_attn = None
from torch.nn.attention import SDPBackend, sdpa_kernel
training_backends = [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]
eval_backends = list(training_backends)
if torch.cuda.get_device_properties(0).major >= 9.0:
# Enable fast CuDNN attention on Hopper.
# This gives NaN on the backward pass for some reason,
# so only use it for evaluation.
eval_backends.append(SDPBackend.CUDNN_ATTENTION)
@contextmanager
def sdpa_attn_ctx(training: bool = False):
with sdpa_kernel(training_backends if training else eval_backends):
yield
@@ -0,0 +1,87 @@
import contextlib
from typing import Any, Iterable, Iterator, Optional
try:
from tqdm import tqdm
except ImportError:
tqdm = None
try:
from ray.experimental.tqdm_ray import tqdm as ray_tqdm
except:
ray_tqdm = None
# Global state
_current_progress_type = "none"
_is_progress_bar_active = False
class DummyProgressBar:
"""A no-op progress bar that mimics tqdm interface"""
def __init__(self, iterable=None, **kwargs):
self.iterable = iterable
def __iter__(self):
return iter(self.iterable)
def update(self, n=1):
pass
def close(self):
pass
def set_description(self, desc):
pass
def get_new_progress_bar(iterable: Optional[Iterable] = None, **kwargs) -> Any:
if not _is_progress_bar_active:
return DummyProgressBar(iterable=iterable, **kwargs)
if _current_progress_type == "tqdm":
if tqdm is None:
raise ImportError("tqdm is required but not installed. Please install tqdm to use the tqdm progress bar.")
return tqdm(iterable=iterable, **kwargs)
elif _current_progress_type == "ray_tqdm":
if ray_tqdm is None:
raise ImportError("ray is required but not installed. Please install ray to use the ray_tqdm progress bar.")
return ray_tqdm(iterable=iterable, **kwargs)
return DummyProgressBar(iterable=iterable, **kwargs)
@contextlib.contextmanager
def progress_bar(type: str = "none", enabled=True):
"""
Context manager for setting progress bar type and options.
Args:
type: Type of progress bar ("none" or "tqdm")
**options: Options to pass to the progress bar (e.g., total, desc)
Raises:
ValueError: If progress bar type is invalid
RuntimeError: If progress bars are nested
Example:
with progress_bar(type="tqdm", total=100):
for i in get_new_progress_bar(range(100)):
process(i)
"""
if type not in ("none", "tqdm", "ray_tqdm"):
raise ValueError("Progress bar type must be 'none' or 'tqdm' or 'ray_tqdm'")
if not enabled:
type = "none"
global _current_progress_type, _is_progress_bar_active
if _is_progress_bar_active:
raise RuntimeError("Nested progress bars are not supported")
_is_progress_bar_active = True
_current_progress_type = type
try:
yield
finally:
_is_progress_bar_active = False
_current_progress_type = "none"
+67
View File
@@ -0,0 +1,67 @@
import os
import subprocess
import tempfile
import time
import numpy as np
from moviepy.editor import ImageSequenceClip
from PIL import Image
from genmo.lib.progress import get_new_progress_bar
class Timer:
def __init__(self):
self.times = {} # Dictionary to store times per stage
def __call__(self, name):
print(f"Timing {name}")
return self.TimerContextManager(self, name)
def print_stats(self):
total_time = sum(self.times.values())
# Print table header
print("{:<20} {:>10} {:>10}".format("Stage", "Time(s)", "Percent"))
for name, t in self.times.items():
percent = (t / total_time) * 100 if total_time > 0 else 0
print("{:<20} {:>10.2f} {:>9.2f}%".format(name, t, percent))
class TimerContextManager:
def __init__(self, outer, name):
self.outer = outer # Reference to the Timer instance
self.name = name
self.start_time = None
def __enter__(self):
self.start_time = time.perf_counter()
return self
def __exit__(self, exc_type, exc_value, traceback):
end_time = time.perf_counter()
elapsed = end_time - self.start_time
self.outer.times[self.name] = self.outer.times.get(self.name, 0) + elapsed
def save_video(final_frames, output_path, fps=30):
assert final_frames.ndim == 4 and final_frames.shape[3] == 3, f"invalid shape: {final_frames} (need t h w c)"
if final_frames.dtype != np.uint8:
final_frames = (final_frames * 255).astype(np.uint8)
ImageSequenceClip(list(final_frames), fps=fps).write_videofile(output_path)
def create_memory_tracker():
import torch
previous = [None] # Use list for mutable closure state
def track(label="all2all"):
current = torch.cuda.memory_allocated() / 1e9
if previous[0] is not None:
diff = current - previous[0]
sign = "+" if diff >= 0 else ""
print(f"GPU memory ({label}): {current:.2f} GB ({sign}{diff:.2f} GB)")
else:
print(f"GPU memory ({label}): {current:.2f} GB")
previous[0] = current # type: ignore
return track
@@ -0,0 +1,721 @@
import os
from typing import Dict, List, Optional, Tuple, Any
import warnings
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.attention import sdpa_kernel
from fastvideo.models.mochi_genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.layers import (
FeedForward,
PatchEmbed,
RMSNorm,
TimestepEmbedder,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.lora import LoraLinear
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
residual_tanh_gated_rmsnorm,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.rope_mixed import (
compute_mixed_rotation,
create_position_matrix,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.utils import (
AttentionPool,
modulate,
pad_and_split_xy,
)
from fastvideo.models.mochi_genmo.mochi_preview.pipelines import compute_packed_indices
from diffusers.models.modeling_utils import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import PeftAdapterMixin
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
if enabled:
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
return fn(*args, **kwargs)
class AsymmetricAttention(nn.Module):
def __init__(
self,
dim_x: int,
dim_y: int,
num_heads: int = 8,
qkv_bias: bool = False,
qk_norm: bool = True,
update_y: bool = True,
out_bias: bool = True,
attention_mode: str = "flash",
softmax_scale: Optional[float] = None,
device: Optional[torch.device] = None,
# Disable LoRA by default ...
qkv_proj_lora_rank: int = 0,
qkv_proj_lora_alpha: int = 0,
qkv_proj_lora_dropout: float = 0.0,
out_proj_lora_rank: int = 0,
out_proj_lora_alpha: int = 0,
out_proj_lora_dropout: float = 0.0,
):
super().__init__()
self.attention_mode = attention_mode
self.dim_x = dim_x
self.dim_y = dim_y
self.num_heads = num_heads
self.head_dim = dim_x // num_heads
self.update_y = update_y
self.softmax_scale = softmax_scale
if dim_x % num_heads != 0:
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
# Input layers.
self.qkv_bias = qkv_bias
qkv_lora_kwargs = dict(
bias=qkv_bias,
device=device,
r=qkv_proj_lora_rank,
lora_alpha=qkv_proj_lora_alpha,
lora_dropout=qkv_proj_lora_dropout,
)
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
# Project text features to match visual features (dim_y -> dim_x)
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
# Query and key normalization for stability.
assert qk_norm
self.q_norm_x = RMSNorm(self.head_dim, device=device)
self.k_norm_x = RMSNorm(self.head_dim, device=device)
self.q_norm_y = RMSNorm(self.head_dim, device=device)
self.k_norm_y = RMSNorm(self.head_dim, device=device)
# Output layers. y features go back down from dim_x -> dim_y.
proj_lora_kwargs = dict(
bias=out_bias,
device=device,
r=out_proj_lora_rank,
lora_alpha=out_proj_lora_alpha,
lora_dropout=out_proj_lora_dropout,
)
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
def run_qkv_y(self, y):
local_heads = self.num_heads
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
q_y, k_y, v_y = qkv_y.unbind(2)
q_y = self.q_norm_y(q_y)
k_y = self.k_norm_y(k_y)
return q_y, k_y, v_y
def prepare_qkv(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor,
scale_y: torch.Tensor,
rope_cos: torch.Tensor,
rope_sin: torch.Tensor,
valid_token_indices: torch.Tensor,
max_seqlen_in_batch: int,
):
# Process visual features
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
assert qkv_x.dtype == torch.bfloat16
B, M, _ = qkv_x.size()
qkv_x = qkv_x.view(B, M, 3, self.num_heads, -1)
qkv_x = qkv_x.permute(2, 0, 1, 3, 4)
# Split qkv_x into q, k, v
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
q_x = self.q_norm_x(q_x)
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
k_x = self.k_norm_x(k_x)
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
# Concatenate streams
B, N, num_heads, head_dim = q_x.size()
D = num_heads * head_dim
# Process text features
if B == 1:
text_seqlen = max_seqlen_in_batch - N
if text_seqlen > 0:
y = y[:, :text_seqlen] # Remove padding tokens.
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
q = torch.cat([q_x, q_y], dim=1)
k = torch.cat([k_x, k_y], dim=1)
v = torch.cat([v_x, v_y], dim=1)
else:
q, k, v = q_x, k_x, v_x
else:
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
indices = valid_token_indices[:, None].expand(-1, D)
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
q = q.view(-1, num_heads, head_dim)
k = k.view(-1, num_heads, head_dim)
v = v.view(-1, num_heads, head_dim)
return q, k, v
@torch.autocast("cuda", enabled=False)
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
out: torch.Tensor = flash_varlen_attn(
q, k, v,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_seqlen_in_batch,
max_seqlen_k=max_seqlen_in_batch,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
) # (total, local_heads, head_dim)
return out.view(total, local_dim)
def sdpa_attention(self, q, k, v):
with sdpa_attn_ctx(training=self.training):
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
)
return out
@torch.autocast("cuda", enabled=False)
def sage_attention(self, q, k, v):
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
def run_attention(
self,
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
*,
B: int,
cu_seqlens: Optional[torch.Tensor] = None,
max_seqlen_in_batch: Optional[int] = None,
):
local_heads = self.num_heads
local_dim = local_heads * self.head_dim
# Check shapes
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
total = q.size(0)
assert k.size(0) == total and v.size(0) == total
if self.attention_mode == "flash":
out = self.flash_attention(
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
else:
assert B == 1, \
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
q = rearrange(q, "(b s) h d -> b h s d", b=B)
k = rearrange(k, "(b s) h d -> b h s d", b=B)
v = rearrange(v, "(b s) h d -> b h s d", b=B)
if self.attention_mode == "sdpa":
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
elif self.attention_mode == "sage":
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
else:
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
out = rearrange(out, "b h s d -> (b s) (h d)")
return out
def post_attention(
self,
out: torch.Tensor,
B: int,
M: int,
L: int,
dtype: torch.dtype,
valid_token_indices: torch.Tensor,
):
"""
Args:
out: (total <= B * (N + L), local_dim)
valid_token_indices: (total <= B * (N + L),)
B: Batch size
M: Number of visual tokens per context parallel rank
L: Number of text tokens
dtype: Data type of the input and output tensors
Returns:
x: (B, N, dim_x) tensor of visual tokens where N = M
y: (B, L, dim_y) tensor of text token features
"""
local_heads = self.num_heads
local_dim = local_heads * self.head_dim
N = M
# Split sequence into visual and text tokens, adding back padding.
if B == 1:
out = out.view(B, -1, local_dim)
if out.size(1) > N:
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
else:
# Empty prompt.
x, y = out, out.new_zeros(B, L, local_dim)
else:
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
assert x.size() == (B, N, local_dim)
assert y.size() == (B, L, local_dim)
# Communicate across context parallel ranks.
x = x.view(B, N, local_heads, self.head_dim)
x = x.view(x.size(0), x.size(1), x.size(2) * x.size(3)) # (B, M, dim_x = num_heads * head_dim)
x = self.proj_x(x)
y = self.proj_y(y)
return x, y
def forward(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
packed_indices: Dict[str, torch.Tensor] = None,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**rope_rotation,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass of asymmetric multi-modal attention.
Args:
x: (B, M, dim_x) tensor of visual tokens
y: (B, L, dim_y) tensor of text token features
packed_indices: Dict with keys for Flash Attention
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
y: (B, L, dim_y) tensor of text token features after multi-modal attention
"""
B, L, _ = y.shape
_, M, _ = x.shape
# Predict a packed QKV tensor from visual and text features.
q, k, v = ck(self.prepare_qkv,
x=x,
y=y,
scale_x=scale_x,
scale_y=scale_y,
rope_cos=rope_rotation.get("rope_cos"),
rope_sin=rope_rotation.get("rope_sin"),
valid_token_indices=packed_indices["valid_token_indices_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
enabled=checkpoint_qkv,
) # (total <= B * (N + L), 3, local_heads, head_dim)
# Self-attention is expensive, so don't checkpoint it.
out = self.run_attention(
q, k, v, B=B,
cu_seqlens=packed_indices["cu_seqlens_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
)
x, y = ck(self.post_attention,
out,
B=B, M=M, L=L,
dtype=v.dtype,
valid_token_indices=packed_indices["valid_token_indices_kv"],
enabled=checkpoint_post_attn,
)
return x, y
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
class AsymmetricJointBlock(nn.Module):
def __init__(
self,
hidden_size_x: int,
hidden_size_y: int,
num_heads: int,
*,
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
update_y: bool = True, # Whether to update text tokens in this block.
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.update_y = update_y
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
if self.update_y:
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
else:
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
# Self-attention:
self.attn = AsymmetricAttention(
hidden_size_x,
hidden_size_y,
num_heads=num_heads,
update_y=update_y,
device=device,
**block_kwargs,
)
# MLP.
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
assert mlp_hidden_dim_x == int(1536 * 8)
self.mlp_x = FeedForward(
in_features=hidden_size_x,
hidden_size=mlp_hidden_dim_x,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
# MLP for text not needed in last block.
if self.update_y:
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
self.mlp_y = FeedForward(
in_features=hidden_size_y,
hidden_size=mlp_hidden_dim_y,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor,
y: torch.Tensor,
# TODO: These could probably just go into attn_kwargs
checkpoint_ff: bool = False,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**attn_kwargs,
):
"""Forward pass of a block.
Args:
x: (B, N, dim) tensor of visual tokens
c: (B, dim) tensor of conditioned features
y: (B, L, dim) tensor of text tokens
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, N, dim) tensor of visual tokens after block
y: (B, L, dim) tensor of text tokens after block
"""
N = x.size(1)
c = F.silu(c)
mod_x = self.mod_x(c)
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
mod_y = self.mod_y(c)
if self.update_y:
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
else:
scale_msa_y = mod_y
# Self-attention block.
x_attn, y_attn = self.attn(
x,
y,
scale_x=scale_msa_x,
scale_y=scale_msa_y,
checkpoint_qkv=checkpoint_qkv,
checkpoint_post_attn=checkpoint_post_attn,
**attn_kwargs,
)
assert x_attn.size(1) == N
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
if self.update_y:
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
# MLP block.
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
if self.update_y:
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
return x, y
def ff_block_x(self, x, scale_x, gate_x):
x_mod = modulated_rmsnorm(x, scale_x)
x_res = self.mlp_x(x_mod)
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
return x
def ff_block_y(self, y, scale_y, gate_y):
y_mod = modulated_rmsnorm(y, scale_y)
y_res = self.mlp_y(y_mod)
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
return y
@torch.compile(disable=not COMPILE_FINAL_LAYER)
class FinalLayer(nn.Module):
"""
The final layer of DiT.
"""
def __init__(
self,
hidden_size,
patch_size,
out_channels,
device: Optional[torch.device] = None,
):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
def forward(self, x, c):
c = F.silu(c)
shift, scale = self.mod(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
return x
class AsymmDiTJoint(ModelMixin, ConfigMixin, PeftAdapterMixin):
"""
Diffusion model with a Transformer backbone.
Ingests text embeddings instead of a label.
"""
@register_to_config
def __init__(
self,
*,
patch_size=2,
in_channels=12,
hidden_size_x=3072,
hidden_size_y=1536,
depth=48,
num_heads=24,
mlp_ratio_x=4.0,
mlp_ratio_y=4.0,
t5_feat_dim: int = 4096,
t5_token_length: int = 256,
patch_embed_bias: bool = True,
timestep_mlp_bias: bool = True,
timestep_scale: float = 1000.0,
use_extended_posenc: bool = False,
rope_theta: float = 10000.0,
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = in_channels
self.patch_size = patch_size
self.num_heads = num_heads
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
self.use_extended_posenc = use_extended_posenc
self.t5_token_length = t5_token_length
self.t5_feat_dim = t5_feat_dim
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
self.timestep_scale = timestep_scale
self.x_embedder = PatchEmbed(
patch_size=patch_size,
in_chans=in_channels,
embed_dim=hidden_size_x,
bias=patch_embed_bias,
device=device,
)
# Conditionings
# Timestep
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
# Caption Pooling (T5)
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
# Dense Embedding Projection (T5)
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
# Initialize pos_frequencies as an empty parameter.
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
# for depth 48:
# b = 0: AsymmetricJointBlock, update_y=True
# b = 1: AsymmetricJointBlock, update_y=True
# ...
# b = 46: AsymmetricJointBlock, update_y=True
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
blocks = []
for b in range(depth):
# Joint multi-modal block
update_y = b < depth - 1
block = AsymmetricJointBlock(
hidden_size_x,
hidden_size_y,
num_heads,
mlp_ratio_x=mlp_ratio_x,
mlp_ratio_y=mlp_ratio_y,
update_y=update_y,
device=device,
**block_kwargs,
)
blocks.append(block)
self.blocks = nn.ModuleList(blocks)
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: (B, C=12, T, H, W) tensor of visual tokens
Returns:
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
"""
return self.x_embedder(x) # Convert BcTHW to BCN
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
def prepare(
self,
x: torch.Tensor,
sigma: torch.Tensor,
t5_feat: torch.Tensor,
t5_mask: torch.Tensor,
):
"""Prepare input and conditioning embeddings."""
# Visual patch embeddings with positional encoding.
T, H, W = x.shape[-3:]
pH, pW = H // self.patch_size, W // self.patch_size
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
assert x.ndim == 3
B = x.size(0)
# Construct position array of size [N, 3].
# pos[:, 0] is the frame index for each location,
# pos[:, 1] is the row index for each location, and
# pos[:, 2] is the column index for each location.
N = T * pH * pW
assert x.size(1) == N
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
rope_cos, rope_sin = compute_mixed_rotation(
freqs=self.pos_frequencies, pos=pos
) # Each are (N, num_heads, dim // 2)
# Global vector embedding for conditionings.
c_t = self.t_embedder(1 - sigma) # (B, D)
# Pool T5 tokens using attention pooler
# Note encoder_hidden_states[1] contains T5 token features.
assert (
t5_feat.size(1) == self.t5_token_length
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for encoder_hidden_states."
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
c = c_t + t5_y_pool
encoder_hidden_states = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
return x, c, encoder_hidden_states, rope_cos, rope_sin
def forward(
self,
hidden_states: torch.Tensor, # [1, 12, 7, 60, 106]
timestep: torch.Tensor, # [1]
encoder_hidden_states: torch.Tensor, # [1, 256, 4096]
encoder_attention_mask: torch.Tensor, # [1, 256]
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
rope_cos: torch.Tensor = None,
rope_sin: torch.Tensor = None,
num_ff_checkpoint: int = 48, # 48
num_qkv_checkpoint: int = 48, # 48
num_post_attn_checkpoint: int = 0, # 0
):
"""Forward pass of DiT.
Args:
hidden_states: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
sigma: (B,) tensor of noise standard deviations
encoder_hidden_states: List((B, L, encoder_hidden_states_dim) tensor of caption token features. For SDXL text encoders: L=77, encoder_hidden_states_dim=2048)
encoder_attention_mask: List((B, L) boolean tensor indicating which tokens are not padding)
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
{'cu_seqlens_kv': tensor([ 0, 11230], device='cuda:0', dtype=torch.int32),
'max_seqlen_in_batch_kv': 11230,
'valid_token_indices_kv': tensor([ 0, 1, 2, ..., 11227, 11228, 11229], device='cuda:0')}
"""
sigma = timestep / self.timestep_scale
num_latent_toks = np.prod(hidden_states.shape[-3:])
packed_indices = compute_packed_indices(hidden_states.device, encoder_attention_mask, int(num_latent_toks))
_, _, T, H, W = hidden_states.shape
if self.pos_frequencies.dtype != torch.float32:
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
# Have to call sdpa_kernel outside of a torch.compile region.
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
hidden_states, c, encoder_hidden_states, rope_cos, rope_sin = self.prepare(hidden_states, sigma, encoder_hidden_states, encoder_attention_mask) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
del encoder_attention_mask
for i, block in enumerate(self.blocks):
hidden_states, encoder_hidden_states = block( # [1, 11130, 3072], [1, 256, 1536]
hidden_states,
c,
encoder_hidden_states,
rope_cos=rope_cos,
rope_sin=rope_sin,
packed_indices=packed_indices,
checkpoint_ff=i < num_ff_checkpoint,
checkpoint_qkv=i < num_qkv_checkpoint,
checkpoint_post_attn=i < num_post_attn_checkpoint,
) # (B, M, D), (B, L, D)
del encoder_hidden_states # Final layers don't use dense text features.
hidden_states = self.final_layer(hidden_states, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
hidden_states = rearrange( # [1, 12, 7, 60, 106]
hidden_states,
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
T=T,
hp=H // self.patch_size,
wp=W // self.patch_size,
p1=self.patch_size,
p2=self.patch_size,
c=self.out_channels,
)
attn_outputs_list = None
return (-hidden_states, attn_outputs_list)
@@ -0,0 +1,737 @@
import os
from typing import Dict, List, Optional, Tuple
import warnings
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.attention import sdpa_kernel
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
from genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
from genmo.mochi_preview.dit.joint_model.layers import (
FeedForward,
PatchEmbed,
RMSNorm,
TimestepEmbedder,
)
from genmo.mochi_preview.dit.joint_model.lora import LoraLinear
from genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
from genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
residual_tanh_gated_rmsnorm,
)
from genmo.mochi_preview.dit.joint_model.rope_mixed import (
compute_mixed_rotation,
create_position_matrix,
)
from genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
from genmo.mochi_preview.dit.joint_model.utils import (
AttentionPool,
modulate,
pad_and_split_xy,
)
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
if enabled:
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
return fn(*args, **kwargs)
class AsymmetricAttention(nn.Module):
def __init__(
self,
dim_x: int,
dim_y: int,
num_heads: int = 8,
qkv_bias: bool = True,
qk_norm: bool = False,
update_y: bool = True,
out_bias: bool = True,
attention_mode: str = "flash",
softmax_scale: Optional[float] = None,
device: Optional[torch.device] = None,
# Disable LoRA by default ...
qkv_proj_lora_rank: int = 0,
qkv_proj_lora_alpha: int = 0,
qkv_proj_lora_dropout: float = 0.0,
out_proj_lora_rank: int = 0,
out_proj_lora_alpha: int = 0,
out_proj_lora_dropout: float = 0.0,
):
super().__init__()
self.attention_mode = attention_mode
self.dim_x = dim_x
self.dim_y = dim_y
self.num_heads = num_heads
self.head_dim = dim_x // num_heads
self.update_y = update_y
self.softmax_scale = softmax_scale
if dim_x % num_heads != 0:
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
# Input layers.
self.qkv_bias = qkv_bias
qkv_lora_kwargs = dict(
bias=qkv_bias,
device=device,
r=qkv_proj_lora_rank,
lora_alpha=qkv_proj_lora_alpha,
lora_dropout=qkv_proj_lora_dropout,
)
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
# Project text features to match visual features (dim_y -> dim_x)
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
# Query and key normalization for stability.
assert qk_norm
self.q_norm_x = RMSNorm(self.head_dim, device=device)
self.k_norm_x = RMSNorm(self.head_dim, device=device)
self.q_norm_y = RMSNorm(self.head_dim, device=device)
self.k_norm_y = RMSNorm(self.head_dim, device=device)
# Output layers. y features go back down from dim_x -> dim_y.
proj_lora_kwargs = dict(
bias=out_bias,
device=device,
r=out_proj_lora_rank,
lora_alpha=out_proj_lora_alpha,
lora_dropout=out_proj_lora_dropout,
)
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
def run_qkv_y(self, y):
cp_rank, cp_size = cp.get_cp_rank_size()
local_heads = self.num_heads // cp_size
if cp.is_cp_active():
# Only predict local heads.
assert not self.qkv_bias
W_qkv_y = self.qkv_y.weight.view(3, self.num_heads, self.head_dim, self.dim_y)
W_qkv_y = W_qkv_y.narrow(1, cp_rank * local_heads, local_heads)
W_qkv_y = W_qkv_y.reshape(3 * local_heads * self.head_dim, self.dim_y)
qkv_y = F.linear(y, W_qkv_y, None) # (B, L, 3 * local_h * head_dim)
else:
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
q_y, k_y, v_y = qkv_y.unbind(2)
q_y = self.q_norm_y(q_y)
k_y = self.k_norm_y(k_y)
return q_y, k_y, v_y
def prepare_qkv(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor,
scale_y: torch.Tensor,
rope_cos: torch.Tensor,
rope_sin: torch.Tensor,
valid_token_indices: torch.Tensor,
max_seqlen_in_batch: int,
):
# Process visual features
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N / cp_group_size
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
assert qkv_x.dtype == torch.bfloat16
qkv_x = cp.all_to_all_collect_tokens(qkv_x, self.num_heads) # (3, B, N, local_h, head_dim)
# Split qkv_x into q, k, v
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
q_x = self.q_norm_x(q_x)
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
k_x = self.k_norm_x(k_x)
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
# Concatenate streams
B, N, num_heads, head_dim = q_x.size()
D = num_heads * head_dim
# Process text features
if B == 1:
text_seqlen = max_seqlen_in_batch - N
if text_seqlen > 0:
y = y[:, :text_seqlen] # Remove padding tokens.
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
q = torch.cat([q_x, q_y], dim=1)
k = torch.cat([k_x, k_y], dim=1)
v = torch.cat([v_x, v_y], dim=1)
else:
q, k, v = q_x, k_x, v_x
else:
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
indices = valid_token_indices[:, None].expand(-1, D)
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
q = q.view(-1, num_heads, head_dim)
k = k.view(-1, num_heads, head_dim)
v = v.view(-1, num_heads, head_dim)
return q, k, v
@torch.autocast("cuda", enabled=False)
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
out: torch.Tensor = flash_varlen_attn(
q, k, v,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_seqlen_in_batch,
max_seqlen_k=max_seqlen_in_batch,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
) # (total, local_heads, head_dim)
return out.view(total, local_dim)
def sdpa_attention(self, q, k, v):
with sdpa_attn_ctx(training=self.training):
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
)
return out
@torch.autocast("cuda", enabled=False)
def sage_attention(self, q, k, v):
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
def run_attention(
self,
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
*,
B: int,
cu_seqlens: Optional[torch.Tensor] = None,
max_seqlen_in_batch: Optional[int] = None,
):
_, cp_size = cp.get_cp_rank_size()
assert self.num_heads % cp_size == 0
local_heads = self.num_heads // cp_size
local_dim = local_heads * self.head_dim
# Check shapes
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
total = q.size(0)
assert k.size(0) == total and v.size(0) == total
if self.attention_mode == "flash":
out = self.flash_attention(
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
else:
assert B == 1, \
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
q = rearrange(q, "(b s) h d -> b h s d", b=B)
k = rearrange(k, "(b s) h d -> b h s d", b=B)
v = rearrange(v, "(b s) h d -> b h s d", b=B)
if self.attention_mode == "sdpa":
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
elif self.attention_mode == "sage":
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
else:
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
out = rearrange(out, "b h s d -> (b s) (h d)")
return out
def post_attention(
self,
out: torch.Tensor,
B: int,
M: int,
L: int,
dtype: torch.dtype,
valid_token_indices: torch.Tensor,
):
"""
Args:
out: (total <= B * (N + L), local_dim)
valid_token_indices: (total <= B * (N + L),)
B: Batch size
M: Number of visual tokens per context parallel rank
L: Number of text tokens
dtype: Data type of the input and output tensors
Returns:
x: (B, N, dim_x) tensor of visual tokens where N = M * cp_size
y: (B, L, dim_y) tensor of text token features
"""
_, cp_size = cp.get_cp_rank_size()
local_heads = self.num_heads // cp_size
local_dim = local_heads * self.head_dim
N = M * cp_size
# Split sequence into visual and text tokens, adding back padding.
if B == 1:
out = out.view(B, -1, local_dim)
if out.size(1) > N:
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
else:
# Empty prompt.
x, y = out, out.new_zeros(B, L, local_dim)
else:
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
assert x.size() == (B, N, local_dim)
assert y.size() == (B, L, local_dim)
# Communicate across context parallel ranks.
x = x.view(B, N, local_heads, self.head_dim)
x = cp.all_to_all_collect_heads(x) # (B, M, dim_x = num_heads * head_dim)
if cp.is_cp_active():
y = cp.all_gather(y) # (cp_size * B, L, local_heads * head_dim)
y = rearrange(y, "(G B) L D -> B L (G D)", G=cp_size, D=local_dim) # (B, L, dim_x)
x = self.proj_x(x)
y = self.proj_y(y)
return x, y
def forward(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
packed_indices: Dict[str, torch.Tensor] = None,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**rope_rotation,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass of asymmetric multi-modal attention.
Args:
x: (B, M, dim_x) tensor of visual tokens
y: (B, L, dim_y) tensor of text token features
packed_indices: Dict with keys for Flash Attention
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
y: (B, L, dim_y) tensor of text token features after multi-modal attention
"""
B, L, _ = y.shape
_, M, _ = x.shape
# Predict a packed QKV tensor from visual and text features.
q, k, v = ck(self.prepare_qkv,
x=x,
y=y,
scale_x=scale_x,
scale_y=scale_y,
rope_cos=rope_rotation.get("rope_cos"),
rope_sin=rope_rotation.get("rope_sin"),
valid_token_indices=packed_indices["valid_token_indices_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
enabled=checkpoint_qkv,
) # (total <= B * (N + L), 3, local_heads, head_dim)
# Self-attention is expensive, so don't checkpoint it.
out = self.run_attention(
q, k, v, B=B,
cu_seqlens=packed_indices["cu_seqlens_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
)
x, y = ck(self.post_attention,
out,
B=B, M=M, L=L,
dtype=v.dtype,
valid_token_indices=packed_indices["valid_token_indices_kv"],
enabled=checkpoint_post_attn,
)
return x, y
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
class AsymmetricJointBlock(nn.Module):
def __init__(
self,
hidden_size_x: int,
hidden_size_y: int,
num_heads: int,
*,
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
update_y: bool = True, # Whether to update text tokens in this block.
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.update_y = update_y
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
if self.update_y:
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
else:
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
# Self-attention:
self.attn = AsymmetricAttention(
hidden_size_x,
hidden_size_y,
num_heads=num_heads,
update_y=update_y,
device=device,
**block_kwargs,
)
# MLP.
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
assert mlp_hidden_dim_x == int(1536 * 8)
self.mlp_x = FeedForward(
in_features=hidden_size_x,
hidden_size=mlp_hidden_dim_x,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
# MLP for text not needed in last block.
if self.update_y:
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
self.mlp_y = FeedForward(
in_features=hidden_size_y,
hidden_size=mlp_hidden_dim_y,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor,
y: torch.Tensor,
# TODO: These could probably just go into attn_kwargs
checkpoint_ff: bool = False,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**attn_kwargs,
):
"""Forward pass of a block.
Args:
x: (B, N, dim) tensor of visual tokens
c: (B, dim) tensor of conditioned features
y: (B, L, dim) tensor of text tokens
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, N, dim) tensor of visual tokens after block
y: (B, L, dim) tensor of text tokens after block
"""
N = x.size(1)
c = F.silu(c)
mod_x = self.mod_x(c)
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
mod_y = self.mod_y(c)
if self.update_y:
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
else:
scale_msa_y = mod_y
# Self-attention block.
x_attn, y_attn = self.attn(
x,
y,
scale_x=scale_msa_x,
scale_y=scale_msa_y,
checkpoint_qkv=checkpoint_qkv,
checkpoint_post_attn=checkpoint_post_attn,
**attn_kwargs,
)
assert x_attn.size(1) == N
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
if self.update_y:
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
# MLP block.
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
if self.update_y:
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
return x, y
def ff_block_x(self, x, scale_x, gate_x):
x_mod = modulated_rmsnorm(x, scale_x)
x_res = self.mlp_x(x_mod)
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
return x
def ff_block_y(self, y, scale_y, gate_y):
y_mod = modulated_rmsnorm(y, scale_y)
y_res = self.mlp_y(y_mod)
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
return y
@torch.compile(disable=not COMPILE_FINAL_LAYER)
class FinalLayer(nn.Module):
"""
The final layer of DiT.
"""
def __init__(
self,
hidden_size,
patch_size,
out_channels,
device: Optional[torch.device] = None,
):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
def forward(self, x, c):
c = F.silu(c)
shift, scale = self.mod(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
return x
class AsymmDiTJoint(nn.Module):
"""
Diffusion model with a Transformer backbone.
Ingests text embeddings instead of a label.
"""
def __init__(
self,
*,
patch_size=2,
in_channels=4,
hidden_size_x=1152,
hidden_size_y=1152,
depth=48,
num_heads=16,
mlp_ratio_x=8.0,
mlp_ratio_y=4.0,
t5_feat_dim: int = 4096,
t5_token_length: int = 256,
patch_embed_bias: bool = True,
timestep_mlp_bias: bool = True,
timestep_scale: Optional[float] = None,
use_extended_posenc: bool = False,
rope_theta: float = 10000.0,
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = in_channels
self.patch_size = patch_size
self.num_heads = num_heads
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
self.use_extended_posenc = use_extended_posenc
self.t5_token_length = t5_token_length
self.t5_feat_dim = t5_feat_dim
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
self.x_embedder = PatchEmbed(
patch_size=patch_size,
in_chans=in_channels,
embed_dim=hidden_size_x,
bias=patch_embed_bias,
device=device,
)
# Conditionings
# Timestep
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
# Caption Pooling (T5)
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
# Dense Embedding Projection (T5)
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
# Initialize pos_frequencies as an empty parameter.
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
# for depth 48:
# b = 0: AsymmetricJointBlock, update_y=True
# b = 1: AsymmetricJointBlock, update_y=True
# ...
# b = 46: AsymmetricJointBlock, update_y=True
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
blocks = []
for b in range(depth):
# Joint multi-modal block
update_y = b < depth - 1
block = AsymmetricJointBlock(
hidden_size_x,
hidden_size_y,
num_heads,
mlp_ratio_x=mlp_ratio_x,
mlp_ratio_y=mlp_ratio_y,
update_y=update_y,
device=device,
**block_kwargs,
)
blocks.append(block)
self.blocks = nn.ModuleList(blocks)
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: (B, C=12, T, H, W) tensor of visual tokens
Returns:
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
"""
return self.x_embedder(x) # Convert BcTHW to BCN
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
def prepare(
self,
x: torch.Tensor,
sigma: torch.Tensor,
t5_feat: torch.Tensor,
t5_mask: torch.Tensor,
):
"""Prepare input and conditioning embeddings."""
# Visual patch embeddings with positional encoding.
T, H, W = x.shape[-3:]
pH, pW = H // self.patch_size, W // self.patch_size
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
assert x.ndim == 3
B = x.size(0)
# Construct position array of size [N, 3].
# pos[:, 0] is the frame index for each location,
# pos[:, 1] is the row index for each location, and
# pos[:, 2] is the column index for each location.
N = T * pH * pW
assert x.size(1) == N
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
rope_cos, rope_sin = compute_mixed_rotation(
freqs=self.pos_frequencies, pos=pos
) # Each are (N, num_heads, dim // 2)
# Global vector embedding for conditionings.
c_t = self.t_embedder(1 - sigma) # (B, D)
# Pool T5 tokens using attention pooler
# Note y_feat[1] contains T5 token features.
assert (
t5_feat.size(1) == self.t5_token_length
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for y_feat."
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
c = c_t + t5_y_pool
y_feat = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
return x, c, y_feat, rope_cos, rope_sin
def forward(
self,
x: torch.Tensor, # [1, 12, 7, 60, 106]
sigma: torch.Tensor, # [1]
y_feat: List[torch.Tensor], # [0][1, 256, 4096]
y_mask: List[torch.Tensor], # [0][1, 256]
packed_indices: Dict[str, torch.Tensor] = None,
rope_cos: torch.Tensor = None,
rope_sin: torch.Tensor = None,
num_ff_checkpoint: int = 0, # 48
num_qkv_checkpoint: int = 0, # 48
num_post_attn_checkpoint: int = 0, # 0
):
"""Forward pass of DiT.
Args:
x: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
sigma: (B,) tensor of noise standard deviations
y_feat: List((B, L, y_feat_dim) tensor of caption token features. For SDXL text encoders: L=77, y_feat_dim=2048)
y_mask: List((B, L) boolean tensor indicating which tokens are not padding)
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
"""
_, _, T, H, W = x.shape
if self.pos_frequencies.dtype != torch.float32:
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
# Have to call sdpa_kernel outside of a torch.compile region.
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
x, c, y_feat, rope_cos, rope_sin = self.prepare(x, sigma, y_feat[0], y_mask[0]) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
del y_mask
cp_rank, cp_size = cp.get_cp_rank_size()
N = x.size(1)
M = N // cp_size
assert N % cp_size == 0, f"Visual sequence length ({x.shape[1]}) must be divisible by cp_size ({cp_size})."
if cp_size > 1:
x = x.narrow(1, cp_rank * M, M)
assert self.num_heads % cp_size == 0
local_heads = self.num_heads // cp_size
rope_cos = rope_cos.narrow(1, cp_rank * local_heads, local_heads)
rope_sin = rope_sin.narrow(1, cp_rank * local_heads, local_heads)
for i, block in enumerate(self.blocks):
x, y_feat = block( # [1, 11130, 3072], [1, 256, 1536]
x,
c,
y_feat,
rope_cos=rope_cos,
rope_sin=rope_sin,
packed_indices=packed_indices,
checkpoint_ff=i < num_ff_checkpoint,
checkpoint_qkv=i < num_qkv_checkpoint,
checkpoint_post_attn=i < num_post_attn_checkpoint,
) # (B, M, D), (B, L, D)
del y_feat # Final layers don't use dense text features.
x = self.final_layer(x, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
patch = x.size(2)
x = cp.all_gather(x)
x = rearrange(x, "(G B) M P -> B (G M) P", G=cp_size, P=patch) # [1, 11130, 48]
x = rearrange( # [1, 12, 7, 60, 106]
x,
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
T=T,
hp=H // self.patch_size,
wp=W // self.patch_size,
p1=self.patch_size,
p2=self.patch_size,
c=self.out_channels,
)
return x
@@ -0,0 +1,158 @@
from typing import Tuple
import torch
import torch.distributed as dist
from einops import rearrange
_CONTEXT_PARALLEL_GROUP = None
_CONTEXT_PARALLEL_RANK = None
_CONTEXT_PARALLEL_GROUP_SIZE = None
_CONTEXT_PARALLEL_GROUP_RANKS = None
def get_cp_rank_size() -> Tuple[int, int]:
if _CONTEXT_PARALLEL_GROUP:
assert isinstance(_CONTEXT_PARALLEL_RANK, int) and isinstance(_CONTEXT_PARALLEL_GROUP_SIZE, int)
return _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE
else:
return 0, 1
def local_shard(x: torch.Tensor, dim: int = 2) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
return x
cp_rank, cp_size = get_cp_rank_size()
return x.tensor_split(cp_size, dim=dim)[cp_rank]
def set_cp_group(cp_group, ranks, global_rank):
global _CONTEXT_PARALLEL_GROUP, _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE, _CONTEXT_PARALLEL_GROUP_RANKS
if _CONTEXT_PARALLEL_GROUP is not None:
raise RuntimeError("CP group already initialized.")
_CONTEXT_PARALLEL_GROUP = cp_group
_CONTEXT_PARALLEL_RANK = dist.get_rank(cp_group)
_CONTEXT_PARALLEL_GROUP_SIZE = dist.get_world_size(cp_group)
_CONTEXT_PARALLEL_GROUP_RANKS = ranks
assert _CONTEXT_PARALLEL_RANK == ranks.index(
global_rank
), f"Rank mismatch: {global_rank} in {ranks} does not have position {_CONTEXT_PARALLEL_RANK} "
assert _CONTEXT_PARALLEL_GROUP_SIZE == len(
ranks
), f"Group size mismatch: {_CONTEXT_PARALLEL_GROUP_SIZE} != len({ranks})"
def get_cp_group():
if _CONTEXT_PARALLEL_GROUP is None:
raise RuntimeError("CP group not initialized")
return _CONTEXT_PARALLEL_GROUP
def is_cp_active():
return _CONTEXT_PARALLEL_GROUP is not None
class AllGatherIntoTensorFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x: torch.Tensor, reduce_dtype, group: dist.ProcessGroup):
ctx.reduce_dtype = reduce_dtype
ctx.group = group
ctx.batch_size = x.size(0)
group_size = dist.get_world_size(group)
x = x.contiguous()
output = torch.empty(group_size * x.size(0), *x.shape[1:], dtype=x.dtype, device=x.device)
dist.all_gather_into_tensor(output, x, group=group)
return output
def all_gather(tensor: torch.Tensor) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
return tensor
return AllGatherIntoTensorFunction.apply(tensor, torch.float32, _CONTEXT_PARALLEL_GROUP)
@torch.compiler.disable()
def _all_to_all_single(output, input, group):
# Disable compilation since torch compile changes contiguity.
assert input.is_contiguous(), "Input tensor must be contiguous."
assert output.is_contiguous(), "Output tensor must be contiguous."
return dist.all_to_all_single(output, input, group=group)
class CollectTokens(torch.autograd.Function):
@staticmethod
def forward(ctx, qkv: torch.Tensor, group: dist.ProcessGroup, num_heads: int):
"""Redistribute heads and receive tokens.
Args:
qkv: query, key or value. Shape: [B, M, 3 * num_heads * head_dim]
Returns:
qkv: shape: [3, B, N, local_heads, head_dim]
where M is the number of local tokens,
N = cp_size * M is the number of global tokens,
local_heads = num_heads // cp_size is the number of local heads.
"""
ctx.group = group
ctx.num_heads = num_heads
cp_size = dist.get_world_size(group)
assert num_heads % cp_size == 0
ctx.local_heads = num_heads // cp_size
qkv = rearrange(
qkv,
"B M (qkv G h d) -> G M h B (qkv d)",
qkv=3,
G=cp_size,
h=ctx.local_heads,
).contiguous()
output_chunks = torch.empty_like(qkv)
_all_to_all_single(output_chunks, qkv, group=group)
return rearrange(output_chunks, "G M h B (qkv d) -> qkv B (G M) h d", qkv=3)
def all_to_all_collect_tokens(x: torch.Tensor, num_heads: int) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
# Move QKV dimension to the front.
# B M (3 H d) -> 3 B M H d
B, M, _ = x.size()
x = x.view(B, M, 3, num_heads, -1)
return x.permute(2, 0, 1, 3, 4)
return CollectTokens.apply(x, _CONTEXT_PARALLEL_GROUP, num_heads)
class CollectHeads(torch.autograd.Function):
@staticmethod
def forward(ctx, x: torch.Tensor, group: dist.ProcessGroup):
"""Redistribute tokens and receive heads.
Args:
x: Output of attention. Shape: [B, N, local_heads, head_dim]
Returns:
Shape: [B, M, num_heads * head_dim]
"""
ctx.group = group
ctx.local_heads = x.size(2)
ctx.head_dim = x.size(3)
group_size = dist.get_world_size(group)
x = rearrange(x, "B (G M) h D -> G h M B D", G=group_size).contiguous()
output = torch.empty_like(x)
_all_to_all_single(output, x, group=group)
del x
return rearrange(output, "G h M B D -> B M (G h D)")
def all_to_all_collect_heads(x: torch.Tensor) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
# Merge heads.
return x.view(x.size(0), x.size(1), x.size(2) * x.size(3))
return CollectHeads.apply(x, _CONTEXT_PARALLEL_GROUP)
@@ -0,0 +1,179 @@
import collections.abc
import math
from itertools import repeat
from typing import Callable, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
# From PyTorch internals
def _ntuple(n):
def parse(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
return tuple(x)
return tuple(repeat(x, n))
return parse
to_2tuple = _ntuple(2)
class TimestepEmbedder(nn.Module):
def __init__(
self,
hidden_size: int,
frequency_embedding_size: int = 256,
*,
bias: bool = True,
timestep_scale: Optional[float] = None,
device: Optional[torch.device] = None,
):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, hidden_size, bias=bias, device=device),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
)
self.frequency_embedding_size = frequency_embedding_size
self.timestep_scale = timestep_scale
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
half = dim // 2
freqs = torch.arange(start=0, end=half, dtype=torch.float32, device=t.device)
freqs.mul_(-math.log(max_period) / half).exp_()
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t):
if self.timestep_scale is not None:
t = t * self.timestep_scale
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
t_emb = self.mlp(t_freq)
return t_emb
class PooledCaptionEmbedder(nn.Module):
def __init__(
self,
caption_feature_dim: int,
hidden_size: int,
*,
bias: bool = True,
device: Optional[torch.device] = None,
):
super().__init__()
self.caption_feature_dim = caption_feature_dim
self.hidden_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(caption_feature_dim, hidden_size, bias=bias, device=device),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
)
def forward(self, x):
return self.mlp(x)
class FeedForward(nn.Module):
def __init__(
self,
in_features: int,
hidden_size: int,
multiple_of: int,
ffn_dim_multiplier: Optional[float],
device: Optional[torch.device] = None,
):
super().__init__()
# keep parameter count and computation constant compared to standard FFN
hidden_size = int(2 * hidden_size / 3)
# custom dim factor multiplier
if ffn_dim_multiplier is not None:
hidden_size = int(ffn_dim_multiplier * hidden_size)
hidden_size = multiple_of * ((hidden_size + multiple_of - 1) // multiple_of)
self.hidden_dim = hidden_size
self.w1 = nn.Linear(in_features, 2 * hidden_size, bias=False, device=device)
self.w2 = nn.Linear(hidden_size, in_features, bias=False, device=device)
def forward(self, x):
# assert self.w1.weight.dtype == torch.bfloat16, f"FFN weight dtype {self.w1.weight.dtype} != bfloat16"
x, gate = self.w1(x).chunk(2, dim=-1)
x = self.w2(F.silu(x) * gate)
return x
class PatchEmbed(nn.Module):
def __init__(
self,
patch_size: int = 16,
in_chans: int = 3,
embed_dim: int = 768,
norm_layer: Optional[Callable] = None,
flatten: bool = True,
bias: bool = True,
dynamic_img_pad: bool = False,
device: Optional[torch.device] = None,
):
super().__init__()
self.patch_size = to_2tuple(patch_size)
self.flatten = flatten
self.dynamic_img_pad = dynamic_img_pad
self.proj = nn.Conv2d(
in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
device=device,
)
assert norm_layer is None
self.norm = norm_layer(embed_dim, device=device) if norm_layer else nn.Identity()
def forward(self, x):
B, _C, T, H, W = x.shape
if not self.dynamic_img_pad:
assert (
H % self.patch_size[0] == 0
), f"Input height ({H}) should be divisible by patch size ({self.patch_size[0]})."
assert (
W % self.patch_size[1] == 0
), f"Input width ({W}) should be divisible by patch size ({self.patch_size[1]})."
else:
pad_h = (self.patch_size[0] - H % self.patch_size[0]) % self.patch_size[0]
pad_w = (self.patch_size[1] - W % self.patch_size[1]) % self.patch_size[1]
x = F.pad(x, (0, pad_w, 0, pad_h))
x = rearrange(x, "B C T H W -> (B T) C H W", B=B, T=T)
x = self.proj(x)
# Flatten temporal and spatial dimensions.
if not self.flatten:
raise NotImplementedError("Must flatten output.")
x = rearrange(x, "(B T) C H W -> B (T H W) C", B=B, T=T)
x = self.norm(x)
return x
class RMSNorm(torch.nn.Module):
def __init__(self, hidden_size, eps=1e-5, device=None):
super().__init__()
self.eps = eps
self.weight = torch.nn.Parameter(torch.empty(hidden_size, device=device))
self.register_parameter("bias", None)
def forward(self, x):
# assert self.weight.dtype == torch.float32, f"RMSNorm weight dtype {self.weight.dtype} != float32"
x_fp32 = x.float()
x_normed = x_fp32 * torch.rsqrt(x_fp32.pow(2).mean(-1, keepdim=True) + self.eps)
return (x_normed * self.weight).type_as(x)
@@ -0,0 +1,112 @@
#! /usr/bin/env python3
import math
from typing import Dict, List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class LoRALayer:
def __init__(
self,
r: int,
lora_alpha: int,
lora_dropout: float,
merge_weights: bool,
):
self.r = r
self.lora_alpha = lora_alpha
if lora_dropout > 0.0:
self.lora_dropout = nn.Dropout(p=lora_dropout)
else:
self.lora_dropout = lambda x: x
self.merged = False
self.merge_weights = merge_weights
def mark_only_lora_as_trainable(model: nn.Module, bias: str = "none") -> None:
assert bias == "none", f"Only bias='none' is supported"
for n, p in model.named_parameters():
if "lora_" not in n:
p.requires_grad = False
def lora_state_dict(model: nn.Module, bias: str = "none") -> Dict[str, torch.Tensor]:
assert bias == "none", f"Only bias='none' is supported"
my_state_dict = model.state_dict()
return {k: my_state_dict[k] for k in my_state_dict if "lora_" in k}
class LoraLinear(nn.Linear, LoRALayer):
# LoRA implemented in a dense layer
def __init__(
self,
in_features: int,
out_features: int,
r: int = 0,
lora_alpha: int = 1,
lora_dropout: float = 0.0,
fan_in_fan_out: bool = False, # Set this to True if the layer to replace stores weight like (fan_in, fan_out)
merge_weights: bool = True,
**kwargs,
):
nn.Linear.__init__(self, in_features, out_features, **kwargs)
LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)
self.fan_in_fan_out = fan_in_fan_out
# Actual trainable parameters
if r > 0:
self.lora_A = nn.Parameter(self.weight.new_zeros((r, in_features)).to(torch.float32))
self.lora_B = nn.Parameter(self.weight.new_zeros((out_features, r)).to(torch.float32))
self.scaling = self.lora_alpha / self.r
# Freezing the pre-trained weight matrix
self.weight.requires_grad = False
self.reset_parameters()
if fan_in_fan_out:
self.weight.data = self.weight.data.transpose(0, 1)
def reset_parameters(self):
nn.Linear.reset_parameters(self)
if hasattr(self, "lora_A"):
# initialize B the same way as the default for nn.Linear and A to zero
# this is different than what is described in the paper but should not affect performance
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
nn.init.zeros_(self.lora_B)
def train(self, mode: bool = True):
def T(w):
return w.transpose(0, 1) if self.fan_in_fan_out else w
nn.Linear.train(self, mode)
if mode:
if self.merge_weights and self.merged:
# Make sure that the weights are not merged
if self.r > 0:
self.weight.data -= T(self.lora_B @ self.lora_A) * self.scaling
self.merged = False
else:
if self.merge_weights and not self.merged:
# Merge the weights and mark it
if self.r > 0:
self.weight.data += T(self.lora_B @ self.lora_A) * self.scaling
self.merged = True
def forward(self, x: torch.Tensor):
def T(w):
return w.transpose(0, 1) if self.fan_in_fan_out else w
if self.r > 0 and not self.merged:
result = F.linear(x, T(self.weight), bias=self.bias)
x = self.lora_dropout(x)
x = x @ self.lora_A.transpose(0, 1)
x = x @ self.lora_B.transpose(0, 1)
x = x * self.scaling
return result + x
else:
return F.linear(x, T(self.weight), bias=self.bias)
@@ -0,0 +1,15 @@
import torch
def modulated_rmsnorm(x, scale, eps=1e-6):
dtype = x.dtype
x = x.float()
# Compute RMS
mean_square = x.pow(2).mean(-1, keepdim=True)
inv_rms = torch.rsqrt(mean_square + eps)
# Normalize and modulate
x_normed = x * inv_rms
x_modulated = x_normed * (1 + scale.unsqueeze(1).float())
return x_modulated.to(dtype)
@@ -0,0 +1,20 @@
import torch
def residual_tanh_gated_rmsnorm(x, x_res, gate, eps=1e-6):
# Convert to fp32 for precision
x_res = x_res.float()
# Compute RMS
mean_square = x_res.pow(2).mean(-1, keepdim=True)
scale = torch.rsqrt(mean_square + eps)
# Apply tanh to gate
tanh_gate = torch.tanh(gate).unsqueeze(1)
# Normalize and apply gated scaling
x_normed = x_res * scale * tanh_gate
# Apply residual connection
output = x + x_normed.type_as(x)
return output
@@ -0,0 +1,88 @@
import functools
import math
import torch
def centers(start: float, stop, num, dtype=None, device=None):
"""linspace through bin centers.
Args:
start (float): Start of the range.
stop (float): End of the range.
num (int): Number of points.
dtype (torch.dtype): Data type of the points.
device (torch.device): Device of the points.
Returns:
centers (Tensor): Centers of the bins. Shape: (num,).
"""
edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device)
return (edges[:-1] + edges[1:]) / 2
@functools.lru_cache(maxsize=1)
def create_position_matrix(
T: int,
pH: int,
pW: int,
device: torch.device,
dtype: torch.dtype,
*,
target_area: float = 36864,
):
"""
Args:
T: int - Temporal dimension
pH: int - Height dimension after patchify
pW: int - Width dimension after patchify
Returns:
pos: [T * pH * pW, 3] - position matrix
"""
with torch.no_grad():
# Create 1D tensors for each dimension
t = torch.arange(T, dtype=dtype)
# Positionally interpolate to area 36864.
# (3072x3072 frame with 16x16 patches = 192x192 latents).
# This automatically scales rope positions when the resolution changes.
# We use a large target area so the model is more sensitive
# to changes in the learned pos_frequencies matrix.
scale = math.sqrt(target_area / (pW * pH))
w = centers(-pW * scale / 2, pW * scale / 2, pW)
h = centers(-pH * scale / 2, pH * scale / 2, pH)
# Use meshgrid to create 3D grids
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
# Stack and reshape the grids.
pos = torch.stack([grid_t, grid_h, grid_w], dim=-1) # [T, pH, pW, 3]
pos = pos.view(-1, 3) # [T * pH * pW, 3]
pos = pos.to(dtype=dtype, device=device)
return pos
def compute_mixed_rotation(
freqs: torch.Tensor,
pos: torch.Tensor,
):
"""
Project each 3-dim position into per-head, per-head-dim 1D frequencies.
Args:
freqs: [3, num_heads, num_freqs] - learned rotation frequency (for t, row, col) for each head position
pos: [N, 3] - position of each token
num_heads: int
Returns:
freqs_cos: [N, num_heads, num_freqs] - cosine components
freqs_sin: [N, num_heads, num_freqs] - sine components
"""
with torch.autocast("cuda", enabled=False):
assert freqs.ndim == 3
freqs_sum = torch.einsum("Nd,dhf->Nhf", pos.to(freqs), freqs)
freqs_cos = torch.cos(freqs_sum)
freqs_sin = torch.sin(freqs_sum)
return freqs_cos, freqs_sin
@@ -0,0 +1,34 @@
# Based on Llama3 Implementation.
import torch
def apply_rotary_emb_qk_real(
xqk: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
) -> torch.Tensor:
"""
Apply rotary embeddings to input tensors using the given frequency tensor without complex numbers.
Args:
xqk (torch.Tensor): Query and/or Key tensors to apply rotary embeddings. Shape: (B, S, *, num_heads, D)
Can be either just query or just key, or both stacked along some batch or * dim.
freqs_cos (torch.Tensor): Precomputed cosine frequency tensor.
freqs_sin (torch.Tensor): Precomputed sine frequency tensor.
Returns:
torch.Tensor: The input tensor with rotary embeddings applied.
"""
assert xqk.dtype == torch.bfloat16
# Split the last dimension into even and odd parts
xqk_even = xqk[..., 0::2]
xqk_odd = xqk[..., 1::2]
# Apply rotation
cos_part = (xqk_even * freqs_cos - xqk_odd * freqs_sin).type_as(xqk)
sin_part = (xqk_even * freqs_sin + xqk_odd * freqs_cos).type_as(xqk)
# Interleave the results back into the original shape
out = torch.stack([cos_part, sin_part], dim=-1).flatten(-2)
assert out.dtype == torch.bfloat16
return out
@@ -0,0 +1,109 @@
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
def modulate(x, shift, scale):
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor:
"""
Pool tokens in x using mask.
NOTE: We assume x does not require gradients.
Args:
x: (B, L, D) tensor of tokens.
mask: (B, L) boolean tensor indicating which tokens are not padding.
Returns:
pooled: (B, D) tensor of pooled tokens.
"""
assert x.size(1) == mask.size(1) # Expected mask to have same length as tokens.
assert x.size(0) == mask.size(0) # Expected mask to have same batch size as tokens.
mask = mask[:, :, None].to(dtype=x.dtype)
mask = mask / mask.sum(dim=1, keepdim=True).clamp(min=1)
pooled = (x * mask).sum(dim=1, keepdim=keepdim)
return pooled
class AttentionPool(nn.Module):
def __init__(
self,
embed_dim: int,
num_heads: int,
output_dim: int = None,
device: Optional[torch.device] = None,
):
"""
Args:
spatial_dim (int): Number of tokens in sequence length.
embed_dim (int): Dimensionality of input tokens.
num_heads (int): Number of attention heads.
output_dim (int): Dimensionality of output tokens. Defaults to embed_dim.
"""
super().__init__()
self.num_heads = num_heads
self.to_kv = nn.Linear(embed_dim, 2 * embed_dim, device=device)
self.to_q = nn.Linear(embed_dim, embed_dim, device=device)
self.to_out = nn.Linear(embed_dim, output_dim or embed_dim, device=device)
def forward(self, x, mask):
"""
Args:
x (torch.Tensor): (B, L, D) tensor of input tokens.
mask (torch.Tensor): (B, L) boolean tensor indicating which tokens are not padding.
NOTE: We assume x does not require gradients.
Returns:
x (torch.Tensor): (B, D) tensor of pooled tokens.
"""
D = x.size(2)
# Construct attention mask, shape: (B, 1, num_queries=1, num_keys=1+L).
attn_mask = mask[:, None, None, :].bool() # (B, 1, 1, L).
attn_mask = F.pad(attn_mask, (1, 0), value=True) # (B, 1, 1, 1+L).
# Average non-padding token features. These will be used as the query.
x_pool = pool_tokens(x, mask, keepdim=True) # (B, 1, D)
# Concat pooled features to input sequence.
x = torch.cat([x_pool, x], dim=1) # (B, L+1, D)
# Compute queries, keys, values. Only the mean token is used to create a query.
kv = self.to_kv(x) # (B, L+1, 2 * D)
q = self.to_q(x[:, 0]) # (B, D)
# Extract heads.
head_dim = D // self.num_heads
kv = kv.unflatten(2, (2, self.num_heads, head_dim)) # (B, 1+L, 2, H, head_dim)
kv = kv.transpose(1, 3) # (B, H, 2, 1+L, head_dim)
k, v = kv.unbind(2) # (B, H, 1+L, head_dim)
q = q.unflatten(1, (self.num_heads, head_dim)) # (B, H, head_dim)
q = q.unsqueeze(2) # (B, H, 1, head_dim)
# Compute attention.
x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) # (B, H, 1, head_dim)
# Concatenate heads and run output.
x = x.squeeze(2).flatten(1, 2) # (B, D = H * head_dim)
x = self.to_out(x)
return x
def pad_and_split_xy(xy, indices, B, N, L, dtype) -> Tuple[torch.Tensor, torch.Tensor]:
D = xy.size(1)
# Pad sequences to (B, N + L, dim).
assert indices.ndim == 1
indices = indices.unsqueeze(1).expand(-1, D) # (total,) -> (total, num_heads * head_dim)
output = torch.zeros(B * (N + L), D, device=xy.device, dtype=dtype)
output = torch.scatter(output, 0, indices, xy)
xy = output.view(B, N + L, D)
# Split visual and text tokens along the sequence length.
return torch.tensor_split(xy, (N,), dim=1)
@@ -0,0 +1,682 @@
import json
import os
import random
from abc import ABC, abstractmethod
from contextlib import contextmanager
from functools import partial
from typing import Any, Dict, List, Literal, Optional, Union, cast
import numpy as np
import ray
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from einops import repeat
from safetensors import safe_open
from safetensors.torch import load_file
from torch import nn
from torch.distributed.fsdp import (
BackwardPrefetch,
MixedPrecision,
ShardingStrategy,
)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import (
lambda_auto_wrap_policy,
transformer_auto_wrap_policy,
)
from transformers import T5EncoderModel, T5Tokenizer
from transformers.models.t5.modeling_t5 import T5Block
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
from genmo.lib.progress import get_new_progress_bar, progress_bar
from genmo.lib.utils import Timer
from genmo.mochi_preview.vae.models import (
Decoder,
Encoder,
decode_latents,
decode_latents_tiled_full,
decode_latents_tiled_spatial,
)
from genmo.mochi_preview.vae.vae_stats import dit_latents_to_vae_latents
def load_to_cpu(p, weights_only=True):
if p.endswith(".safetensors"):
return load_file(p)
else:
assert p.endswith(".pt")
return torch.load(p, map_location="cpu", weights_only=weights_only)
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
if linear_steps is None:
linear_steps = num_steps // 2
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
quadratic_steps = num_steps - linear_steps
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
const = quadratic_coef * (linear_steps**2)
quadratic_sigma_schedule = [
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, num_steps)
]
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule + [1.0]
sigma_schedule = [1.0 - x for x in sigma_schedule]
return sigma_schedule
T5_MODEL = "google/t5-v1_1-xxl"
MAX_T5_TOKEN_LENGTH = 256
def setup_fsdp_sync(model, device_id, *, param_dtype, auto_wrap_policy) -> FSDP:
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(
param_dtype=param_dtype,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
),
auto_wrap_policy=auto_wrap_policy,
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
limit_all_gathers=True,
device_id=device_id,
sync_module_states=True,
use_orig_params=True,
)
torch.cuda.synchronize()
return model
class ModelFactory(ABC):
def __init__(self, **kwargs):
self.kwargs = kwargs
@abstractmethod
def get_model(self, *, local_rank: int, device_id: Union[int, Literal["cpu"]], world_size: int) -> Any:
assert isinstance(device_id, int) or device_id == "cpu", "device_id must be an integer or 'cpu'"
# FSDP does not work when the model is on the CPU
if device_id == "cpu":
assert world_size == 1, "CPU offload only supports single-GPU inference"
class T5ModelFactory(ModelFactory):
def __init__(self, model_dir=None):
super().__init__()
self.model_dir = model_dir or T5_MODEL
def get_model(self, *, local_rank, device_id, world_size):
super().get_model(local_rank=local_rank, device_id=device_id, world_size=world_size)
model = T5EncoderModel.from_pretrained(self.model_dir)
if world_size > 1:
model = setup_fsdp_sync(
model,
device_id=device_id,
param_dtype=torch.float32,
auto_wrap_policy=partial(
transformer_auto_wrap_policy,
transformer_layer_cls={
T5Block,
},
),
)
elif isinstance(device_id, int):
model = model.to(torch.device(f"cuda:{device_id}")) # type: ignore
return model.eval()
class DitModelFactory(ModelFactory):
def __init__(
self, *,
model_path: str,
model_dtype: str,
lora_path: Optional[str] = None,
attention_mode: Optional[str] = None
):
# Infer attention mode if not specified
if attention_mode is None:
from genmo.lib.attn_imports import flash_varlen_attn # type: ignore
attention_mode = "sdpa" if flash_varlen_attn is None else "flash"
print(f"Attention mode: {attention_mode}")
super().__init__(
model_path=model_path,
lora_path=lora_path,
model_dtype=model_dtype,
attention_mode=attention_mode
)
def get_model(
self,
*,
local_rank,
device_id,
world_size,
model_kwargs=None,
patch_model_fns=None,
strict_load=True,
load_checkpoint=True,
fast_init=True,
):
from genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
if not model_kwargs:
model_kwargs = {}
lora_sd = None
lora_path = self.kwargs["lora_path"]
if lora_path is not None:
if lora_path.endswith(".safetensors"):
lora_sd = {}
with safe_open(lora_path, framework="pt") as f:
for k in f.keys():
lora_sd[k] = f.get_tensor(k)
lora_kwargs = json.loads(f.metadata()["kwargs"])
print(f"Loaded LoRA kwargs: {lora_kwargs}")
else:
lora = load_to_cpu(lora_path, weights_only=False)
lora_sd, lora_kwargs = lora["state_dict"], lora["kwargs"]
model_kwargs.update(cast(dict, lora_kwargs))
model_args = dict(
depth=48,
patch_size=2,
num_heads=24,
hidden_size_x=3072,
hidden_size_y=1536,
mlp_ratio_x=4.0,
mlp_ratio_y=4.0,
in_channels=12,
qk_norm=True,
qkv_bias=False,
out_bias=True,
patch_embed_bias=True,
timestep_mlp_bias=True,
timestep_scale=1000.0,
t5_feat_dim=4096,
t5_token_length=256,
rope_theta=10000.0,
attention_mode=self.kwargs["attention_mode"],
**model_kwargs,
)
if fast_init:
model: nn.Module = torch.nn.utils.skip_init(AsymmDiTJoint, **model_args)
else:
model: nn.Module = AsymmDiTJoint(**model_args)
for fn in patch_model_fns or []:
model = fn(model)
# FSDP syncs weights from rank 0 to all other ranks
if local_rank == 0 and load_checkpoint:
model_path = self.kwargs["model_path"]
sd = load_to_cpu(model_path)
# Load the state dictionary and capture the return value
load_result = model.load_state_dict(sd, strict=strict_load)
if not strict_load:
# Print mismatched keys
missing_keys = [k for k in load_result.missing_keys if ".lora_" not in k]
if missing_keys:
print(f"Missing keys from {model_path}: {missing_keys}")
if load_result.unexpected_keys:
print(f"Unexpected keys from {model_path}: {load_result.unexpected_keys}")
if lora_sd:
model.load_state_dict(lora_sd, strict=strict_load) # type: ignore
if world_size > 1:
assert self.kwargs["model_dtype"] == "bf16", "FP8 is not supported for multi-GPU inference"
model = setup_fsdp_sync(
model,
device_id=device_id,
param_dtype=torch.float32,
auto_wrap_policy=partial(
lambda_auto_wrap_policy,
lambda_fn=lambda m: m in model.blocks,
),
)
elif isinstance(device_id, int):
model = model.to(torch.device(f"cuda:{device_id}"))
return model.eval()
class DecoderModelFactory(ModelFactory):
def __init__(self, *, model_path: str):
super().__init__(model_path=model_path)
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
# TODO(ved): Set flag for torch.compile
# TODO(ved): Use skip_init
decoder = Decoder(
out_channels=3,
base_channels=128,
channel_multipliers=[1, 2, 4, 6],
temporal_expansions=[1, 2, 3],
spatial_expansions=[2, 2, 2],
num_res_blocks=[3, 3, 4, 6, 3],
latent_dim=12,
has_attention=[False, False, False, False, False],
output_norm=False,
nonlinearity="silu",
output_nonlinearity="silu",
causal=True,
)
# VAE is not FSDP-wrapped
state_dict = load_file(self.kwargs["model_path"])
decoder.load_state_dict(state_dict, strict=True)
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
decoder.eval().to(device)
return decoder
class EncoderModelFactory(ModelFactory):
def __init__(self, *, model_path: str):
super().__init__(model_path=model_path)
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
# TODO(ved): Set flag for torch.compile
# TODO(ved): Use skip_init
# We don't FSDP the encoder b/c it is small
encoder = Encoder(
in_channels=15,
base_channels=64,
channel_multipliers=[1, 2, 4, 6],
num_res_blocks=[3, 3, 4, 6, 3],
latent_dim=12,
temporal_reductions=[1, 2, 3],
spatial_reductions=[2, 2, 2],
prune_bottlenecks=[False, False, False, False, False],
has_attentions=[False, True, True, True, True],
affine=True,
bias=True,
input_is_conv_1x1=True,
padding_mode="replicate",
)
state_dict = load_file(self.kwargs["model_path"])
encoder.load_state_dict(state_dict, strict=True)
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
encoder.eval().to(device)
return encoder
def get_conditioning(
tokenizer: T5Tokenizer,
encoder: Encoder,
device: torch.device,
batch_inputs: bool,
*,
prompt: str,
negative_prompt: str,
):
if batch_inputs:
return dict(
batched=get_conditioning_for_prompts(
tokenizer, encoder, device, [prompt, negative_prompt]
)
)
else:
cond_input = get_conditioning_for_prompts(tokenizer, encoder, device, [prompt])
null_input = get_conditioning_for_prompts(tokenizer, encoder, device, [negative_prompt])
return dict(cond=cond_input, null=null_input)
def get_conditioning_for_prompts(tokenizer, encoder, device, prompts: List[str]):
assert len(prompts) in [1, 2] # [neg] or [pos] or [pos, neg]
B = len(prompts)
t5_toks = tokenizer(
prompts,
padding="max_length",
truncation=True,
max_length=MAX_T5_TOKEN_LENGTH,
return_tensors="pt",
return_attention_mask=True,
)
caption_input_ids_t5 = t5_toks["input_ids"]
caption_attention_mask_t5 = t5_toks["attention_mask"].bool()
del t5_toks
assert caption_input_ids_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
assert caption_attention_mask_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
# Special-case empty negative prompt by zero-ing it
if prompts[-1] == "":
caption_input_ids_t5[-1] = 0
caption_attention_mask_t5[-1] = False
caption_input_ids_t5 = caption_input_ids_t5.to(device, non_blocking=True)
caption_attention_mask_t5 = caption_attention_mask_t5.to(device, non_blocking=True)
y_mask = [caption_attention_mask_t5]
y_feat = [encoder(caption_input_ids_t5, caption_attention_mask_t5).last_hidden_state.detach()]
# Sometimes returns a tensor, othertimes a tuple, not sure why
# See: https://huggingface.co/genmo/mochi-1-preview/discussions/3
assert tuple(y_feat[-1].shape) == (B, MAX_T5_TOKEN_LENGTH, 4096)
assert y_feat[-1].dtype == torch.float32
return dict(y_mask=y_mask, y_feat=y_feat)
def compute_packed_indices(
device: torch.device, text_mask: torch.Tensor, num_latents: int
) -> Dict[str, Union[torch.Tensor, int]]:
"""
Based on https://github.com/Dao-AILab/flash-attention/blob/765741c1eeb86c96ee71a3291ad6968cfbf4e4a1/flash_attn/bert_padding.py#L60-L80
Args:
num_latents: Number of latent tokens
text_mask: (B, L) List of boolean tensor indicating which text tokens are not padding.
Returns:
packed_indices: Dict with keys for Flash Attention:
- valid_token_indices_kv: up to (B * (N + L),) tensor of valid token indices (non-padding)
in the packed sequence.
- cu_seqlens_kv: (B + 1,) tensor of cumulative sequence lengths in the packed sequence.
- max_seqlen_in_batch_kv: int of the maximum sequence length in the batch.
"""
# Create an expanded token mask saying which tokens are valid across both visual and text tokens.
PATCH_SIZE = 2
num_visual_tokens = num_latents // (PATCH_SIZE**2)
assert num_visual_tokens > 0
mask = F.pad(text_mask, (num_visual_tokens, 0), value=True) # (B, N + L)
seqlens_in_batch = mask.sum(dim=-1, dtype=torch.int32) # (B,)
valid_token_indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten() # up to (B * (N + L),)
assert valid_token_indices.size(0) >= text_mask.size(0) * num_visual_tokens # At least (B * N,)
cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
max_seqlen_in_batch = seqlens_in_batch.max().item()
return {
"cu_seqlens_kv": cu_seqlens.to(device, non_blocking=True),
"max_seqlen_in_batch_kv": cast(int, max_seqlen_in_batch),
"valid_token_indices_kv": valid_token_indices.to(device, non_blocking=True),
}
def assert_eq(x, y, msg=None):
assert x == y, f"{msg or 'Assertion failed'}: {x} != {y}"
def sample_model(device, dit, conditioning, **args):
random.seed(args["seed"])
np.random.seed(args["seed"])
torch.manual_seed(args["seed"])
generator = torch.Generator(device=device)
generator.manual_seed(args["seed"])
w, h, t = args["width"], args["height"], args["num_frames"]
sample_steps = args["num_inference_steps"]
cfg_schedule = args["cfg_schedule"]
sigma_schedule = args["sigma_schedule"]
assert_eq(len(cfg_schedule), sample_steps, "cfg_schedule must have length sample_steps")
assert_eq((t - 1) % 6, 0, "t - 1 must be divisible by 6")
assert_eq(
len(sigma_schedule),
sample_steps + 1,
"sigma_schedule must have length sample_steps + 1",
)
B = 1
SPATIAL_DOWNSAMPLE = 8
TEMPORAL_DOWNSAMPLE = 6
IN_CHANNELS = 12
latent_t = ((t - 1) // TEMPORAL_DOWNSAMPLE) + 1
latent_w, latent_h = w // SPATIAL_DOWNSAMPLE, h // SPATIAL_DOWNSAMPLE
z = torch.randn(
(B, IN_CHANNELS, latent_t, latent_h, latent_w),
device=device,
dtype=torch.float32,
)
num_latents = latent_t * latent_h * latent_w
cond_batched = cond_text = cond_null = None
if "cond" in conditioning:
cond_text = conditioning["cond"]
cond_null = conditioning["null"]
cond_text["packed_indices"] = compute_packed_indices(device, cond_text["y_mask"][0], num_latents)
cond_null["packed_indices"] = compute_packed_indices(device, cond_null["y_mask"][0], num_latents)
else:
cond_batched = conditioning["batched"]
cond_batched["packed_indices"] = compute_packed_indices(device, cond_batched["y_mask"][0], num_latents)
z = repeat(z, "b ... -> (repeat b) ...", repeat=2)
def model_fn(*, z, sigma, cfg_scale):
if cond_batched:
with torch.autocast("cuda", dtype=torch.bfloat16):
out = dit(z, sigma, **cond_batched)
out_cond, out_uncond = torch.chunk(out, chunks=2, dim=0)
else:
nonlocal cond_text, cond_null
with torch.autocast("cuda", dtype=torch.bfloat16):
out_cond = dit(z, sigma, **cond_text)
out_uncond = dit(z, sigma, **cond_null)
assert out_cond.shape == out_uncond.shape
out_uncond = out_uncond.to(z)
out_cond = out_cond.to(z)
return out_uncond + cfg_scale * (out_cond - out_uncond)
# Euler sampler w/ customizable sigma schedule & cfg scale
for i in get_new_progress_bar(range(0, sample_steps), desc="Sampling"):
sigma = sigma_schedule[i]
dsigma = sigma - sigma_schedule[i + 1]
# `pred` estimates `z_0 - eps`.
pred = model_fn(
z=z,
sigma=torch.full([B] if cond_text else [B * 2], sigma, device=z.device),
cfg_scale=cfg_schedule[i],
)
assert pred.dtype == torch.float32
z = z + dsigma * pred
z = z[:B] if cond_batched else z
return dit_latents_to_vae_latents(z)
@contextmanager
def move_to_device(model: nn.Module, target_device, *, enabled=True):
if not enabled:
yield
return
og_device = next(model.parameters()).device
if og_device == target_device:
print(f"move_to_device is a no-op model is already on {target_device}")
else:
print(f"moving model from {og_device} -> {target_device}")
model.to(target_device)
yield
if og_device != target_device:
print(f"moving model from {target_device} -> {og_device}")
model.to(og_device)
def t5_tokenizer(model_dir=None):
return T5Tokenizer.from_pretrained(model_dir or T5_MODEL, legacy=False)
class MochiSingleGPUPipeline:
def __init__(
self,
*,
text_encoder_factory: ModelFactory,
dit_factory: ModelFactory,
decoder_factory: ModelFactory,
cpu_offload: Optional[bool] = False,
decode_type: str = "full",
decode_args: Optional[Dict[str, Any]] = None,
fast_init=True,
strict_load=True
):
self.device = torch.device("cuda:0")
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
t = Timer()
self.cpu_offload = cpu_offload
self.decode_args = decode_args or {}
self.decode_type = decode_type
init_id = "cpu" if cpu_offload else 0
with t("load_text_encoder"):
self.text_encoder = text_encoder_factory.get_model(
local_rank=0,
device_id=init_id,
world_size=1,
)
with t("load_dit"):
self.dit = dit_factory.get_model(local_rank=0, device_id=init_id, world_size=1, fast_init=fast_init, strict_load=strict_load) # type: ignore
with t("load_vae"):
self.decoder = decoder_factory.get_model(local_rank=0, device_id=init_id, world_size=1)
t.print_stats()
def __call__(self, batch_cfg, prompt, negative_prompt, **kwargs):
with torch.inference_mode():
print_max_memory = lambda: print(
f"Max memory reserved: {torch.cuda.max_memory_reserved() / 1024**3:.2f} GB"
)
print_max_memory()
with move_to_device(self.text_encoder, self.device):
conditioning = get_conditioning(
tokenizer=self.tokenizer,
encoder=self.text_encoder,
device=self.device,
batch_inputs=batch_cfg,
prompt=prompt,
negative_prompt=negative_prompt,
)
print_max_memory()
with move_to_device(self.dit, self.device):
latents = sample_model(self.device, self.dit, conditioning, **kwargs)
print_max_memory()
with move_to_device(self.decoder, self.device):
if self.decode_type == "tiled_full":
frames = decode_latents_tiled_full(
self.decoder, latents, **self.decode_args)
elif self.decode_type == "tiled_spatial":
frames = decode_latents_tiled_spatial(
self.decoder, latents, **self.decode_args,
num_tiles_w=4, num_tiles_h=2)
else:
frames = decode_latents(self.decoder, latents)
print_max_memory()
return frames.cpu().numpy()
def cast_dit(model, dtype):
for name, module in model.named_modules():
if isinstance(module, nn.Linear):
assert any(
n in name for n in ["mlp", "t5", "mod_", "attn.qkv_", "attn.proj_", "final_layer"]
), f"Unexpected linear layer: {name}"
module.to(dtype=dtype)
elif isinstance(module, nn.Conv2d):
assert "x_embedder.proj" in name, f"Unexpected conv2d layer: {name}"
module.to(dtype=dtype)
return model
### ALL CODE BELOW HERE IS FOR MULTI-GPU MODE ###
# In multi-gpu mode, all models must belong to a device which has a predefined context parallel group
# So it doesn't make sense to work with models individually
class MultiGPUContext:
def __init__(
self,
*,
text_encoder_factory,
dit_factory,
decoder_factory,
device_id,
local_rank,
world_size,
):
t = Timer()
self.device = torch.device(f"cuda:{device_id}")
print(f"Initializing rank {local_rank+1}/{world_size}")
assert world_size > 1, f"Multi-GPU mode requires world_size > 1, got {world_size}"
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "29500"
with t("init_process_group"):
dist.init_process_group(
"nccl",
rank=local_rank,
world_size=world_size,
device_id=self.device, # force non-lazy init
)
pg = dist.group.WORLD
cp.set_cp_group(pg, list(range(world_size)), local_rank)
distributed_kwargs = dict(local_rank=local_rank, device_id=device_id, world_size=world_size)
self.world_size = world_size
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
with t("load_text_encoder"):
self.text_encoder = text_encoder_factory.get_model(**distributed_kwargs)
with t("load_dit"):
self.dit = dit_factory.get_model(**distributed_kwargs)
with t("load_vae"):
self.decoder = decoder_factory.get_model(**distributed_kwargs)
self.local_rank = local_rank
t.print_stats()
def run(self, *, fn, **kwargs):
return fn(self, **kwargs)
class MochiMultiGPUPipeline:
def __init__(
self,
*,
text_encoder_factory: ModelFactory,
dit_factory: ModelFactory,
decoder_factory: ModelFactory,
world_size: int,
):
ray.init()
RemoteClass = ray.remote(MultiGPUContext)
self.ctxs = [
RemoteClass.options(num_gpus=1).remote(
text_encoder_factory=text_encoder_factory,
dit_factory=dit_factory,
decoder_factory=decoder_factory,
world_size=world_size,
device_id=0,
local_rank=i,
)
for i in range(world_size)
]
for ctx in self.ctxs:
ray.get(ctx.__ray_ready__.remote())
def __call__(self, **kwargs):
def sample(ctx, *, batch_cfg, prompt, negative_prompt, **kwargs):
with progress_bar(type="ray_tqdm", enabled=ctx.local_rank == 0), torch.inference_mode():
conditioning = get_conditioning(
ctx.tokenizer,
ctx.text_encoder,
ctx.device,
batch_cfg,
prompt=prompt,
negative_prompt=negative_prompt,
)
latents = sample_model(ctx.device, ctx.dit, conditioning=conditioning, **kwargs)
if ctx.local_rank == 0:
torch.save(latents, "latents.pt")
frames = decode_latents(ctx.decoder, latents)
return frames.cpu().numpy()
return ray.get([ctx.run.remote(fn=sample, **kwargs, show_progress=i == 0) for i, ctx in enumerate(self.ctxs)])[
0
]
@@ -0,0 +1,155 @@
from typing import Tuple, Union
import torch
import torch.distributed as dist
import torch.nn.functional as F
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
def cast_tuple(t, length=1):
return t if isinstance(t, tuple) else ((t,) * length)
def cp_pass_frames(x: torch.Tensor, frames_to_send: int) -> torch.Tensor:
"""
Forward pass that handles communication between ranks for inference.
Args:
x: Tensor of shape (B, C, T, H, W)
frames_to_send: int, number of frames to communicate between ranks
Returns:
output: Tensor of shape (B, C, T', H, W)
"""
cp_rank, cp_world_size = cp.get_cp_rank_size()
if frames_to_send == 0 or cp_world_size == 1:
return x
group = cp.get_cp_group()
global_rank = dist.get_rank()
# Send to next rank
if cp_rank < cp_world_size - 1:
assert x.size(2) >= frames_to_send
tail = x[:, :, -frames_to_send:].contiguous()
dist.send(tail, global_rank + 1, group=group)
# Receive from previous rank
if cp_rank > 0:
B, C, _, H, W = x.shape
recv_buffer = torch.empty(
(B, C, frames_to_send, H, W),
dtype=x.dtype,
device=x.device,
)
dist.recv(recv_buffer, global_rank - 1, group=group)
x = torch.cat([recv_buffer, x], dim=2)
return x
def _pad_to_max(x: torch.Tensor, max_T: int) -> torch.Tensor:
if max_T > x.size(2):
pad_T = max_T - x.size(2)
pad_dims = (0, 0, 0, 0, 0, pad_T)
return F.pad(x, pad_dims)
return x
def gather_all_frames(x: torch.Tensor) -> torch.Tensor:
"""
Gathers all frames from all processes for inference.
Args:
x: Tensor of shape (B, C, T, H, W)
Returns:
output: Tensor of shape (B, C, T_total, H, W)
"""
cp_rank, cp_size = cp.get_cp_rank_size()
if cp_size == 1:
return x
cp_group = cp.get_cp_group()
# Ensure the tensor is contiguous for collective operations
x = x.contiguous()
# Get the local time dimension size
local_T = x.size(2)
local_T_tensor = torch.tensor([local_T], device=x.device, dtype=torch.int64)
# Gather all T sizes from all processes
all_T = [torch.zeros(1, dtype=torch.int64, device=x.device) for _ in range(cp_size)]
dist.all_gather(all_T, local_T_tensor, group=cp_group)
all_T = [t.item() for t in all_T]
# Pad the tensor at the end of the time dimension to match max_T
max_T = max(all_T)
x = _pad_to_max(x, max_T).contiguous()
# Prepare a list to hold the gathered tensors
gathered_x = [torch.zeros_like(x).contiguous() for _ in range(cp_size)]
# Perform the all_gather operation
dist.all_gather(gathered_x, x, group=cp_group)
# Slice each gathered tensor back to its original T size
for idx, t_size in enumerate(all_T):
gathered_x[idx] = gathered_x[idx][:, :, :t_size]
return torch.cat(gathered_x, dim=2)
def excessive_memory_usage(input: torch.Tensor, max_gb: float = 2.0) -> bool:
"""Estimate memory usage based on input tensor size and data type."""
element_size = input.element_size() # Size in bytes of each element
memory_bytes = input.numel() * element_size
memory_gb = memory_bytes / 1024**3
return memory_gb > max_gb
class ContextParallelCausalConv3d(torch.nn.Conv3d):
def __init__(
self,
in_channels,
out_channels,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]],
**kwargs,
):
kernel_size = cast_tuple(kernel_size, 3)
stride = cast_tuple(stride, 3)
height_pad = (kernel_size[1] - 1) // 2
width_pad = (kernel_size[2] - 1) // 2
super().__init__(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=stride,
dilation=(1, 1, 1),
padding=(0, height_pad, width_pad),
**kwargs,
)
def forward(self, x: torch.Tensor):
cp_rank, cp_world_size = cp.get_cp_rank_size()
context_size = self.kernel_size[0] - 1
if cp_rank == 0:
mode = "constant" if self.padding_mode == "zeros" else self.padding_mode
x = F.pad(x, (0, 0, 0, 0, context_size, 0), mode=mode)
if cp_world_size == 1:
return super().forward(x)
if all(s == 1 for s in self.stride):
# Receive some frames from previous rank.
x = cp_pass_frames(x, context_size)
return super().forward(x)
# Less efficient implementation for strided convs.
# All gather x, infer and chunk.
x = gather_all_frames(x) # [B, C, k - 1 + global_T, H, W]
x = super().forward(x)
x_chunks = x.tensor_split(cp_world_size, dim=2)
assert len(x_chunks) == cp_world_size
return x_chunks[cp_rank]
@@ -0,0 +1,35 @@
"""Container for latent space posterior."""
import torch
class LatentDistribution:
def __init__(self, mean: torch.Tensor, logvar: torch.Tensor):
"""Initialize latent distribution.
Args:
mean: Mean of the distribution. Shape: [B, C, T, H, W].
logvar: Logarithm of variance of the distribution. Shape: [B, C, T, H, W].
"""
assert mean.shape == logvar.shape
self.mean = mean
self.logvar = logvar
def sample(self, temperature=1.0, generator: torch.Generator = None, noise=None):
if temperature == 0.0:
return self.mean
if noise is None:
noise = torch.randn(self.mean.shape, device=self.mean.device, dtype=self.mean.dtype, generator=generator)
else:
assert noise.device == self.mean.device
noise = noise.to(self.mean.dtype)
if temperature != 1.0:
raise NotImplementedError(f"Temperature {temperature} is not supported.")
# Just Gaussian sample with no scaling of variance.
return noise * torch.exp(self.logvar * 0.5) + self.mean
def mode(self):
return self.mean
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,67 @@
import torch
# Channel-wise mean and standard deviation of VAE encoder latents
STATS = {
"mean": torch.Tensor(
[
-0.06730895953510081,
-0.038011381506090416,
-0.07477820912866141,
-0.05565264470995561,
0.012767231469026969,
-0.04703542746246419,
0.043896967884726704,
-0.09346305707025976,
-0.09918314763016893,
-0.008729793427399178,
-0.011931556316503654,
-0.0321993391887285,
]
),
"std": torch.Tensor(
[
0.9263795028493863,
0.9248894543193766,
0.9393059390890617,
0.959253732819592,
0.8244560132752793,
0.917259975397747,
0.9294154431013696,
1.3720942357788521,
0.881393668867029,
0.9168315692124348,
0.9185249279345552,
0.9274757570805041,
]
),
}
def dit_latents_to_vae_latents(dit_outputs: torch.Tensor) -> torch.Tensor:
"""Unnormalize latents output by Mochi's DiT to be compatible with VAE.
Run this on sampled latents before calling the VAE decoder.
Args:
latents (torch.Tensor): [B, C_z, T_z, H_z, W_z], float
Returns:
torch.Tensor: [B, C_z, T_z, H_z, W_z], float
"""
mean = STATS["mean"][:, None, None, None]
std = STATS["std"][:, None, None, None]
assert dit_outputs.ndim == 5
assert dit_outputs.size(1) == mean.size(0) == std.size(0)
return dit_outputs * std.to(dit_outputs) + mean.to(dit_outputs)
def vae_latents_to_dit_latents(vae_latents: torch.Tensor):
"""Normalize latents output by the VAE encoder to be compatible with Mochi's DiT.
E.g, for fine-tuning or video-to-video.
"""
mean = STATS["mean"][:, None, None, None]
std = STATS["std"][:, None, None, None]
assert vae_latents.ndim == 5
assert vae_latents.size(1) == mean.size(0) == std.size(0)
return (vae_latents - mean.to(vae_latents)) / std.to(vae_latents)
+24 -28
View File
@@ -55,7 +55,6 @@ from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -185,7 +184,6 @@ class MochiAttention(nn.Module):
**kwargs,
)
class MochiAttnProcessor2_0:
"""Attention processor used in Mochi."""
@@ -279,6 +277,7 @@ class MochiAttnProcessor2_0:
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# H
query = torch.cat([query, encoder_query], dim=1).unsqueeze(2)
key = torch.cat([key, encoder_key], dim=1).unsqueeze(2)
@@ -288,9 +287,8 @@ class MochiAttnProcessor2_0:
attn_mask = encoder_attention_mask[:, :].bool()
attn_mask = F.pad(attn_mask, (sequence_length, 0), value=True)
hidden_states = flash_attn_no_pad(
qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None
)
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
# hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask = None, dropout_p=0.0, is_causal=False)
@@ -651,11 +649,11 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
output_attn=False,
hidden_states: torch.Tensor, # [2, 12, 28, 60, 106]
encoder_hidden_states: torch.Tensor, # [2, 256, 4096]
timestep: torch.LongTensor, # [2]
encoder_attention_mask: torch.Tensor, #[2, 256]
output_attn = False,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
) -> torch.Tensor:
@@ -680,28 +678,28 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p = self.config.patch_size
post_patch_height = height // p
post_patch_width = width // p
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
timestep = 1000 - timestep
temb, encoder_hidden_states = self.time_embed(
temb, encoder_hidden_states = self.time_embed( # [2, 3072], [2, 256, 1536]
timestep,
encoder_hidden_states,
encoder_attention_mask,
hidden_dtype=hidden_states.dtype,
)
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
hidden_states = self.patch_embed(hidden_states)
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [56, 12, 60, 106]
hidden_states = self.patch_embed(hidden_states) # [56, 1590, 3072]
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2) # [2, 44520, 3072]
image_rotary_emb = self.rope(
self.pos_frequencies,
num_frames,
image_rotary_emb = self.rope( #[0][44520, 24, 64]
self.pos_frequencies, #[3, 24, 64]
num_frames, # 28
post_patch_height,
post_patch_width,
device=hidden_states.device,
@@ -735,7 +733,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states, attn_outputs = block(
hidden_states, encoder_hidden_states, attn_outputs = block( # [2, 44520, 3072], [2, 256, 1536],
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
@@ -744,16 +742,14 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
output_attn=output_attn,
)
attn_outputs_list.append(attn_outputs)
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(
batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1
)
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5)
output = hidden_states.reshape(batch_size, -1, num_frames, height, width)
hidden_states = self.proj_out(hidden_states) #[2, 44520, 48]
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1) # [2, 28, 30, 53, 2, 2, 12]
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5) # [2, 12, 28, 30, 2, 53, 2]
output = hidden_states.reshape(batch_size, -1, num_frames, height, width) # [2, 12, 28, 60, 106]
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
+18 -6
View File
@@ -22,7 +22,8 @@ from typing import Dict
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import convert_unet_state_dict_to_peft
from fastvideo.distill.solver import PCMFMScheduler
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
from safetensors.torch import load_file
def initialize_distributed():
local_rank = int(os.getenv("RANK", 0))
@@ -53,16 +54,27 @@ def main(args):
args.linear_threshold,
args.linear_range,
)
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
mochi_genmo = False
if mochi_genmo:
model_path = "/root/weights/dit.safetensors"
state_dcit = load_file(model_path)
transformer = AsymmDiTJoint()
transformer.load_state_dict(state_dcit)
# from IPython import embed
# embed()
transformer.config.in_channels = 12
print("load gennmo mochi successfully")
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/"
)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder='transformer/')
pipe = MochiPipeline.from_pretrained(
args.model_path, transformer=transformer, scheduler=scheduler
)
pipe.enable_vae_tiling()
+8 -1
View File
@@ -214,9 +214,10 @@ def main(args):
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
f
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
# keep the master weight to float32
weight_type = torch.float32 if args.master_weight_type == 'fp32' else torch.bfloat16
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="transformer",
@@ -728,6 +729,12 @@ if __name__ == "__main__":
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
parser.add_argument(
"--Mochi_type",
type=str,
default="hf",
help="Choose Mochi model between hf and genmo(original mochi).",
)
args = parser.parse_args()
main(args)
+1 -1
View File
@@ -15,7 +15,7 @@ if __name__ == "__main__":
type=str,
help="The local directory to download the repository to",
)
parser.add_argument(
parser.add_argument(di
"--repo_type",
type=str,
help="The type of repository to download (dataset or model)",