#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
#           This file was automatically generated from src/transformers/models/neomme/modular_neomme.py.
#               Do NOT edit this file manually as any edits will be overwritten by the generation of
#             the file from the modular. If any change should be done, please apply the change to the
#                          modular_neomme.py file directly. One of our CI enforces this.
#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
# Copyright 2026 H Company and the HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from collections.abc import Callable
from dataclasses import dataclass

import torch
import torch.nn.functional as F
from torch import nn

from ... import initialization as init
from ...activations import ACT2FN
from ...masking_utils import create_bidirectional_mask, create_bidirectional_sliding_window_mask
from ...modeling_layers import GradientCheckpointingLayer
from ...modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling, MaskedLMOutput
from ...modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
from ...processing_utils import Unpack
from ...utils import TransformersKwargs, auto_docstring, torch_compilable_check
from ...utils.deprecation import deprecate_kwarg
from ...utils.generic import can_return_tuple, maybe_autocast
from ...utils.output_capturing import capture_outputs
from .configuration_neomme import NeoMMEConfig


class NeoMMERMSNorm(nn.Module):
    def __init__(self, dim: int, eps: float = 1e-6, with_scale: bool = True):
        super().__init__()
        self.eps = eps
        self.with_scale = with_scale

        if self.with_scale:
            self.weight = nn.Parameter(torch.ones(dim), requires_grad=True)

    def _norm(self, hidden_states: torch.Tensor):
        mean_squared = hidden_states.pow(2).mean(-1, keepdim=True) + self.eps
        # Use torch.pow() (over torch.sqrt() or torch.rsqrt()) to address compiler differences between Torch and JAX
        return hidden_states * torch.pow(mean_squared, -0.5)

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        normed_output = self._norm(hidden_states.float())
        if self.with_scale:
            normed_output = normed_output * self.weight.float()
        return normed_output.type_as(hidden_states)


class NeoMMEEmbeddings(nn.Module):
    """Factorized (ALBERT-style) token embeddings: `vocab_size -> embedding_rank -> hidden_size`."""

    def __init__(self, config: NeoMMEConfig):
        super().__init__()
        self.word_embeddings = nn.Embedding(config.vocab_size, config.embedding_rank)
        self.embedding_projection = nn.Linear(config.embedding_rank, config.hidden_size, bias=False)

    def forward(self, input_ids: torch.LongTensor) -> torch.Tensor:
        return self.embedding_projection(self.word_embeddings(input_ids))


class NeoMMEPatchEmbeddings(nn.Module):
    """Patch stem that maps flattened image patches to hidden size."""

    def __init__(self, config: NeoMMEConfig):
        super().__init__()
        self.norm = nn.LayerNorm(config.patch_dim)
        self.up_proj = nn.Linear(config.patch_dim, config.hidden_size * 2, bias=False)
        self.act_fn = nn.GELU()
        self.down_proj = nn.Linear(config.hidden_size * 2, config.hidden_size, bias=True)

    def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
        hidden_states = self.up_proj(self.norm(pixel_values))
        return self.down_proj(self.act_fn(hidden_states))


