Ref2VA/video_vae/base_module.py
| 1 | # SPDX-License-Identifier: Apache-2.0 |
| 2 | # Transformer building blocks for the MiniMax H3 visual VAE ViT decoder. |
| 3 | import math |
| 4 | import os |
| 5 | import torch |
| 6 | import torch.nn as nn |
| 7 | from typing import Optional |
| 8 | from diffusers.utils import logging |
| 9 | from diffusers.utils.torch_utils import maybe_allow_in_graph |
| 10 | |
| 11 | from .attention import Attention |
| 12 | |
| 13 | logger = logging.get_logger(__name__) # pylint: disable=invalid-name |
| 14 | |
| 15 | |
| 16 | def _env_flag(name, default="0"): |
| 17 | value = os.environ.get(name, default) |
| 18 | return str(value).strip().lower() in ("1", "true", "yes", "on") |
| 19 | |
| 20 | |
| 21 | def _env_optional_bool(name, default=""): |
| 22 | value = str(os.environ.get(name, default)).strip().lower() |
| 23 | if value in ("", "default", "auto", "none", "unset"): |
| 24 | return None |
| 25 | return value not in ("0", "false", "no", "off", "disabled") |
| 26 | |
| 27 | |
| 28 | def _vit_torch_compile_kwargs(prefix): |
| 29 | kwargs = {} |
| 30 | backend = os.environ.get(f"{prefix}_BACKEND", "inductor").strip() |
| 31 | mode = os.environ.get(f"{prefix}_MODE", "reduce-overhead").strip() |
| 32 | if backend and backend.lower() not in ("default", "none"): |
| 33 | kwargs["backend"] = backend |
| 34 | if mode and mode.lower() not in ("default", "none"): |
| 35 | kwargs["mode"] = mode |
| 36 | kwargs["fullgraph"] = _env_flag(f"{prefix}_FULLGRAPH", "0") |
| 37 | dynamic = _env_optional_bool(f"{prefix}_DYNAMIC") |
| 38 | if dynamic is not None: |
| 39 | kwargs["dynamic"] = dynamic |
| 40 | return kwargs |
| 41 | |
| 42 | |
| 43 | |
| 44 | |
| 45 | def _vit_norm_input(module, hidden_states): |
| 46 | if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1"): |
| 47 | return hidden_states.float() |
| 48 | return hidden_states.to(getattr(module.weight, "dtype", hidden_states.dtype)) |
| 49 | |
| 50 | |
| 51 | |
| 52 | |
| 53 | |
| 54 | |
| 55 | class FeedForward(nn.Module): |
| 56 | def __init__( |
| 57 | self, |
| 58 | dim: int, |
| 59 | dim_out: Optional[int] = None, |
| 60 | mult: int = 4, |
| 61 | activation_fn: str = "silu", |
| 62 | bias: bool = True, |
| 63 | use_gated: bool = True, |
| 64 | glu_balanced: bool = False, |
| 65 | ): |
| 66 | super().__init__() |
| 67 | ratio = 2 / 3 if (use_gated and glu_balanced) else 1 |
| 68 | inner_dim = round(dim * mult * ratio) |
| 69 | dim_out = dim_out if dim_out is not None else dim |
| 70 | self.use_gated = use_gated |
| 71 | |
| 72 | if use_gated: |
| 73 | self.w1 = nn.Linear(dim, inner_dim * 2, bias=bias) |
| 74 | else: |
| 75 | self.w1 = nn.Linear(dim, inner_dim, bias=bias) |
| 76 | |
| 77 | if activation_fn == "silu": |
| 78 | self.act_fn = nn.SiLU() |
| 79 | elif activation_fn == "gelu": |
| 80 | self.act_fn = nn.GELU() |
| 81 | elif activation_fn == "gelu-approximate": |
| 82 | self.act_fn = nn.GELU(approximate="tanh") |
| 83 | else: |
| 84 | raise ValueError(f"Unsupported activation function: {activation_fn}") |
| 85 | |
| 86 | self.w2 = nn.Linear(inner_dim, dim_out, bias=bias) |
| 87 | self._compile_forward_enabled = _env_flag( |
| 88 | "MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE", "0" |
| 89 | ) |
| 90 | self._compile_forward_fatal = _env_flag( |
| 91 | "MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE_FATAL", "0" |
| 92 | ) |
| 93 | self._compiled_forward = None |
| 94 | |
| 95 | def _forward_impl(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| 96 | hidden_states = self.w1(hidden_states) |
| 97 | |
| 98 | if self.use_gated: |
| 99 | gate, hidden_states = hidden_states.chunk(2, dim=-1) |
| 100 | hidden_states = self.act_fn(gate) * hidden_states |
| 101 | else: |
| 102 | hidden_states = self.act_fn(hidden_states) |
| 103 | |
| 104 | hidden_states = self.w2(hidden_states) |
| 105 | return hidden_states |
| 106 | |
| 107 | def _get_forward_impl(self): |
| 108 | if not self._compile_forward_enabled: |
| 109 | return self._forward_impl |
| 110 | if self._compiled_forward is not None: |
| 111 | return self._compiled_forward |
| 112 | if not hasattr(torch, "compile"): |
| 113 | message = "torch.compile is unavailable; falling back to eager ViT FeedForward" |
| 114 | if self._compile_forward_fatal: |
| 115 | raise RuntimeError(message) |
| 116 | logger.warning(f"[ViTFeedForward] {message}") |
| 117 | self._compile_forward_enabled = False |
| 118 | return self._forward_impl |
| 119 | |
| 120 | kwargs = _vit_torch_compile_kwargs("MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE") |
| 121 | try: |
| 122 | self._compiled_forward = torch.compile(self._forward_impl, **kwargs) |
| 123 | logger.info(f"[ViTFeedForward] torch.compile enabled kwargs={kwargs}") |
| 124 | except Exception as exc: |
| 125 | if self._compile_forward_fatal: |
| 126 | raise |
| 127 | logger.warning( |
| 128 | f"[ViTFeedForward] torch.compile setup failed: {type(exc).__name__}: {exc}; " |
| 129 | "falling back to eager" |
| 130 | ) |
| 131 | self._compile_forward_enabled = False |
| 132 | self._compiled_forward = None |
| 133 | return self._forward_impl |
| 134 | return self._compiled_forward |
| 135 | |
| 136 | def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| 137 | forward_impl = self._get_forward_impl() |
| 138 | try: |
| 139 | return forward_impl(hidden_states) |
| 140 | except Exception as exc: |
| 141 | if ( |
| 142 | self._compile_forward_enabled |
| 143 | and self._compiled_forward is not None |
| 144 | and forward_impl is self._compiled_forward |
| 145 | and not self._compile_forward_fatal |
| 146 | ): |
| 147 | logger.warning( |
| 148 | f"[ViTFeedForward] compiled forward failed: {type(exc).__name__}: {exc}; " |
| 149 | "disabling compile and retrying eager" |
| 150 | ) |
| 151 | self._compile_forward_enabled = False |
| 152 | self._compiled_forward = None |
| 153 | return self._forward_impl(hidden_states) |
| 154 | raise |
| 155 | |
| 156 | |
| 157 | class RotaryEmbeddingND(nn.Module): |
| 158 | def __init__(self, dim, rotary_base=10000, n_dim=3, use_angle=False): |
| 159 | super().__init__() |
| 160 | self.dim = dim |
| 161 | self.n_dim = n_dim |
| 162 | |
| 163 | if dim % (2 * n_dim) != 0: |
| 164 | raise ValueError( |
| 165 | f"head_dim {dim} must be divisible by 2 * n_dim {2 * n_dim}" |
| 166 | ) |
| 167 | |
| 168 | if use_angle: |
| 169 | self.angle_scale = 2.0 * math.pi |
| 170 | else: |
| 171 | self.angle_scale = 1.0 |
| 172 | |
| 173 | inv_freq = 1 / rotary_base ** torch.arange( |
| 174 | 0, 1, 2 * n_dim / dim, dtype=torch.float32 |
| 175 | ) |
| 176 | self.register_buffer("inv_freq", inv_freq, persistent=False) |
| 177 | |
| 178 | def forward(self, img_ids): |
| 179 | B, N, D = img_ids.shape |
| 180 | if D != self.n_dim: |
| 181 | raise ValueError(f"Expected {self.n_dim} dimensions, got {D}") |
| 182 | |
| 183 | with torch.autocast("cuda", enabled=False): |
| 184 | angles = ( |
| 185 | self.angle_scale |
| 186 | * img_ids[:, :, :, None] |
| 187 | * self.inv_freq.to(img_ids.device)[None, None, None, :] |
| 188 | ) |
| 189 | angles = angles.flatten(2, 3) |
| 190 | angles = angles.tile(2) |
| 191 | angles = angles.unsqueeze(2) |
| 192 | |
| 193 | cos = torch.cos(angles) |
| 194 | sin = torch.sin(angles) |
| 195 | |
| 196 | return cos.to(dtype=img_ids.dtype), sin.to(dtype=img_ids.dtype) |
| 197 | |
| 198 | |
| 199 | @maybe_allow_in_graph |
| 200 | class TransformerBlock(nn.Module): |
| 201 | def __init__( |
| 202 | self, |
| 203 | heads: int, |
| 204 | dim_head: int, |
| 205 | embed_dim: Optional[int] = None, |
| 206 | ffn_glu_balanced: bool = False, |
| 207 | norm_type: str = "layer_norm", |
| 208 | norm_affine: bool = True, |
| 209 | qk_norm_type: str = "rms_norm", |
| 210 | qk_norm_affine: bool = False, |
| 211 | ffn_activation_fn: str = "silu", |
| 212 | ffn_use_gated: bool = True, |
| 213 | use_scale: bool = True, |
| 214 | bias: bool = True, |
| 215 | eps: float = 1e-5, |
| 216 | **kwargs, |
| 217 | ): |
| 218 | super().__init__() |
| 219 | dim = embed_dim if embed_dim is not None else dim_head * heads |
| 220 | self.use_scale = use_scale |
| 221 | |
| 222 | if norm_type == "layer_norm": |
| 223 | norm_class = nn.LayerNorm |
| 224 | elif norm_type == "rms_norm": |
| 225 | norm_class = nn.RMSNorm |
| 226 | else: |
| 227 | raise ValueError(f"unknown norm_type {norm_type}") |
| 228 | |
| 229 | self.norm1 = norm_class( |
| 230 | dim, |
| 231 | elementwise_affine=norm_affine, |
| 232 | eps=eps, |
| 233 | ) |
| 234 | self.attn = Attention( |
| 235 | heads=heads, |
| 236 | dim_head=dim_head, |
| 237 | embed_dim=dim, |
| 238 | qk_norm_type=qk_norm_type, |
| 239 | qk_norm_affine=qk_norm_affine, |
| 240 | bias=bias, |
| 241 | eps=eps, |
| 242 | **kwargs, |
| 243 | ) |
| 244 | if use_scale: |
| 245 | self.scale1 = nn.Parameter(torch.zeros(dim)) |
| 246 | |
| 247 | self.norm2 = norm_class( |
| 248 | dim, |
| 249 | elementwise_affine=norm_affine, |
| 250 | eps=eps, |
| 251 | ) |
| 252 | self.ff = FeedForward( |
| 253 | dim=dim, |
| 254 | activation_fn=ffn_activation_fn, |
| 255 | bias=bias, |
| 256 | use_gated=ffn_use_gated, |
| 257 | glu_balanced=ffn_glu_balanced, |
| 258 | ) |
| 259 | if use_scale: |
| 260 | self.scale2 = nn.Parameter(torch.zeros(dim)) |
| 261 | |
| 262 | def forward( |
| 263 | self, |
| 264 | hidden_states: torch.FloatTensor, |
| 265 | rotary_pos_emb: Optional[torch.FloatTensor] = None, |
| 266 | pack_info: dict = {}, |
| 267 | ): |
| 268 | norm_hidden_states = self.norm1(_vit_norm_input(self.norm1, hidden_states)).to(hidden_states.dtype) |
| 269 | attn_output = self.attn(norm_hidden_states, rotary_pos_emb, pack_info) |
| 270 | if self.use_scale: |
| 271 | hidden_states = hidden_states + attn_output * self.scale1 |
| 272 | else: |
| 273 | hidden_states = hidden_states + attn_output |
| 274 | |
| 275 | norm_hidden_states = self.norm2(_vit_norm_input(self.norm2, hidden_states)).to(hidden_states.dtype) |
| 276 | ff_output = self.ff(norm_hidden_states) |
| 277 | if self.use_scale: |
| 278 | hidden_states = hidden_states + ff_output * self.scale2 |
| 279 | else: |
| 280 | hidden_states = hidden_states + ff_output |
| 281 | |
| 282 | return hidden_states |
| 283 | |