| |
| |
| import math |
| import os |
| import torch |
| import torch.nn as nn |
| from typing import Optional |
| from diffusers.utils import logging |
| from diffusers.utils.torch_utils import maybe_allow_in_graph |
|
|
| from .attention import Attention |
|
|
| logger = logging.get_logger(__name__) |
|
|
|
|
| def _env_flag(name, default="0"): |
| value = os.environ.get(name, default) |
| return str(value).strip().lower() in ("1", "true", "yes", "on") |
|
|
|
|
| def _env_optional_bool(name, default=""): |
| value = str(os.environ.get(name, default)).strip().lower() |
| if value in ("", "default", "auto", "none", "unset"): |
| return None |
| return value not in ("0", "false", "no", "off", "disabled") |
|
|
|
|
| def _vit_torch_compile_kwargs(prefix): |
| kwargs = {} |
| backend = os.environ.get(f"{prefix}_BACKEND", "inductor").strip() |
| mode = os.environ.get(f"{prefix}_MODE", "reduce-overhead").strip() |
| if backend and backend.lower() not in ("default", "none"): |
| kwargs["backend"] = backend |
| if mode and mode.lower() not in ("default", "none"): |
| kwargs["mode"] = mode |
| kwargs["fullgraph"] = _env_flag(f"{prefix}_FULLGRAPH", "0") |
| dynamic = _env_optional_bool(f"{prefix}_DYNAMIC") |
| if dynamic is not None: |
| kwargs["dynamic"] = dynamic |
| return kwargs |
|
|
|
|
|
|
|
|
| def _vit_norm_input(module, hidden_states): |
| if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1"): |
| return hidden_states.float() |
| return hidden_states.to(getattr(module.weight, "dtype", hidden_states.dtype)) |
|
|
|
|
|
|
|
|
|
|
|
|
| class FeedForward(nn.Module): |
| def __init__( |
| self, |
| dim: int, |
| dim_out: Optional[int] = None, |
| mult: int = 4, |
| activation_fn: str = "silu", |
| bias: bool = True, |
| use_gated: bool = True, |
| glu_balanced: bool = False, |
| ): |
| super().__init__() |
| ratio = 2 / 3 if (use_gated and glu_balanced) else 1 |
| inner_dim = round(dim * mult * ratio) |
| dim_out = dim_out if dim_out is not None else dim |
| self.use_gated = use_gated |
|
|
| if use_gated: |
| self.w1 = nn.Linear(dim, inner_dim * 2, bias=bias) |
| else: |
| self.w1 = nn.Linear(dim, inner_dim, bias=bias) |
|
|
| if activation_fn == "silu": |
| self.act_fn = nn.SiLU() |
| elif activation_fn == "gelu": |
| self.act_fn = nn.GELU() |
| elif activation_fn == "gelu-approximate": |
| self.act_fn = nn.GELU(approximate="tanh") |
| else: |
| raise ValueError(f"Unsupported activation function: {activation_fn}") |
|
|
| self.w2 = nn.Linear(inner_dim, dim_out, bias=bias) |
| self._compile_forward_enabled = _env_flag( |
| "MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE", "0" |
| ) |
| self._compile_forward_fatal = _env_flag( |
| "MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE_FATAL", "0" |
| ) |
| self._compiled_forward = None |
|
|
| def _forward_impl(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| hidden_states = self.w1(hidden_states) |
|
|
| if self.use_gated: |
| gate, hidden_states = hidden_states.chunk(2, dim=-1) |
| hidden_states = self.act_fn(gate) * hidden_states |
| else: |
| hidden_states = self.act_fn(hidden_states) |
|
|
| hidden_states = self.w2(hidden_states) |
| return hidden_states |
|
|
| def _get_forward_impl(self): |
| if not self._compile_forward_enabled: |
| return self._forward_impl |
| if self._compiled_forward is not None: |
| return self._compiled_forward |
| if not hasattr(torch, "compile"): |
| message = "torch.compile is unavailable; falling back to eager ViT FeedForward" |
| if self._compile_forward_fatal: |
| raise RuntimeError(message) |
| logger.warning(f"[ViTFeedForward] {message}") |
| self._compile_forward_enabled = False |
| return self._forward_impl |
|
|
| kwargs = _vit_torch_compile_kwargs("MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE") |
| try: |
| self._compiled_forward = torch.compile(self._forward_impl, **kwargs) |
| logger.info(f"[ViTFeedForward] torch.compile enabled kwargs={kwargs}") |
| except Exception as exc: |
| if self._compile_forward_fatal: |
| raise |
| logger.warning( |
| f"[ViTFeedForward] torch.compile setup failed: {type(exc).__name__}: {exc}; " |
| "falling back to eager" |
| ) |
| self._compile_forward_enabled = False |
| self._compiled_forward = None |
| return self._forward_impl |
| return self._compiled_forward |
|
|
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| forward_impl = self._get_forward_impl() |
| try: |
| return forward_impl(hidden_states) |
| except Exception as exc: |
| if ( |
| self._compile_forward_enabled |
| and self._compiled_forward is not None |
| and forward_impl is self._compiled_forward |
| and not self._compile_forward_fatal |
| ): |
| logger.warning( |
| f"[ViTFeedForward] compiled forward failed: {type(exc).__name__}: {exc}; " |
| "disabling compile and retrying eager" |
| ) |
| self._compile_forward_enabled = False |
| self._compiled_forward = None |
| return self._forward_impl(hidden_states) |
| raise |
|
|
|
|
| class RotaryEmbeddingND(nn.Module): |
| def __init__(self, dim, rotary_base=10000, n_dim=3, use_angle=False): |
| super().__init__() |
| self.dim = dim |
| self.n_dim = n_dim |
|
|
| if dim % (2 * n_dim) != 0: |
| raise ValueError( |
| f"head_dim {dim} must be divisible by 2 * n_dim {2 * n_dim}" |
| ) |
|
|
| if use_angle: |
| self.angle_scale = 2.0 * math.pi |
| else: |
| self.angle_scale = 1.0 |
|
|
| inv_freq = 1 / rotary_base ** torch.arange( |
| 0, 1, 2 * n_dim / dim, dtype=torch.float32 |
| ) |
| self.register_buffer("inv_freq", inv_freq, persistent=False) |
|
|
| def forward(self, img_ids): |
| B, N, D = img_ids.shape |
| if D != self.n_dim: |
| raise ValueError(f"Expected {self.n_dim} dimensions, got {D}") |
|
|
| with torch.autocast("cuda", enabled=False): |
| angles = ( |
| self.angle_scale |
| * img_ids[:, :, :, None] |
| * self.inv_freq.to(img_ids.device)[None, None, None, :] |
| ) |
| angles = angles.flatten(2, 3) |
| angles = angles.tile(2) |
| angles = angles.unsqueeze(2) |
|
|
| cos = torch.cos(angles) |
| sin = torch.sin(angles) |
|
|
| return cos.to(dtype=img_ids.dtype), sin.to(dtype=img_ids.dtype) |
|
|
|
|
| @maybe_allow_in_graph |
| class TransformerBlock(nn.Module): |
| def __init__( |
| self, |
| heads: int, |
| dim_head: int, |
| embed_dim: Optional[int] = None, |
| ffn_glu_balanced: bool = False, |
| norm_type: str = "layer_norm", |
| norm_affine: bool = True, |
| qk_norm_type: str = "rms_norm", |
| qk_norm_affine: bool = False, |
| ffn_activation_fn: str = "silu", |
| ffn_use_gated: bool = True, |
| use_scale: bool = True, |
| bias: bool = True, |
| eps: float = 1e-5, |
| **kwargs, |
| ): |
| super().__init__() |
| dim = embed_dim if embed_dim is not None else dim_head * heads |
| self.use_scale = use_scale |
|
|
| if norm_type == "layer_norm": |
| norm_class = nn.LayerNorm |
| elif norm_type == "rms_norm": |
| norm_class = nn.RMSNorm |
| else: |
| raise ValueError(f"unknown norm_type {norm_type}") |
|
|
| self.norm1 = norm_class( |
| dim, |
| elementwise_affine=norm_affine, |
| eps=eps, |
| ) |
| self.attn = Attention( |
| heads=heads, |
| dim_head=dim_head, |
| embed_dim=dim, |
| qk_norm_type=qk_norm_type, |
| qk_norm_affine=qk_norm_affine, |
| bias=bias, |
| eps=eps, |
| **kwargs, |
| ) |
| if use_scale: |
| self.scale1 = nn.Parameter(torch.zeros(dim)) |
|
|
| self.norm2 = norm_class( |
| dim, |
| elementwise_affine=norm_affine, |
| eps=eps, |
| ) |
| self.ff = FeedForward( |
| dim=dim, |
| activation_fn=ffn_activation_fn, |
| bias=bias, |
| use_gated=ffn_use_gated, |
| glu_balanced=ffn_glu_balanced, |
| ) |
| if use_scale: |
| self.scale2 = nn.Parameter(torch.zeros(dim)) |
|
|
| def forward( |
| self, |
| hidden_states: torch.FloatTensor, |
| rotary_pos_emb: Optional[torch.FloatTensor] = None, |
| pack_info: dict = {}, |
| ): |
| norm_hidden_states = self.norm1(_vit_norm_input(self.norm1, hidden_states)).to(hidden_states.dtype) |
| attn_output = self.attn(norm_hidden_states, rotary_pos_emb, pack_info) |
| if self.use_scale: |
| hidden_states = hidden_states + attn_output * self.scale1 |
| else: |
| hidden_states = hidden_states + attn_output |
|
|
| norm_hidden_states = self.norm2(_vit_norm_input(self.norm2, hidden_states)).to(hidden_states.dtype) |
| ff_output = self.ff(norm_hidden_states) |
| if self.use_scale: |
| hidden_states = hidden_states + ff_output * self.scale2 |
| else: |
| hidden_states = hidden_states + ff_output |
|
|
| return hidden_states |
|
|