# SPDX-License-Identifier: Apache-2.0
# Transformer building blocks for the MiniMax H3 visual VAE ViT decoder.
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__)  # pylint: disable=invalid-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