class NeoMMERotaryEmbedding(nn.Module):
    """Two-axis interleaved M-RoPE with per-layer-type frequency spectra."""

    @deprecate_kwarg("device", version="5.18")
    def __init__(self, config: NeoMMEConfig, device=None):
        super().__init__()
        self.max_seq_len_cached = config.max_position_embeddings
        self.original_max_seq_len = config.max_position_embeddings
        self.config = config
        self.layer_types = sorted(set(config.layer_types))
        self.rope_type = {}
        for layer_type in self.layer_types:
            rope_params = self.config.rope_parameters[layer_type]
            if rope_params is None:
                continue

            self.rope_type[layer_type] = rope_params["rope_type"]
            rope_init_fn: Callable = self.compute_default_rope_parameters
            if self.rope_type[layer_type] != "default":
                rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type[layer_type]]
            curr_inv_freq, curr_attention_scaling = rope_init_fn(self.config, device, layer_type=layer_type)
            setattr(self, f"{layer_type}_inv_freq", nn.Buffer(curr_inv_freq, persistent=False))
            setattr(self, f"{layer_type}_original_inv_freq", nn.Buffer(curr_inv_freq.clone(), persistent=False))
            setattr(self, f"{layer_type}_attention_scaling", curr_attention_scaling)

    @staticmethod
    @deprecate_kwarg("device", version="5.18")
    def compute_default_rope_parameters(
        config: NeoMMEConfig, device=None, layer_type: str | None = None, **kwargs
    ) -> tuple[torch.Tensor, float]:
        """
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            layer_type (`str`, *optional*):
                The current layer type if the model has different RoPE parameters per type.
                Should not be used unless `config.layer_types is not None`
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        """
        base = config.rope_parameters[layer_type]["rope_theta"]
        # key difference to gemma3: partial rope
        partial_rotary_factor = config.rope_parameters[layer_type].get("partial_rotary_factor", 1.0)
        head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
        dim = int(head_dim * partial_rotary_factor)

        attention_factor = 1.0  # Unused in this type of RoPE

        # Compute the inverse frequencies
        inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float) / dim))
        return inv_freq.to(device), attention_factor

    @torch.no_grad()
    @dynamic_rope_update  # power user: used with advanced RoPE types (e.g. dynamic rope)
    def forward(
        self, x: torch.Tensor, position_ids: torch.LongTensor, layer_type: str | None = None
    ) -> tuple[torch.Tensor, torch.Tensor]:
        inv_freq = getattr(self, f"{layer_type}_inv_freq")
        attention_scaling = getattr(self, f"{layer_type}_attention_scaling")

        inv_freq_expanded = inv_freq[None, None, :, None].float().expand(2, position_ids.shape[1], -1, 1)
        position_ids_expanded = position_ids[:, :, None, :].float()  # (2, batch, 1, seq_len)

        device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
        with maybe_autocast(device_type=device_type, enabled=False):
            # (2, batch, seq_len, rotary_dim // 2)
            freqs = (inv_freq_expanded @ position_ids_expanded).transpose(2, 3)
            cos = freqs.cos() * attention_scaling
            sin = freqs.sin() * attention_scaling

        cos = self.recomposition_frequencies(cos)
        sin = self.recomposition_frequencies(sin)
        return cos.to(x.dtype), sin.to(x.dtype)

    def recomposition_frequencies(self, freq):
        """
        Recompose the frequencies into the final spatial layout used per each grid.
        """
        # in contrast to the H-H-W-W layout, row/col here interleave as row0-col0-row1-col1
        # i.e. `mrope_section = [head_dim//4, head_dim//4]`
        freq_row, freq_col = freq[0][..., 0::2], freq[1][..., 1::2]
        angles = torch.stack([freq_row, freq_col], dim=-1).flatten(-2)
        return torch.cat([angles, angles], dim=-1)


def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
    """
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    """
    batch, num_key_value_heads, slen, head_dim = hidden_states.shape
    if n_rep == 1:
        return hidden_states
    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)


class NeoMMEExclusiveSelfAttention(nn.Module):
    def __init__(self, config: NeoMMEConfig):
        super().__init__()
        self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
        self.alpha = nn.Parameter(torch.zeros(config.num_attention_heads))

    def forward(self, attn_output: torch.Tensor, value_states: torch.Tensor) -> torch.Tensor:
        value_states = repeat_kv(value_states, self.num_key_value_groups).transpose(1, 2)
        value_unit = F.normalize(value_states.float(), dim=-1).to(attn_output.dtype)
        projection = (attn_output * value_unit).sum(-1, keepdim=True)
        scale = torch.tanh(self.alpha).to(attn_output.dtype).view(1, 1, -1, 1)
        return attn_output - (scale * projection) * value_unit


class NeoMMESigmoidGatedProjection(nn.Module):
    def __init__(self, config: NeoMMEConfig):
        super().__init__()
        projection_size = config.num_attention_heads * config.head_dim
        self.gate_proj = nn.Linear(config.hidden_size, projection_size, bias=False)
        self.o_proj = nn.Linear(projection_size, config.hidden_size, bias=config.attention_bias)

    def forward(self, attn_output: torch.Tensor, hidden_states: torch.Tensor) -> torch.Tensor:
        return self.o_proj(attn_output * torch.sigmoid(self.gate_proj(hidden_states)))


