Spaces:
Running on Zero
Running on Zero
File size: 5,966 Bytes
3c58630 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | from __future__ import annotations
import math
import torch
import torch.nn as nn
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
Matches the JiT-style sinusoidal embedding + MLP.
"""
def __init__(self, hidden_size: int, frequency_embedding_size: int = 256) -> None:
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=True),
)
@staticmethod
def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10_000) -> torch.Tensor:
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half
)
args = t[:, None].float() * freqs[None]
emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
return emb
def forward(self, t: torch.Tensor) -> torch.Tensor:
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
return self.mlp(t_freq)
class LabelEmbedder(nn.Module):
"""
Embeds class labels. Reserve index `num_classes` for the dropped/unconditional label.
"""
def __init__(self, num_classes: int, hidden_size: int) -> None:
super().__init__()
self.num_classes = int(num_classes)
self.embedding_table = nn.Embedding(self.num_classes + 1, hidden_size)
def forward(self, labels: torch.Tensor) -> torch.Tensor:
return self.embedding_table(labels)
class RotaryEmbedding1D(nn.Module):
"""
Temporal RoPE applied to Q/K using frame indices only.
- x is expected to be (B, H, S, D) with D even.
"""
def __init__(self, dim: int, *, max_frames: int, theta: float = 10_000.0) -> None:
super().__init__()
if dim % 2 != 0:
raise ValueError(f"RoPE dim must be even, got dim={dim}")
self.dim = int(dim)
self.max_frames = int(max_frames)
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False) # (D/2,)
self._cos = None
self._sin = None
self._cache_key = None
def set_token_layout(
self,
*,
num_frames: int,
num_tokens_per_frame: int,
num_register_tokens: int = 0,
device,
dtype: torch.dtype,
) -> None:
num_frames = int(num_frames)
num_tokens_per_frame = int(num_tokens_per_frame)
num_register_tokens = int(num_register_tokens)
if num_frames > self.max_frames:
raise ValueError(f"num_frames={num_frames} exceeds max_frames={self.max_frames}")
if num_frames <= 0 or num_tokens_per_frame <= 0 or num_register_tokens < 0:
raise ValueError(
f"Invalid layout: num_frames={num_frames}, num_tokens_per_frame={num_tokens_per_frame}, "
f"num_register_tokens={num_register_tokens}"
)
cache_key = (num_frames, num_tokens_per_frame, num_register_tokens, str(device), str(dtype))
if self._cache_key == cache_key and self._cos is not None and self._sin is not None:
return
# Compute RoPE for frame indices only: every token in a frame shares that frame index.
half = self.dim // 2
frame_pos = torch.arange(num_frames, device=device, dtype=torch.float32) # (F,)
inv_freq = self.inv_freq.to(device=device, dtype=torch.float32) # (D/2,)
freqs = torch.einsum("i,j->ij", frame_pos, inv_freq) # (F, D/2)
freqs = freqs[:, :half]
emb = torch.cat([freqs, freqs], dim=-1) # (F, D)
cos_f = emb.cos()
sin_f = emb.sin()
seq_len = num_frames * num_tokens_per_frame
token_frame_for_sequence = (torch.arange(seq_len, device=device) // num_tokens_per_frame).long()
if num_register_tokens > 0:
token_frame = torch.cat(
[
torch.zeros((num_register_tokens,), device=device, dtype=torch.long),
token_frame_for_sequence,
],
dim=0,
)
else:
token_frame = token_frame_for_sequence
cos = cos_f.index_select(0, token_frame).to(dtype=dtype) # (S_total, D)
sin = sin_f.index_select(0, token_frame).to(dtype=dtype) # (S_total, D)
self._cos = cos.unsqueeze(0).unsqueeze(0) # (1,1,S,D)
self._sin = sin.unsqueeze(0).unsqueeze(0) # (1,1,S,D)
self._cache_key = cache_key
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.ndim != 4:
raise ValueError(f"Expected (B,H,S,D), got {tuple(x.shape)}")
if x.shape[-1] != self.dim:
raise ValueError(f"Expected dim={self.dim}, got {x.shape[-1]}")
if self._cos is None or self._sin is None:
raise RuntimeError("RotaryEmbedding1D layout is not set. Call set_token_layout(...) before forward().")
seq_len = x.shape[2]
cos = self._cos.to(device=x.device, dtype=x.dtype)
sin = self._sin.to(device=x.device, dtype=x.dtype)
if cos.shape[2] != seq_len:
raise ValueError(f"Cached RoPE seq_len={cos.shape[2]} does not match x seq_len={seq_len}")
# Broadcast (1,1,S,D) -> (B,H,S,D).
cos = cos.expand(x.shape[0], -1, -1, -1)
sin = sin.expand(x.shape[0], -1, -1, -1)
return (x * cos) + (_rotate_half(x) * sin)
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
# HuggingFace/LLaMA-style RoPE: treat the last dim as (d/2 real | d/2 imag).
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
|