Compare commits
24
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d7b07adc08 | ||
|
|
8a41e61cec | ||
|
|
60ce6a62a6 | ||
|
|
7b2ca9abec | ||
|
|
b537e01a88 | ||
|
|
a233b58a6c | ||
|
|
d24b25a3e1 | ||
|
|
2a7c147c4e | ||
|
|
0c1c939d59 | ||
|
|
881e1f130a | ||
|
|
00e899cd90 | ||
|
|
9f4151526b | ||
|
|
1ce3983d68 | ||
|
|
ee241cfa4d | ||
|
|
8106ee3f3d | ||
|
|
4dde52be9c | ||
|
|
58abda5c09 | ||
|
|
3e6019415a | ||
|
|
ccaf43c195 | ||
|
|
351538db29 | ||
|
|
361b24612d | ||
|
|
6ad03bea79 | ||
|
|
86e1f88877 | ||
|
|
e212a9c6b9 |
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
+20
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)",
|
||||
|
||||
Reference in New Issue
Block a user