def rotate_half(x):
    """Rotates half the hidden dims of the input."""
    x1 = x[..., : x.shape[-1] // 2]
    x2 = x[..., x.shape[-1] // 2 :]
    return torch.cat((-x2, x1), dim=-1)


def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):
    """Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    """
    cos = cos.unsqueeze(unsqueeze_dim)
    sin = sin.unsqueeze(unsqueeze_dim)

    # Keep half or full tensor for later concatenation
    rotary_dim = cos.shape[-1]
    q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:]
    k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:]

    # Apply rotary embeddings on the first half or full tensor
    q_embed = (q_rot * cos) + (rotate_half(q_rot) * sin)
    k_embed = (k_rot * cos) + (rotate_half(k_rot) * sin)

    # Concatenate back to full shape
    q_embed = torch.cat([q_embed, q_pass], dim=-1)
    k_embed = torch.cat([k_embed, k_pass], dim=-1)
    return q_embed, k_embed


def eager_attention_forward(
    module: nn.Module,
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    attention_mask: torch.Tensor | None,
    scaling: float,
    dropout: float = 0.0,
    **kwargs: Unpack[TransformersKwargs],
):
    key_states = repeat_kv(key, module.num_key_value_groups)
    value_states = repeat_kv(value, module.num_key_value_groups)

    attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling
    if attention_mask is not None:
        attn_weights = attn_weights + attention_mask

    attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)
    attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
    attn_output = torch.matmul(attn_weights, value_states)
    attn_output = attn_output.transpose(1, 2).contiguous()

    return attn_output, attn_weights


class NeoMMEAttention(nn.Module):
    """Bidirectional grouped-query attention with QK-norm, M-RoPE, and a sigmoid output gate.

    QK-norm runs before rotary embedding, value embeddings are added after rotation, and
    exclusive self-attention is applied before the output gate.
    """

    def __init__(self, config: NeoMMEConfig, layer_idx: int):
        super().__init__()
        self.config = config
        self.layer_idx = layer_idx
        self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
        self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
        self.scaling = self.head_dim**-0.5
        self.attention_dropout = config.attention_dropout
        self.is_causal = False

        self.q_proj = nn.Linear(
            config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias
        )
        self.k_proj = nn.Linear(
            config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
        )
        self.v_proj = nn.Linear(
            config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
        )
        # Parent LlamaAttention already sets: layer_idx, num_heads, num_key_value_heads, num_key_value_groups, head_dim
        # We only add NeoMME-specific attributes
        self.is_local_attention = config.layer_types[layer_idx] == "sliding_attention"

        # `sliding_window` is a HALF-width (`abs(i - j) <= window`). The flash-attention path
        # builds an inclusive symmetric band of `sliding_window - 1` per side, hence the `+ 1`.
        self.sliding_window = (
            config.per_layer_config[layer_idx].sliding_window + 1 if self.is_local_attention else None
        )

        self.attention_type = config.layer_types[layer_idx]
        self.num_attention_heads = config.num_attention_heads
        self.q_norm = NeoMMERMSNorm(config.head_dim, config.norm_eps, with_scale=False)
        self.k_norm = NeoMMERMSNorm(config.head_dim, config.norm_eps, with_scale=False)
        self.exclusive_self_attention = NeoMMEExclusiveSelfAttention(config)
        self.output_projection = NeoMMESigmoidGatedProjection(config)

    def forward(
        self,
        hidden_states: torch.Tensor,
        position_embeddings: tuple[torch.Tensor, torch.Tensor],
        attention_mask: torch.Tensor | None = None,
        value_embeds: torch.Tensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> tuple[torch.Tensor, torch.Tensor | None]:
        input_shape = hidden_states.shape[:-1]
        hidden_shape = (*input_shape, -1, self.head_dim)

        query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
        key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
        value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)

        query_states = self.q_norm(query_states)
        key_states = self.k_norm(key_states)

        cos, sin = position_embeddings
        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
        if value_embeds is not None:
            value_states = value_states + value_embeds.view(hidden_shape).transpose(1, 2)

        attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
            self.config._attn_implementation, eager_attention_forward
        )

        attn_output, attn_weights = attention_interface(
            self,
            query_states,
            key_states,
            value_states,
            attention_mask,
            dropout=self.attention_dropout if self.training else 0.0,
            scaling=self.scaling,
            sliding_window=self.sliding_window,
            **kwargs,
        )

        attn_output = self.exclusive_self_attention(attn_output, value_states)
        attn_output = attn_output.reshape(*input_shape, -1)
        attn_output = self.output_projection(attn_output, hidden_states)
        return attn_output, attn_weights


