Spaces:
Running on Zero
Running on Zero
File size: 9,865 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 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 | from __future__ import annotations
import os
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
_SDPA_DEBUG_PRINTED = False
def _configure_torch_sdpa() -> None:
cuda_backend = getattr(torch.backends, "cuda", None)
if cuda_backend is None:
return
settings = {
"enable_flash_sdp": True,
"enable_mem_efficient_sdp": True,
"enable_math_sdp": False,
}
for name, value in settings.items():
fn = getattr(cuda_backend, name, None)
if callable(fn):
fn(value)
_configure_torch_sdpa()
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6) -> None:
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
# Keep output dtype stable under AMP: bf16/fp16 * fp32 promotes to fp32,
# which can make PyTorch fast CUDA SDPA unavailable.
return x.to(dtype) * self.weight.to(dtype)
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
if shift.dtype != x.dtype:
shift = shift.to(dtype=x.dtype)
if scale.dtype != x.dtype:
scale = scale.to(dtype=x.dtype)
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
class SwiGLUFFN(nn.Module):
def __init__(self, dim: int, hidden_dim: int, drop: float = 0.0, bias: bool = True) -> None:
super().__init__()
hidden_dim = int(hidden_dim * 2 / 3)
self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=bias)
self.w3 = nn.Linear(hidden_dim, dim, bias=bias)
self.drop = nn.Dropout(drop)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x12 = self.w12(x)
x1, x2 = x12.chunk(2, dim=-1)
hidden = F.silu(x1) * x2
return self.w3(self.drop(hidden))
class Attention(nn.Module):
def __init__(
self,
dim: int,
num_heads: int,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
qkv_bias: bool = True,
qk_norm: bool = True,
) -> None:
super().__init__()
if dim % num_heads != 0:
raise ValueError(f"dim={dim} must be divisible by num_heads={num_heads}")
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.q_norm = RMSNorm(self.head_dim) if qk_norm else nn.Identity()
self.k_norm = RMSNorm(self.head_dim) if qk_norm else nn.Identity()
self.attn_drop = float(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
def _sdpa(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
*,
dropout_p: float,
is_causal: bool,
attn_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
def _debug_enabled() -> bool:
v = os.environ.get("JMT4D_SDPA_DEBUG", "0").strip().lower()
return v not in ("", "0", "false", "no", "off")
def _debug_print(msg: str) -> None:
global _SDPA_DEBUG_PRINTED
if _SDPA_DEBUG_PRINTED:
return
_SDPA_DEBUG_PRINTED = True
rank = os.environ.get("RANK", "0")
print(f"[attention][rank={rank}] {msg}", flush=True)
def _raise_fast_sdpa_unavailable(cause: Optional[BaseException] = None) -> None:
detail = (
"CUDA fast attention unavailable: this demo requires PyTorch flash SDPA or "
"memory-efficient CUDA SDPA and exits instead of using math/chunked fallback. "
f"attention_device={q.device} attention_dtype={q.dtype} attention_shape={tuple(q.shape)}. "
"Check that CUDA is available, run with --device cuda --amp bf16 or --amp fp16, "
"and use a CUDA PyTorch wheel on a supported NVIDIA GPU."
)
if not q.is_cuda:
detail += f" The attention tensor is on {q.device}, not CUDA."
elif q.dtype not in (torch.float16, torch.bfloat16):
detail += " Fast CUDA SDPA requires float16 or bfloat16 attention tensors."
elif not hasattr(torch.backends.cuda, "sdp_kernel"):
detail += " This PyTorch build does not expose torch.backends.cuda.sdp_kernel."
if cause is None:
raise RuntimeError(detail)
raise RuntimeError(detail) from cause
def _try_sdpa_with_sdp_kernel(*, flash: bool, mem_efficient: bool, math_backend: bool) -> torch.Tensor:
with torch.backends.cuda.sdp_kernel(
enable_flash=bool(flash),
enable_mem_efficient=bool(mem_efficient),
enable_math=bool(math_backend),
):
return F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal
)
if not q.is_cuda or q.dtype not in (torch.float16, torch.bfloat16) or not hasattr(torch.backends.cuda, "sdp_kernel"):
_raise_fast_sdpa_unavailable()
try:
if _debug_enabled():
try:
out = _try_sdpa_with_sdp_kernel(flash=True, mem_efficient=False, math_backend=False)
_debug_print(
"CUDA fast attention active: using PyTorch flash SDPA "
f"(external flash-attn package is not required). dtype={q.dtype} q={tuple(q.shape)}"
)
return out
except Exception:
out = _try_sdpa_with_sdp_kernel(flash=False, mem_efficient=True, math_backend=False)
_debug_print(
"CUDA fast attention active: using PyTorch memory-efficient SDPA "
f"(flash SDPA was unavailable). dtype={q.dtype} q={tuple(q.shape)}"
)
return out
return _try_sdpa_with_sdp_kernel(flash=True, mem_efficient=True, math_backend=False)
except Exception as exc:
_raise_fast_sdpa_unavailable(exc)
def forward(self, x: torch.Tensor, rope, src_key_padding_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
bsz, seq_len, dim = x.shape
qkv = self.qkv(x).reshape(bsz, seq_len, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
q, k, v = qkv[0], qkv[1], qkv[2] # (B, H, S, Hd)
q = self.q_norm(q)
k = self.k_norm(k)
if rope is not None:
q = rope(q)
k = rope(k)
attn_mask = None
if src_key_padding_mask is not None:
if src_key_padding_mask.ndim != 2 or src_key_padding_mask.shape != (bsz, seq_len):
raise ValueError(f"src_key_padding_mask must be (B,S)={bsz, seq_len}, got {tuple(src_key_padding_mask.shape)}")
# RenderFormer-style semantics: True means "valid / keep".
keep = src_key_padding_mask.to(device=q.device, dtype=torch.bool).view(bsz, 1, 1, seq_len)
attn_mask = keep.expand(bsz, self.num_heads, 1, seq_len) # (B,H,1,S)
# scaled_dot_product_attention expects (B, H, S, Hd)
dropout_p = self.attn_drop if self.training else 0.0
x = self._sdpa(q, k, v, dropout_p=float(dropout_p), is_causal=False, attn_mask=attn_mask)
x = x.transpose(1, 2).reshape(bsz, seq_len, dim)
x = self.proj(x)
return self.proj_drop(x)
class DiTBlock(nn.Module):
"""
DiT-style transformer block with AdaLN (shift/scale/gates).
"""
def __init__(
self,
hidden_size: int,
num_heads: int,
mlp_ratio: float = 4.0,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
) -> None:
super().__init__()
self.norm1 = RMSNorm(hidden_size, eps=1e-6)
self.attn = Attention(
hidden_size, num_heads=num_heads, attn_drop=attn_drop, proj_drop=proj_drop, qkv_bias=True, qk_norm=True
)
self.norm2 = RMSNorm(hidden_size, eps=1e-6)
mlp_hidden = int(hidden_size * mlp_ratio)
self.mlp = SwiGLUFFN(hidden_size, mlp_hidden, drop=proj_drop)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 6 * hidden_size, bias=True),
)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor,
rope=None,
src_key_padding_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1)
x = x + gate_msa.unsqueeze(1) * self.attn(
modulate(self.norm1(x), shift_msa, scale_msa),
rope=rope,
src_key_padding_mask=src_key_padding_mask,
)
x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
return x
class FinalLayer(nn.Module):
def __init__(self, hidden_size: int, out_dim: int) -> None:
super().__init__()
self.norm_final = RMSNorm(hidden_size)
self.linear = nn.Linear(hidden_size, out_dim, bias=True)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True),
)
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
return self.linear(x)
|