class NeoMMEMLP(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.config = config
        self.hidden_size = config.hidden_size
        self.intermediate_size = config.intermediate_size
        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias)
        self.act_fn = ACT2FN[config.hidden_act]

    def forward(self, x):
        return self.down_proj(self.act_fn(self.up_proj(x)))


class NeoMMEEncoderLayer(GradientCheckpointingLayer):
    """Pre-norm encoder layer with initial-state mixing and muP depth scaling."""

    def __init__(self, config: NeoMMEConfig, layer_idx: int):
        super().__init__()
        self.self_attn = NeoMMEAttention(config, layer_idx)
        self.mlp = NeoMMEMLP(config)
        self.lambdas = nn.Parameter(torch.tensor([1.0, 0.0]))
        self.input_layernorm = NeoMMERMSNorm(config.hidden_size, config.norm_eps, with_scale=False)
        self.post_attention_layernorm = NeoMMERMSNorm(config.hidden_size, config.norm_eps, with_scale=False)
        self.residual_multiplier = config.residual_multiplier
        self.attention_type = config.layer_types[layer_idx]

    def forward(
        self,
        hidden_states: torch.Tensor,
        initial_hidden_states: torch.Tensor,
        value_embeds: torch.Tensor | None = None,
        position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
        attention_mask: torch.Tensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> torch.Tensor:
        mixed_states = self.lambdas[0] * hidden_states + self.lambdas[1] * initial_hidden_states
        normed_states = self.input_layernorm(mixed_states)
        attn_output, _ = self.self_attn(
            normed_states,
            position_embeddings=position_embeddings,
            attention_mask=attention_mask,
            value_embeds=value_embeds,
            **kwargs,
        )
        hidden_states = hidden_states + self.residual_multiplier * attn_output

        normed_states = self.post_attention_layernorm(hidden_states)
        mlp_output = self.mlp(normed_states)
        return hidden_states + self.residual_multiplier * mlp_output


@auto_docstring
class NeoMMEPreTrainedModel(PreTrainedModel):
    config: NeoMMEConfig
    base_model_prefix = "model"
    supports_gradient_checkpointing = True
    _no_split_modules = ["NeoMMEEncoderLayer"]
    _skip_keys_device_placement = ["past_key_values"]
    _supports_flash_attn = True
    _supports_sdpa = True
    _supports_flex_attn = True

    _can_compile_fullgraph = True
    _supports_attention_backend = True
    _can_record_outputs = {
        "hidden_states": NeoMMEEncoderLayer,
        "attentions": NeoMMEAttention,
    }
    input_modalities = ("image", "text")

    def get_input_embeddings(self) -> nn.Embedding:
        backbone = getattr(self, self.base_model_prefix, self)
        return backbone.embed_tokens.word_embeddings

    def set_input_embeddings(self, value: nn.Embedding) -> None:
        backbone = getattr(self, self.base_model_prefix, self)
        backbone.embed_tokens.word_embeddings = value

    @torch.no_grad()
    def _init_weights(self, module: nn.Module):
        # `apply` visits children before parents, so the NeoMME-specific parent initialization below runs last.
        super()._init_weights(module)

        if isinstance(module, NeoMMEEmbeddings):
            init.normal_(module.word_embeddings.weight, mean=0.0, std=self.config.embedding_rank**-0.5)
        elif isinstance(module, NeoMMEExclusiveSelfAttention):
            init.zeros_(module.alpha)
        elif isinstance(module, NeoMMESigmoidGatedProjection):
            # Zero-init so the attention residual contributes nothing at initialization.
            init.zeros_(module.o_proj.weight)
        elif isinstance(module, NeoMMEMLP):
            init.zeros_(module.down_proj.weight)
        elif isinstance(module, NeoMMEEncoderLayer):
            init.copy_(module.lambdas, torch.tensor([1.0, 0.0]))
        elif isinstance(module, NeoMMEModel):
            init.zeros_(module.value_embeddings.weight)
        elif isinstance(module, NeoMMEForMaskedLM):
            init.normal_(module.lm_head.weight, mean=0.0, std=self.config.embedding_rank**-0.5)
        elif isinstance(module, NeoMMERotaryEmbedding):
            for layer_type in module.layer_types:
                rope_init_fn = module.compute_default_rope_parameters
                if module.rope_type[layer_type] != "default":
                    rope_init_fn = ROPE_INIT_FUNCTIONS[module.rope_type[layer_type]]
                inv_freq, _ = rope_init_fn(module.config, layer_type=layer_type)
                init.copy_(getattr(module, f"{layer_type}_inv_freq"), inv_freq)
                init.copy_(getattr(module, f"{layer_type}_original_inv_freq"), inv_freq)

    def _resize_token_embeddings(
        self, new_num_tokens: int, pad_to_multiple_of: int | None = None, mean_resizing: bool = True
    ) -> nn.Embedding:
        """Resize word and value embedding tables together."""
        word_embeddings = super()._resize_token_embeddings(new_num_tokens, pad_to_multiple_of, mean_resizing)
        backbone = getattr(self, self.base_model_prefix, self)
        resized = self._get_resized_embeddings(
            backbone.value_embeddings, word_embeddings.weight.shape[0], mean_resizing=mean_resizing
        )
        backbone.value_embeddings = nn.Embedding(resized.num_embeddings, resized.embedding_dim)
        backbone.value_embeddings.weight = resized.weight
        return word_embeddings


@auto_docstring(
    custom_intro="""
    The bare NeoMME model. It encodes text tokens and image patches with one bidirectional Transformer encoder.
    """
)
class NeoMMEModel(NeoMMEPreTrainedModel):
    def __init__(self, config: NeoMMEConfig):
        super().__init__(config)
        self.embed_tokens = NeoMMEEmbeddings(config)
        self.patch_embeddings = NeoMMEPatchEmbeddings(config)
        self.rotary_emb = NeoMMERotaryEmbedding(config)
        self.embedding_norm = NeoMMERMSNorm(config.hidden_size, config.norm_eps, with_scale=False)
        self.final_norm = NeoMMERMSNorm(config.hidden_size, config.norm_eps, with_scale=False)
        self.layers = nn.ModuleList(
            [NeoMMEEncoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
        )

        global_layers = [i for i, layer_type in enumerate(config.layer_types) if layer_type == "full_attention"]
        self.value_embeddings = nn.Embedding(config.vocab_size, config.num_key_value_heads * config.head_dim)
        self.value_embedding_layers = {global_layers[0], global_layers[-1]}
        self.gradient_checkpointing = False
        self.post_init()

    @can_return_tuple
    @auto_docstring(custom_intro="Projects flattened image patches into the model hidden space.")
    def get_image_features(
        self, pixel_values: torch.Tensor, **kwargs: Unpack[TransformersKwargs]
    ) -> BaseModelOutputWithPooling:
        if pixel_values.shape[-1] != self.config.patch_dim:  # trf-ignore: TRF041
            raise ValueError(
                f"pixel_values has patch width {pixel_values.shape[-1]} but the model expects "
                f"{self.config.patch_dim} (= 3 * patch_size ** 2 with patch_size={self.config.patch_size})"
            )
        image_features = self.patch_embeddings(pixel_values.to(self.patch_embeddings.norm.weight.dtype))
        return BaseModelOutputWithPooling(last_hidden_state=image_features, pooler_output=image_features)

    def get_placeholder_mask(self, input_ids: torch.LongTensor, image_features: torch.Tensor) -> torch.BoolTensor:
        """Find patch placeholders and validate that every image feature has a destination."""
        previous_ids = F.pad(input_ids[:, :-1], (1, 0), value=self.config.pad_token_id or 0)  # token IDs shifted right
        image_mask = (input_ids == self.config.image_token_id) & (previous_ids != self.config.document_token_id)

        num_image_tokens = image_mask.sum()
        torch_compilable_check(
            num_image_tokens == image_features.shape[0],
            lambda: f"Got {image_features.shape[0]} image patches for {int(num_image_tokens)} image placeholder tokens",
        )
        return image_mask.unsqueeze(-1).expand(-1, -1, image_features.shape[-1])

    @capture_outputs
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        pixel_values: torch.Tensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> BaseModelOutput:
        r"""
        position_ids (`torch.LongTensor` of shape `(2, batch_size, sequence_length)` or `(batch_size, sequence_length)`, *optional*):
            Positions for the input tokens. [`NeoMMEProcessor`] returns two-axis positions for document images. A
            one-axis position tensor is used for text inputs.
        pixel_values (`torch.Tensor` of shape `(num_patches, 3 * patch_size ** 2)`, *optional*):
            Flattened image patches returned by [`NeoMMEProcessor`]. The model places these patches at image
            placeholders in `input_ids`.
        """
        hidden_states = self.embed_tokens(input_ids)
        if pixel_values is not None:
            image_outputs = self.get_image_features(pixel_values, return_dict=True)
            image_mask = self.get_placeholder_mask(input_ids, image_outputs.pooler_output)
            hidden_states = hidden_states.masked_scatter(image_mask, image_outputs.pooler_output)

        batch_size, seq_len = hidden_states.shape[:2]
        # create 2D positions - text uses the token index for both M-RoPE axes
        if position_ids is None:
            position_ids = torch.arange(seq_len, device=hidden_states.device).expand(batch_size, -1)
        if position_ids.ndim == 2:
            position_ids = position_ids.unsqueeze(0).expand(2, -1, -1)

        # Reuse this normalized input as `initial_hidden_states` in every encoder layer.
        hidden_states = initial_hidden_states = self.embedding_norm(hidden_states)

        if not isinstance(attention_mask_mapping := attention_mask, dict):
            attention_mask_mapping: dict[int, torch.Tensor | None] = {}
            mask_kwargs = {"inputs_embeds": hidden_states, "attention_mask": attention_mask}
            for layer_id in range(self.config.num_hidden_layers):
                per_layer_config = self.config.per_layer_config[layer_id]
                if per_layer_config.sliding_window is not None:
                    attention_mask_mapping[layer_id] = create_bidirectional_sliding_window_mask(
                        config=per_layer_config,
                        **mask_kwargs,
                    )
                else:
                    attention_mask_mapping[layer_id] = create_bidirectional_mask(
                        config=per_layer_config,
                        **mask_kwargs,
                    )

        position_embeddings = {
            layer_type: self.rotary_emb(hidden_states, position_ids, layer_type)
            for layer_type in set(self.config.layer_types)
        }

        value_embeds = self.value_embeddings(input_ids)

        for layer_idx, encoder_layer in enumerate(self.layers):
            # Pass gradient-carrying tensors positionally so reentrant checkpointing tracks them.
            hidden_states = encoder_layer(
                hidden_states,
                initial_hidden_states,
                value_embeds if layer_idx in self.value_embedding_layers else None,
                position_embeddings=position_embeddings[encoder_layer.attention_type],
                attention_mask=attention_mask_mapping[layer_idx],
                **kwargs,
            )

        hidden_states = self.final_norm(hidden_states)
        return BaseModelOutput(last_hidden_state=hidden_states)


@auto_docstring(
    custom_intro="""
    The NeoMME model with a factorized masked token decoder.
    """
)
class NeoMMEForMaskedLM(NeoMMEPreTrainedModel):
    _tied_weights_keys = {
        "lm_head.weight": "model.embed_tokens.word_embeddings.weight",
        "unembedding_projection.weight": "model.embed_tokens.embedding_projection.weight",
    }

    def __init__(self, config: NeoMMEConfig):
        super().__init__(config)
        self.model = NeoMMEModel(config)
        self.unembedding_projection = nn.Linear(config.embedding_rank, config.hidden_size, bias=False)
        self.lm_head = nn.Linear(config.embedding_rank, config.vocab_size, bias=False)
        self.post_init()

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        pixel_values: torch.Tensor | None = None,
        labels: torch.LongTensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> MaskedLMOutput:
        r"""
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for the masked-language-modeling loss. Indices should be in `[0, ..., config.vocab_size - 1]`
            or `-100`; only tokens with a label different from `-100` contribute.
        """
        outputs: BaseModelOutput = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            position_ids=position_ids,
            pixel_values=pixel_values,
            **kwargs,
        )
        hidden_states = outputs.last_hidden_state
        hidden_states = hidden_states @ self.unembedding_projection.weight
        logits = self.lm_head(hidden_states)

        loss = None
        if labels is not None:
            loss = self.loss_function(logits, labels, vocab_size=self.config.vocab_size, **kwargs)

        return MaskedLMOutput(
            loss=loss,
            logits=logits,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
        )


@auto_docstring(
    custom_intro="""
    Output type for [`NeoMMEForRetrieval`].
    """
)
@dataclass
class NeoMMEForRetrievalOutput(BaseModelOutput):
    r"""
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*):
        Retrieval loss. This value is always `None`.
    embeddings (`torch.FloatTensor` of shape `(batch_size, sequence_length, embedding_dim)`, *optional*):
        Normalized token embeddings for late-interaction retrieval. Padding rows are zeroed. Score them with MeanMaxSim.
    dense_embeddings (`torch.FloatTensor` of shape `(batch_size, hidden_size)` or `(batch_size, dense_dim)`, *optional*):
        A normalized mean-pooled embedding for each input. When `dense_dim` is set, the last dimension is `dense_dim`.
        Score them with cosine similarity.
    """

    loss: torch.FloatTensor | None = None
    embeddings: torch.FloatTensor | None = None
    dense_embeddings: torch.FloatTensor | None = None


class NeoMMEMultiVectorHead(nn.Module):
    def __init__(self, config: NeoMMEConfig):
        super().__init__()
        self.proj = nn.Linear(config.hidden_size, config.embedding_dim, bias=False)

    def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
        embeddings = self.proj(hidden_states.to(self.proj.weight.dtype))
        embeddings = F.normalize(embeddings, dim=-1)
        # Use masked_fill because multiplying NaN or Inf by zero would leave padding non-finite.
        return embeddings.masked_fill(~attention_mask.bool().unsqueeze(-1), 0.0)


class NeoMMEDenseHead(nn.Module):
    def forward(
        self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, dense_dim: int | None = None
    ) -> torch.Tensor:
        expanded_mask = attention_mask.unsqueeze(-1).expand(hidden_states.shape).to(hidden_states.dtype)
        pooled = (hidden_states * expanded_mask).sum(1) / expanded_mask.sum(1).clamp_min(1e-9)
        if dense_dim is None:
            return F.normalize(pooled, dim=-1)

        if not 0 < dense_dim <= pooled.shape[-1]:
            raise ValueError(f"dense_dim must be in 1..{pooled.shape[-1]} (the pooled width), got {dense_dim}")
        return F.normalize(pooled[..., :dense_dim], dim=-1)


@auto_docstring(
    custom_intro="""
    The NeoMME model with multi-vector and dense retrieval heads. One forward pass can return token embeddings for
    MaxSim scoring and mean-pooled embeddings for cosine similarity.
    """
)
class NeoMMEForRetrieval(NeoMMEPreTrainedModel):
    def __init__(self, config: NeoMMEConfig):
        super().__init__(config)
        self.model = NeoMMEModel(config)
        self.multi_vector_head = NeoMMEMultiVectorHead(config)
        self.dense_head = NeoMMEDenseHead()
        self.post_init()

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        pixel_values: torch.Tensor | None = None,
        output_multivector: bool = True,
        output_dense: bool = True,
        dense_dim: int | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> NeoMMEForRetrievalOutput:
        r"""
        output_multivector (`bool`, *optional*, defaults to `True`):
            Whether to return token embeddings for late-interaction retrieval.
        output_dense (`bool`, *optional*, defaults to `True`):
            Whether to return one mean-pooled dense embedding per input.
        dense_dim (`int`, *optional*):
            Width of the Matryoshka prefix to return for dense embeddings. The model truncates the pooled vector
            before normalizing it.
        """
        if not (output_multivector or output_dense):
            raise ValueError("At least one of `output_multivector` or `output_dense` must be True")

        outputs: BaseModelOutput = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            position_ids=position_ids,
            pixel_values=pixel_values,
            **kwargs,
        )
        hidden_states = outputs.last_hidden_state
        if attention_mask is None:
            attention_mask = torch.ones(hidden_states.shape[:2], dtype=torch.bool, device=hidden_states.device)

        embeddings = self.multi_vector_head(hidden_states, attention_mask) if output_multivector else None
        dense_embeddings = self.dense_head(hidden_states, attention_mask, dense_dim) if output_dense else None
        return NeoMMEForRetrievalOutput(
            embeddings=embeddings,
            dense_embeddings=dense_embeddings,
            last_hidden_state=hidden_states,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
        )


__all__ = ["NeoMMEForMaskedLM", "NeoMMEForRetrieval", "NeoMMEModel", "NeoMMEPreTrainedModel"]
