#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
#           This file was automatically generated from src/transformers/models/inkling/modular_inkling.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_inkling.py file directly. One of our CI enforces this.
#                🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
# Copyright 2026 the HuggingFace 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.
# coding=utf-8

import math
from collections.abc import Callable
from dataclasses import dataclass

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

from ... import initialization as init
from ...activations import ACT2FN
from ...cache_utils import Cache, DynamicCache
from ...generation import GenerationMixin
from ...integrations import (
    use_experts_implementation,
    use_kernel_forward_from_hub,
    use_kernel_func_from_hub,
    use_kernelized_func,
)
from ...integrations.accelerate import force_accelerate_hooks
from ...masking_utils import create_causal_mask, create_recurrent_attention_mask, create_sliding_window_causal_mask
from ...modeling_layers import GradientCheckpointingLayer
from ...modeling_outputs import BaseModelOutputWithPast, BaseModelOutputWithPooling
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
from ...processing_utils import Unpack
from ...utils import ModelOutput, TransformersKwargs, auto_docstring, can_return_tuple, torch_compilable_check
from ...utils.generic import merge_with_config_defaults
from ...utils.output_capturing import capture_outputs
from .configuration_inkling import InklingAudioConfig, InklingConfig, InklingTextConfig, InklingVisionConfig


@auto_docstring(
    custom_intro="""
    Base class for Inkling outputs, with hidden states and attentions.
    """
)
@dataclass
class InklingModelOutputWithPast(BaseModelOutputWithPast):
    r"""
    image_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        image_hidden_states of the model produced by the vision encoder and after projecting the last hidden state.
    """

    image_hidden_states: torch.FloatTensor | None = None


@auto_docstring(
    custom_intro="""
    Base class for Inkling causal language model (or autoregressive) outputs.
    """
)
@dataclass
class InklingCausalLMOutputWithPast(ModelOutput):
    r"""
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss (for next-token prediction).
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.text_config.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
        `past_key_values` input) to speed up sequential decoding.
    image_hidden_states (`torch.FloatTensor`, *optional*):
        A `torch.FloatTensor` of size `(batch_size, num_images, sequence_length, hidden_size)`.
        image_hidden_states of the model produced by the vision encoder after projecting last hidden state.
    """

    loss: torch.FloatTensor | None = None
    logits: torch.FloatTensor | None = None
    past_key_values: Cache | None = None
    hidden_states: tuple[torch.FloatTensor] | None = None
    attentions: tuple[torch.FloatTensor] | None = None
    image_hidden_states: torch.FloatTensor | None = None


@use_kernel_forward_from_hub("RMSNorm")
class InklingRMSNorm(nn.Module):
    def __init__(self, hidden_size, eps: float = 1e-6) -> None:
        """
        InklingRMSNorm is equivalent to T5LayerNorm
        """
        super().__init__()
        self.weight = nn.Parameter(torch.ones(hidden_size))
        self.variance_epsilon = eps

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        input_dtype = hidden_states.dtype
        hidden_states = hidden_states.to(torch.float32)
        variance = hidden_states.pow(2).mean(-1, keepdim=True)
        hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
        return self.weight * hidden_states.to(input_dtype)

    def extra_repr(self):
        return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"


class InklingRelativeLogits(nn.Module):
    """hidden states conditioned relative position bias. `proj` is a trained bank of bias-vs-distance profiles; each token's
    `relative_states` mixes them into one bias value per backward distance
    (`sglang RelLogitsProj` + the FA4 `score_mod`, materialized densely). The bias is zero
    outside `0 <= distance < rel_extent`; causality and padding stay in the attention mask.
    """

    def __init__(self, d_rel: int, rel_extent: int):
        super().__init__()
        self.rel_extent = rel_extent
        self.proj = nn.Parameter(torch.empty(d_rel, rel_extent))

    def forward(
        self,
        relative_states: torch.Tensor,
        query_positions: torch.Tensor,
        key_positions: torch.Tensor,
    ) -> torch.Tensor:
        # relative_states: [batch, q_len, num_heads, d_rel] -> bias: [batch, num_heads, q_len, kv_len]
        rel_logits = (relative_states @ self.proj).transpose(1, 2)
        distance = (query_positions[:, None] - key_positions[None, :])[None, None, :, :]
        gather_index = distance.clamp(0, self.rel_extent - 1).expand(*rel_logits.shape[:2], -1, -1)
        position_bias = rel_logits.gather(-1, gather_index)
        return position_bias.masked_fill((distance < 0) | (distance >= self.rel_extent), 0.0)


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)


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,
    position_bias: torch.Tensor | None = None,
    **kwargs,
):
    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 position_bias is not None:
        attn_weights = attn_weights + position_bias
    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 InklingAttention(nn.Module):
    def __init__(self, config: InklingTextConfig, layer_idx: int):
        super().__init__()
        self.config = config
        self.layer_idx = layer_idx
        self.is_sliding = config.layer_types[self.layer_idx] == "hybrid_sliding"
        self.head_dim = config.swa_head_dim if self.is_sliding else config.head_dim
        self.num_heads = config.swa_num_attention_heads if self.is_sliding else config.num_attention_heads
        self.num_key_value_heads = config.swa_num_key_value_heads if self.is_sliding else config.num_key_value_heads
        self.num_key_value_groups = self.num_heads // self.num_key_value_heads
        self.sliding_window = config.sliding_window_size if self.is_sliding else None
        self.rel_extent = config.sliding_window_size if self.is_sliding else config.rel_extent
        # q/k are RMS-normalized per head, hence 1/d rather than 1/sqrt(d)
        self.scaling = 1.0 / self.head_dim
        self.attention_dropout = config.attention_dropout
        self.is_causal = True

        self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=False)
        self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.r_proj = nn.Linear(config.hidden_size, self.num_heads * config.d_rel, bias=False)
        self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=False)
        self.k_sconv = InklingShortConvolution(
            self.num_key_value_heads * self.head_dim, config.sconv_kernel_size, layer_idx, conv_idx=0
        )
        self.v_sconv = InklingShortConvolution(
            self.num_key_value_heads * self.head_dim, config.sconv_kernel_size, layer_idx, conv_idx=1
        )
        self.q_norm = InklingRMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.k_norm = InklingRMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.rel_logits_proj = InklingRelativeLogits(config.d_rel, self.rel_extent)

    def forward(
        self,
        hidden_states: torch.Tensor,
        attention_mask: torch.Tensor | None,
        conv_mask: torch.Tensor | None = None,
        past_key_values: Cache | 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)
        key_states = self.k_sconv(self.k_proj(hidden_states), past_key_values=past_key_values, conv_mask=conv_mask)
        value_states = self.v_sconv(self.v_proj(hidden_states), past_key_values=past_key_values, conv_mask=conv_mask)
        relative_states = self.r_proj(hidden_states)

        query_states = self.q_norm(query_states.view(hidden_shape)).transpose(1, 2)
        key_states = self.k_norm(key_states.view(hidden_shape)).transpose(1, 2)
        value_states = value_states.view(hidden_shape).transpose(1, 2)

        q_length = query_states.shape[2]
        if past_key_values is not None:
            # Important to get those values before updating the cache to be correct
            kv_length, kv_offset = past_key_values.get_mask_sizes(q_length, self.layer_idx)
            q_offset = past_key_values.get_query_offset(self.layer_idx)
            # Update the cache
            key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)
        else:
            kv_length = key_states.shape[2]
            q_offset, kv_offset = 0, 0

        kv_positions = torch.arange(kv_length, device=hidden_states.device) + kv_offset
        q_positions = torch.arange(q_length, device=hidden_states.device) + q_offset
        relative_states = relative_states.view(*input_shape, self.num_heads, -1)
        position_bias = self.rel_logits_proj(relative_states, q_positions, kv_positions)

        # original impl applies log scalnig in f32
        if not self.is_sliding and self.config.log_scaling_n_floor is not None:
            effective_n = (q_positions + 1).float()
            tau = 1.0 + self.config.log_scaling_alpha * torch.log(
                (effective_n / self.config.log_scaling_n_floor).clamp(min=1.0)
            )
            tau = tau.view(1, 1, -1, 1)
            query_states = (query_states.float() * tau).to(query_states.dtype)
            position_bias = (position_bias.float() * tau).to(position_bias.dtype)

        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=0.0 if not self.training else self.attention_dropout,
            scaling=self.scaling,
            sliding_window=self.sliding_window,
            position_bias=position_bias,
            **kwargs,
        )

        attn_output = attn_output.reshape(*input_shape, -1).contiguous()
        attn_output = self.o_proj(attn_output)
        return attn_output, attn_weights


class InklingMLP(nn.Module):
    def __init__(self, config: InklingTextConfig):
        super().__init__()
        self.config = config
        self.hidden_size = config.hidden_size
        self.intermediate_size = config.intermediate_size
        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
        self.act_fn = ACT2FN[config.hidden_act]
        self.global_scale = nn.Parameter(torch.ones(1))

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        hidden_states = self.down_proj(self.act_fn(self.gate_proj(hidden_states)) * self.up_proj(hidden_states))
        return hidden_states * self.global_scale


@use_experts_implementation
class InklingExperts(nn.Module):
    """Collection of expert weights stored as 3D tensors."""

    def __init__(self, config: InklingTextConfig):
        super().__init__()
        self.num_experts = config.n_routed_experts
        self.hidden_dim = config.hidden_size
        self.intermediate_dim = config.moe_intermediate_size
        self.gate_up_proj = nn.Parameter(torch.empty(self.num_experts, 2 * self.intermediate_dim, self.hidden_dim))
        self.down_proj = nn.Parameter(torch.empty(self.num_experts, self.hidden_dim, self.intermediate_dim))
        self.act_fn = ACT2FN[config.hidden_act]

    def forward(
        self,
        hidden_states: torch.Tensor,
        top_k_index: torch.Tensor,
        top_k_weights: torch.Tensor,
    ) -> torch.Tensor:
        final_hidden_states = torch.zeros_like(hidden_states)
        with torch.no_grad():
            expert_mask = torch.nn.functional.one_hot(top_k_index, num_classes=self.num_experts)
            expert_mask = expert_mask.permute(2, 1, 0)
            expert_hit = torch.greater(expert_mask.sum(dim=(-1, -2)), 0).nonzero()

        for expert_idx in expert_hit:
            expert_idx = expert_idx[0]
            if expert_idx == self.num_experts:
                continue
            top_k_pos, token_idx = torch.where(expert_mask[expert_idx])
            current_state = hidden_states[token_idx]
            gate, up = nn.functional.linear(current_state, self.gate_up_proj[expert_idx]).chunk(2, dim=-1)
            current_hidden_states = self.act_fn(gate) * up
            current_hidden_states = nn.functional.linear(current_hidden_states, self.down_proj[expert_idx])
            current_hidden_states = current_hidden_states * top_k_weights[token_idx, top_k_pos, None]
            final_hidden_states.index_add_(0, token_idx, current_hidden_states.to(final_hidden_states.dtype))

        return final_hidden_states


class InklingTopkRouter(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.num_experts = config.n_routed_experts
        self.n_shared_experts = config.n_shared_experts
        self.n_total_experts = self.num_experts + self.n_shared_experts
        self.hidden_dim = config.hidden_size
        self.route_scale = config.route_scale
        self.top_k = config.num_experts_per_tok

        self.weight = nn.Parameter(torch.empty(self.n_total_experts, config.hidden_size))
        self.global_scale = nn.Parameter(torch.ones(1))
        self.e_score_correction_bias = nn.Parameter(torch.empty(self.num_experts))

    def forward(self, hidden_states) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        flat = hidden_states.reshape(-1, self.hidden_dim)
        router_logits = F.linear(flat, self.weight)

        # same as `self.route_tokens_to_experts` from before, prob same as our MoE and can be copied
        scores = router_logits.sigmoid()
        routed_scores = scores[..., : -self.n_shared_experts]
        scores_for_choice = routed_scores + self.e_score_correction_bias
        topk_indices = torch.topk(scores_for_choice, self.top_k, dim=-1, sorted=False)[1]

        routed_logits = router_logits[..., : -self.n_shared_experts]
        shared_logits = router_logits[..., -self.n_shared_experts :]
        topk_logits = torch.cat([routed_logits.gather(-1, topk_indices), shared_logits], dim=-1)
        topk_log_probs = F.logsigmoid(topk_logits)
        topk_weights = torch.exp(topk_log_probs - torch.logsumexp(topk_log_probs, dim=-1, keepdim=True))

        topk_weights = topk_weights * self.route_scale * self.global_scale

        shared_gammas = topk_weights[..., -self.n_shared_experts :].contiguous()
        topk_weights = topk_weights[..., : self.top_k].contiguous()

        return routed_logits, topk_weights, topk_indices, shared_gammas


class InklingSharedExperts(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.n_shared_experts = config.n_shared_experts
        intermediate_dim = config.moe_intermediate_size
        # TP loader cuts shards on the raw tensor but validates shapes on the target, so a Transpose
        # conversion op breaks sharded loads. The runtime transpose(1, 2) is not a per-forward
        # cost: it is a stride-metadata view, so the same
        # matmul layout every nn.Linear runs
        self.gate_proj = nn.Parameter(torch.empty(config.n_shared_experts, intermediate_dim, config.hidden_size))
        self.up_proj = nn.Parameter(torch.empty(config.n_shared_experts, intermediate_dim, config.hidden_size))
        self.down_proj = nn.Parameter(torch.empty(config.n_shared_experts, config.hidden_size, intermediate_dim))
        self.act_fn = ACT2FN[config.hidden_act]

    def forward(self, hidden_states, gammas):
        input_shape = hidden_states.shape
        hidden_states = hidden_states.reshape(1, -1, input_shape[-1]).expand(self.n_shared_experts, -1, -1)
        gammas = gammas.reshape(-1, self.n_shared_experts, 1).transpose(0, 1)

        gate = torch.bmm(hidden_states, self.gate_proj.transpose(1, 2))
        up = torch.bmm(hidden_states, self.up_proj.transpose(1, 2))
        activated = self.act_fn(gate) * up * gammas
        down = torch.bmm(activated, self.down_proj.transpose(1, 2))

        out = down.float().sum(dim=0).to(hidden_states.dtype)
        return out.view(input_shape)


class InklingMoE(nn.Module):
    """Gate -> routed experts (+ shared experts), TML flavour."""

    def __init__(self, config):
        super().__init__()
        self.config = config
        self.gate = InklingTopkRouter(config)
        self.experts = InklingExperts(config)
        self.shared_experts = InklingSharedExperts(config)

    def forward(self, hidden_states) -> torch.Tensor:
        residuals = hidden_states
        input_shape = hidden_states.shape
        _, topk_weights, topk_indices, shared_gammas = self.gate(hidden_states)
        hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
        hidden_states = self.experts(hidden_states, topk_indices, topk_weights).view(*input_shape)
        hidden_states = hidden_states + self.shared_experts(residuals, gammas=shared_gammas)
        return hidden_states


def apply_mask_to_padding_states(hidden_states, attention_mask):
    """
    Tunes out the hidden states for padding tokens, see https://github.com/state-spaces/mamba/issues/66
    """
    # NOTE: attention mask is a 2D boolean tensor
    if attention_mask is not None and attention_mask.shape[1] > 1 and attention_mask.shape[0] > 1:
        dtype = hidden_states.dtype
        hidden_states = (hidden_states * attention_mask[:, :, None]).to(dtype)

    return hidden_states


@use_kernel_func_from_hub("causal_conv1d_update")
def causal_conv1d_update(
    hidden_states: torch.Tensor,
    conv_state: torch.Tensor,
    weight: nn.Parameter,
    bias: nn.Parameter | None = None,
    activation: str | None = None,
):
    _, hidden_size, seq_len = hidden_states.shape
    state_len = conv_state.shape[-1]

    hidden_states_new = torch.cat([conv_state, hidden_states], dim=-1).to(weight.dtype)
    conv_state.copy_(hidden_states_new[:, :, -state_len:])
    out = F.conv1d(hidden_states_new, weight.unsqueeze(1), bias, padding=0, groups=hidden_size)
    out = out[:, :, -seq_len:]
    if activation is not None:
        out = ACT2FN[activation](out)
    return out.to(hidden_states.dtype)


@use_kernel_func_from_hub("causal_conv1d_fn")
def causal_conv1d_fn(
    hidden_states: torch.Tensor,
    weight: nn.Parameter,
    bias: nn.Parameter | None = None,
    activation: str | None = None,
    **kwargs,
):
    _, hidden_size, seq_len = hidden_states.shape
    padding = weight.shape[-1] - 1

    out = F.conv1d(
        hidden_states.to(weight.dtype),
        weight=weight.unsqueeze(1),
        bias=bias,
        padding=padding,
        groups=hidden_size,
    )[:, :, :seq_len]
    if activation is not None:
        out = ACT2FN[activation](out)
    return out.to(hidden_states.dtype)


@use_kernelized_func([causal_conv1d_update, causal_conv1d_fn])
class InklingShortConvolution(nn.Module):
    def __init__(self, hidden_size: int, conv_kernel_size: int, layer_idx: int, conv_idx: int):
        super().__init__()
        self.layer_idx = layer_idx
        self.conv_idx = conv_idx
        self.conv_kernel_size = conv_kernel_size

        self.conv1d = nn.Conv1d(
            in_channels=hidden_size,
            out_channels=hidden_size,
            kernel_size=conv_kernel_size,
            groups=hidden_size,
            padding=conv_kernel_size - 1,
            bias=False,
        )

    @force_accelerate_hooks("conv1d")
    def forward(
        self,
        hidden_states: torch.Tensor,
        past_key_values: Cache | None = None,
        conv_mask: torch.Tensor | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ):
        # Keep the computation in fp32
        input_dtype = hidden_states.dtype
        hidden_states = hidden_states.float()

        residual = hidden_states
        hidden_states = apply_mask_to_padding_states(hidden_states, conv_mask)
        seq_len = hidden_states.shape[1]
        hidden_states = hidden_states.transpose(1, 2)

        use_precomputed_states = past_key_values is not None and past_key_values.has_previous_state(
            self.layer_idx, self.conv_idx
        )

        if use_precomputed_states and seq_len == 1 and not past_key_values.layers[self.layer_idx].record_past:
            conv_state = past_key_values.layers[self.layer_idx].conv_states[self.conv_idx]
            # Single-token cached decode: the fused per-step kernel updates the conv state in-place.
            hidden_states = causal_conv1d_update(
                hidden_states, conv_state, self.conv1d.weight.squeeze(1), self.conv1d.bias
            )
        else:
            if past_key_values is not None:
                hidden_states = past_key_values.update_conv_state(
                    hidden_states, self.layer_idx, state_idx=self.conv_idx, conv_kernel_size=self.conv_kernel_size
                )

            hidden_states = causal_conv1d_fn(
                hidden_states, self.conv1d.weight.squeeze(1), self.conv1d.bias, seq_idx=kwargs.get("seq_idx")
            )

            # Drop the additional previous states
            if use_precomputed_states:
                hidden_states = hidden_states[:, :, -seq_len:]

        hidden_states = hidden_states.transpose(1, 2)
        hidden_states = (hidden_states + residual).to(dtype=input_dtype)
        return hidden_states


class InklingDecoderLayer(GradientCheckpointingLayer):
    def __init__(self, config: InklingTextConfig, layer_idx: int):
        super().__init__()
        self.hidden_size = config.hidden_size
        self.self_attn = InklingAttention(config, layer_idx)

        if config.mlp_layer_types[layer_idx] == "sparse":
            self.mlp = InklingMoE(config)
        else:
            self.mlp = InklingMLP(config)

        self.input_layernorm = InklingRMSNorm(config.hidden_size, config.rms_norm_eps)
        self.post_attention_layernorm = InklingRMSNorm(config.hidden_size, config.rms_norm_eps)
        self.layer_type = config.layer_types[layer_idx]
        self.attn_sconv = InklingShortConvolution(
            config.hidden_size, config.conv_kernel_size, layer_idx=layer_idx, conv_idx=2
        )
        self.mlp_sconv = InklingShortConvolution(
            config.hidden_size, config.conv_kernel_size, layer_idx=layer_idx, conv_idx=3
        )

    def forward(
        self,
        hidden_states: torch.Tensor,
        attention_mask: torch.Tensor | None = None,
        conv_mask: torch.Tensor | None = None,
        past_key_values: Cache | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> torch.Tensor:
        residual = hidden_states
        hidden_states = self.input_layernorm(hidden_states)
        hidden_states, _ = self.self_attn(
            hidden_states=hidden_states,
            attention_mask=attention_mask,
            conv_mask=conv_mask,
            past_key_values=past_key_values,
            **kwargs,
        )
        hidden_states = self.attn_sconv(hidden_states, past_key_values=past_key_values, conv_mask=conv_mask)
        hidden_states = residual + hidden_states

        residual = hidden_states
        hidden_states = self.post_attention_layernorm(hidden_states)
        hidden_states = self.mlp(hidden_states)
        hidden_states = self.mlp_sconv(hidden_states, past_key_values=past_key_values, conv_mask=conv_mask)
        hidden_states = residual + hidden_states
        return hidden_states


@auto_docstring
class InklingPreTrainedModel(PreTrainedModel):
    config_class = InklingConfig
    base_model_prefix = "model"
    supports_gradient_checkpointing = True
    _no_split_modules = ["InklingDecoderLayer"]
    _skip_keys_device_placement = ["past_key_values"]
    # The relative position bias flows through the attention interface as a `position_bias` (duh)
    # kwarg that only the eager path consumes; other backends need a score_mod/kernel
    _supports_flash_attn = False
    _supports_sdpa = True
    _supports_flex_attn = True
    _can_compile_fullgraph = False
    _supports_attention_backend = False
    _keys_to_ignore_on_load_unexpected = [r"model\.mtp\..*"]
    _keep_in_fp32_modules_strict = ["attn_sconv", "mlp_sconv", "k_sconv", "v_sconv"]
    _can_record_outputs = {
        "hidden_states": InklingDecoderLayer,
        "attentions": InklingAttention,
    }

    @torch.no_grad()
    def _init_weights(self, module):
        super()._init_weights(module)
        std = self.config.get_text_config().initializer_range
        if isinstance(module, InklingRelativeLogits):
            init.normal_(module.proj, mean=0.0, std=std)
        elif isinstance(module, InklingMLP):
            init.ones_(module.global_scale)
        elif isinstance(module, InklingExperts):
            init.normal_(module.gate_up_proj, mean=0.0, std=std)
            init.normal_(module.down_proj, mean=0.0, std=std)
        elif isinstance(module, InklingTopkRouter):
            init.normal_(module.weight, mean=0.0, std=std)
            init.ones_(module.global_scale)
            init.zeros_(module.e_score_correction_bias)
        elif isinstance(module, InklingSharedExperts):
            init.normal_(module.gate_proj, mean=0.0, std=std)
            init.normal_(module.up_proj, mean=0.0, std=std)
            init.normal_(module.down_proj, mean=0.0, std=std)
        elif isinstance(module, InklingAudioModelEmbeddings):
            # `_init_weights` runs with `self` being either the top model (`InklingConfig`) or the audio
            # sub-model (`InklingAudioConfig`), so resolve the audio config from whichever we have.
            audio_config = getattr(self.config, "audio_config", self.config)
            init.copy_(
                module.audio_tokens_offsets,
                torch.arange(audio_config.n_mel_bins) * audio_config.mel_vocab_size,
            )


@auto_docstring
class InklingTextModel(InklingPreTrainedModel):
    config: InklingTextConfig

    def __init__(self, config: InklingTextConfig):
        super().__init__(config)
        self.padding_idx = config.pad_token_id
        self.vocab_size = config.vocab_size

        self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
        self.layers = nn.ModuleList(
            [InklingDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
        )
        self.norm = InklingRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.embed_norm = InklingRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
        self.gradient_checkpointing = False

        # Initialize weights and apply final processing
        self.post_init()

    @merge_with_config_defaults
    @capture_outputs
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        past_key_values: Cache | None = None,
        inputs_embeds: torch.FloatTensor | None = None,
        use_cache: bool | None = None,
        **kwargs: Unpack[TransformersKwargs],
    ) -> BaseModelOutputWithPast:
        if (input_ids is None) ^ (inputs_embeds is not None):
            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")

        if inputs_embeds is None:
            inputs_embeds = self.embed_norm(self.embed_tokens(input_ids))

        if use_cache and past_key_values is None:
            past_key_values = DynamicCache(config=self.config)

        if position_ids is None:
            past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
            position_ids = torch.arange(inputs_embeds.shape[1], device=inputs_embeds.device) + past_seen_tokens
            position_ids = position_ids.unsqueeze(0)

        # It may already have been prepared by e.g. `generate`
        if not isinstance(causal_mask_mapping := attention_mask, dict):
            mask_kwargs = {
                "config": self.config,
                "inputs_embeds": inputs_embeds,
                "attention_mask": attention_mask,
                "past_key_values": past_key_values,
                "position_ids": position_ids,
            }
            causal_mask_mapping = {
                "full_attention": create_causal_mask(**mask_kwargs),
                "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
                "linear_attention": create_recurrent_attention_mask(**mask_kwargs),
            }

        hidden_states = inputs_embeds
        for i, decoder_layer in enumerate(self.layers):
            attention_type = "full_attention" if self.config.layer_types[i] == "hybrid" else "sliding_attention"
            hidden_states = decoder_layer(
                hidden_states,
                attention_mask=causal_mask_mapping[attention_type],
                conv_mask=causal_mask_mapping["linear_attention"],
                past_key_values=past_key_values,
                **kwargs,
            )

        hidden_states = self.norm(hidden_states)
        return BaseModelOutputWithPast(
            last_hidden_state=hidden_states,
            past_key_values=past_key_values,
        )


@auto_docstring
class InklingForCausalLM(InklingPreTrainedModel, GenerationMixin):
    # `embed` and `unembed` are separate tensors in the checkpoints, never tied
    _tied_weights_keys = {}
    _tp_plan = {"lm_head": "rowwise_split_input"}
    _pp_plan = {"lm_head": (["hidden_states"], ["logits"])}
    config: InklingTextConfig

    def __init__(self, config: InklingTextConfig):
        super().__init__(config)
        self.model = InklingTextModel(config)
        self.vocab_size = config.vocab_size
        self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)

        # Initialize weights and apply final processing
        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,
        past_key_values: Cache | None = None,
        inputs_embeds: torch.FloatTensor | None = None,
        labels: torch.LongTensor | None = None,
        use_cache: bool | None = None,
        logits_to_keep: int | torch.Tensor = 0,
        **kwargs: Unpack[TransformersKwargs],
    ) -> InklingCausalLMOutputWithPast:
        r"""
        Example:

        ```python
        >>> from transformers import AutoTokenizer, InklingForCausalLM

        >>> model = InklingForCausalLM.from_pretrained("google/gemma-2-9b")
        >>> tokenizer = AutoTokenizer.from_pretrained("google/gemma-2-9b")

        >>> prompt = "What is your favorite condiment?"
        >>> inputs = tokenizer(prompt, return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
        >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "What is your favorite condiment?"
        ```"""
        outputs = self.model(
            input_ids=input_ids,
            attention_mask=attention_mask,
            position_ids=position_ids,
            past_key_values=past_key_values,
            inputs_embeds=inputs_embeds,
            use_cache=use_cache,
            **kwargs,
        )

        hidden_states = outputs.last_hidden_state / self.config.logits_mup_width_multiplier
        # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
        slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
        logits = self.lm_head(hidden_states[:, slice_indices, :])
        unpadded_vocab_size = self.config.unpadded_vocab_size
        if unpadded_vocab_size is not None and unpadded_vocab_size < logits.shape[-1]:
            logits = logits[..., :unpadded_vocab_size]

        loss = None
        if labels is not None:
            loss = self.loss_function(logits=logits, labels=labels, vocab_size=logits.shape[-1], **kwargs)

        return InklingCausalLMOutputWithPast(
            loss=loss,
            logits=logits,
            past_key_values=outputs.past_key_values,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
        )


class InklingAudioModelEmbeddings(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.embed_audio_tokens = nn.Embedding((config.num_codebooks * config.codebook_size), config.hidden_size)
        self.register_buffer(
            "audio_tokens_offsets", torch.arange(config.num_codebooks) * config.codebook_size, persistent=False
        )

    def forward(self, input_ids):
        inputs_embeds = self.embed_audio_tokens(input_ids + self.audio_tokens_offsets)
        inputs_embeds = inputs_embeds.sum(dim=-2)
        return inputs_embeds


class InklingAudioModel(InklingPreTrainedModel):
    def __init__(self, config: InklingAudioConfig):
        super().__init__(config)
        self.embed_audio_tokens = InklingAudioModelEmbeddings(config)
        self.norm = InklingRMSNorm(config.text_hidden_size, eps=1e-6)

    def forward(self, audio_input_ids: torch.Tensor) -> torch.Tensor:
        hidden_states = self.embed_audio_tokens(audio_input_ids)
        hidden_states = self.norm(hidden_states)
        return BaseModelOutputWithPooling(
            last_hidden_state=hidden_states,
            pooler_output=hidden_states,
        )


class InklingVisionEncoderLayer(nn.Module):
    def __init__(self, input_dim: int, output_dim: int, t_fold: int, hw_fold: int, add_norm: bool):
        super().__init__()
        self.projection = nn.Linear(input_dim, output_dim, bias=False)
        if add_norm:
            self.layer_norm = InklingRMSNorm(output_dim)
        self.hw_fold = hw_fold
        self.t_fold = t_fold
        self.add_norm = add_norm

    def fold_timespace_to_depth(self, hidden_states: torch.Tensor) -> torch.Tensor:
        """
        Convert a tensor of shape (B, T, H, W, C) to a tensor of shape (B, T // t, H // hw, W //  hw, C * (t * hw**2))
        """
        B, T, H, W, C = hidden_states.shape

        t_new = T // self.t_fold
        h_new = H // self.hw_fold
        w_new = W // self.hw_fold

        hidden_states = hidden_states.reshape(B, t_new, self.t_fold, h_new, self.hw_fold, w_new, self.hw_fold, C)

        hidden_states = hidden_states.permute(0, 1, 3, 5, 2, 4, 6, 7)
        hidden_states = hidden_states.reshape(B, t_new, h_new, w_new, self.t_fold * self.hw_fold * self.hw_fold * C)
        return hidden_states

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        if self.hw_fold > 1 or self.t_fold > 1:
            hidden_states = self.fold_timespace_to_depth(hidden_states)

        hidden_states = self.projection(hidden_states)
        if self.add_norm:
            hidden_states = self.layer_norm(hidden_states)
            hidden_states = F.gelu(hidden_states)
        return hidden_states


def prime_factors(number: int) -> list[int]:
    factors = []

    while number % 2 == 0:
        factors.append(2)
        number //= 2

    for p in range(3, math.isqrt(number) + 1, 2):
        while number % p == 0:
            factors.append(p)
            number //= p

    if number > 1:
        factors.append(number)
    return factors


def plan_out_scales(
    temporal_patch_size: int, patch_size: int, n_layers: int, n_channels: int, device="cpu"
) -> torch.LongTensor:
    """
    Plan out the dimensions for each layer in the HMLP encoder.

    This function determines the progression of dimensions (temporal, height, width, channels)
    for a multi-layer perceptual model that processes image/video patches. It follows these
    principles:
    1. Start with small dimensions and increase to full size
    2. Expand spatial dimensions (height/width) first, then temporal
    3. Increase channel count to avoid information bottlenecks
    4. Round channel dimensions to multiples of 64 for hardware efficiency

    The function computes optimal assignments of scale configurations to layers using either:
    - For n_layers >= len(scales): Individual best matching scales for each layer (allowing duplicates)
    - For n_layers < len(scales): Global optimal assignment via linear_sum_assignment

    The first and last scales are always fixed to ensure the proper input and output dimensions.

    Args:
        temporal_patch_size: Temporal dimension of input patches
        patch_size: Spatial dimension (height/width) of input patches
        n_layers: Number of layers in the encoder
        n_channels: Number of input channels (default: 3 for RGB)

    Returns:
        torch.LongTensor of shape `(n_layers + 1, 4)` where the last dim holds values for (t, h, w, c) grids.
    """
    h = torch.cumprod(torch.tensor(prime_factors(patch_size)[::-1], device=device), dim=0)
    t = torch.cumprod(torch.tensor(prime_factors(temporal_patch_size)[::-1], device=device), dim=0)

    h_ch = torch.ceil(h**2 * n_channels / 64).int() * 64
    t_ch = torch.ceil(h[-1] ** 2 * n_channels * t).int() * 64

    base = torch.tensor([[1, 1, 1, n_channels]], device=device)
    spatial = torch.stack([torch.ones_like(h), h, h, h_ch], dim=1)
    temporal = torch.stack([t, torch.full_like(t, h[-1]), torch.full_like(t, h[-1]), t_ch], dim=1)
    scales = torch.cat([base, spatial, temporal], dim=0)

    size_reduction = torch.prod(scales[:, :-1], dim=1).float()

    total_elements = patch_size * patch_size * temporal_patch_size * n_channels
    log_ideal_scales = torch.linspace(
        0, torch.log(torch.tensor(total_elements, device=device)), n_layers + 1, device=device
    )
    cost_matrix = torch.abs(log_ideal_scales.unsqueeze(1) - torch.log(size_reduction).unsqueeze(0))

    if n_layers >= scales.shape[0]:
        idxs = torch.argmin(cost_matrix, dim=1)
    else:
        from scipy.optimize import linear_sum_assignment

        _, idxs_np = linear_sum_assignment(cost_matrix.cpu().numpy())
        idxs = torch.tensor(idxs_np, device=device)
        # idxs = torch.softmax(-cost_matrix * 10, dim=1).argmax(dim=1)

    idxs[0] = 0
    idxs[-1] = scales.shape[0] - 1
    return scales[idxs]


class InklingVisionModel(InklingPreTrainedModel):
    def __init__(self, config: InklingVisionConfig):
        super().__init__(config)
        self.scales = plan_out_scales(
            config.temporal_patch_size,
            config.patch_size,
            config.num_hidden_layers,
            config.num_channels,
        )

        # num_hidden_layers - 1 to encoder and the last to proj to text hidden dim
        self.encoder_layers = nn.ModuleList()
        for i, (start_scale, end_scale) in enumerate(zip(self.scales[:-1], self.scales[1:])):
            shuffle_mult = (
                (end_scale[0] // start_scale[0]) * (end_scale[1] // start_scale[1]) * (end_scale[2] // start_scale[2])
            )
            output_dim = config.text_hidden_size if i == config.num_hidden_layers - 1 else end_scale[3]
            hw_fold = end_scale[1] // start_scale[1]
            t_fold = end_scale[0] // start_scale[0]
            self.encoder_layers.append(
                InklingVisionEncoderLayer(
                    input_dim=start_scale[3] * shuffle_mult,
                    output_dim=output_dim,
                    hw_fold=hw_fold,
                    t_fold=t_fold,
                    add_norm=i != config.num_hidden_layers - 1,
                )
            )

        self.final_norm = InklingRMSNorm(config.text_hidden_size)
        self.post_init()

    def forward(self, pixel_values: torch.Tensor, **kwargs: Unpack[TransformersKwargs]) -> torch.Tensor:
        num_patches = pixel_values.shape[0]
        hidden_states = pixel_values
        for layer in self.encoder_layers:
            hidden_states = layer(hidden_states=hidden_states)

        hidden_states = self.final_norm(hidden_states)
        hidden_states = hidden_states.reshape(num_patches, -1)
        return BaseModelOutputWithPooling(
            last_hidden_state=hidden_states,
            pooler_output=hidden_states,
        )


@auto_docstring(
    custom_intro="""
    The Base Inkling model which consists of a vision backbone and a language model without language modeling head.,
    """
)
class InklingModel(InklingPreTrainedModel):
    # we are filtering the logits/labels so we shouldn't divide the loss based on num_items_in_batch
    accepts_loss_kwargs = False

    def __init__(self, config: InklingConfig):
        super().__init__(config)
        self.vocab_size = config.text_config.vocab_size
        self.language_model = InklingTextModel(config.text_config)
        self.audio_tower = InklingAudioModel(config.audio_config)
        self.vision_tower = InklingVisionModel(config.vision_config)
        self.post_init()

    @can_return_tuple
    @auto_docstring(custom_intro="Projects the last hidden state from the vision model into language model space.")
    def get_image_features(
        self, pixel_values: torch.FloatTensor, **kwargs: Unpack[TransformersKwargs]
    ) -> tuple | BaseModelOutputWithPooling:
        return self.vision_tower(pixel_values=pixel_values, **kwargs)

    @can_return_tuple
    @auto_docstring(custom_intro="Projects discretized dMel bin tokens into the language model space.")
    def get_audio_features(
        self,
        audio_input_ids: torch.LongTensor,
        audio_input_ids_mask: torch.Tensor | None = None,
    ) -> tuple | BaseModelOutputWithPooling:
        r"""
        audio_input_ids (`torch.LongTensor` of shape `(num_audios, max_num_frames, n_mel_bins)`):
            Batch of (padded) dMel bin tokens produced by [`InklingProcessor`].
        audio_input_ids_mask (`torch.Tensor` of shape `(num_audios, max_num_frames)`, *optional*):
            Mask marking valid (non-padding) frames. When provided, only valid frames are encoded so that the
            number of returned audio embeddings matches the number of audio placeholder tokens.
        """
        if audio_input_ids_mask is not None:
            audio_input_ids = audio_input_ids[audio_input_ids_mask.bool()]
        else:
            audio_input_ids = audio_input_ids.reshape(-1, audio_input_ids.shape[-1])
        return self.audio_tower(audio_input_ids)

    def get_placeholder_mask(
        self,
        input_ids: torch.LongTensor,
        inputs_embeds: torch.FloatTensor,
        features: torch.FloatTensor,
        token_id: int,
    ):
        """
        Obtains a multimodal placeholder mask from `input_ids` or `inputs_embeds` for the given `token_id`, and checks
        that the placeholder token count matches the length of `features`. If the lengths differ, an error is raised.
        """
        if input_ids is None:
            special_mask = inputs_embeds == self.get_input_embeddings()(
                torch.tensor(token_id, dtype=torch.long, device=inputs_embeds.device)
            )
            special_mask = special_mask.all(-1)
        else:
            special_mask = input_ids == token_id

        n_tokens = special_mask.sum()
        special_mask = special_mask.unsqueeze(-1).expand_as(inputs_embeds).to(inputs_embeds.device)
        torch_compilable_check(
            inputs_embeds[special_mask].numel() == features.numel(),
            f"Multimodal features and placeholder tokens do not match, tokens: {n_tokens}, features: {features.shape[0]}",
        )
        return special_mask

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        pixel_values: torch.FloatTensor | None = None,
        audio_input_ids: torch.LongTensor | None = None,
        audio_input_ids_mask: torch.Tensor | None = None,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        past_key_values: Cache | None = None,
        token_type_ids: torch.LongTensor | None = None,
        inputs_embeds: torch.FloatTensor | None = None,
        labels: torch.LongTensor | None = None,
        use_cache: bool | None = None,
        **lm_kwargs: Unpack[TransformersKwargs],
    ) -> tuple | InklingModelOutputWithPast:
        r"""
        audio_input_ids (`torch.LongTensor` of shape `(num_audios, max_num_frames, n_mel_bins)`, *optional*):
            Batch of (padded) discretized dMel bin tokens produced by [`InklingProcessor`].
        audio_input_ids_mask (`torch.Tensor` of shape `(num_audios, max_num_frames)`, *optional*):
            Mask marking valid (non-padding) audio frames in `audio_input_ids`.
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.text_config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.text_config.vocab_size]`.

        Example:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoProcessor, InklingForConditionalGeneration

        >>> model = InklingForConditionalGeneration.from_pretrained("google/inkling2-3b-mix-224")
        >>> processor = AutoProcessor.from_pretrained("google/inkling2-3b-mix-224")

        >>> prompt = "Where is the cat standing?"
        >>> url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/pipeline-cat-chonk.jpeg"
        >>> with httpx.stream("GET", url) as response:
        ...     image = Image.open(BytesIO(response.read()))

        >>> inputs = processor(images=image, text=prompt,  return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(**inputs,)
        >>> processor.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "Where is the cat standing?\nsnow"
        ```"""
        if (input_ids is None) ^ (inputs_embeds is not None):
            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")

        if inputs_embeds is None:
            inputs_embeds = self.language_model.embed_norm(self.get_input_embeddings()(input_ids))

        # Merge text and images
        if pixel_values is not None:
            image_features = self.get_image_features(pixel_values).pooler_output
            image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
            special_image_mask = self.get_placeholder_mask(
                input_ids, inputs_embeds, image_features, self.config.image_token_id
            )
            inputs_embeds = inputs_embeds.masked_scatter(special_image_mask, image_features)

        # Merge text and audio
        audio_features = None
        if audio_input_ids is not None:
            audio_features = self.get_audio_features(audio_input_ids, audio_input_ids_mask).last_hidden_state
            audio_features = audio_features.to(inputs_embeds.device, inputs_embeds.dtype)
            special_audio_mask = self.get_placeholder_mask(
                input_ids, inputs_embeds, audio_features, self.config.audio_token_id
            )
            inputs_embeds = inputs_embeds.masked_scatter(special_audio_mask, audio_features)

        # It may already have been prepared by e.g. `generate`
        if not isinstance(causal_mask_mapping := attention_mask, dict):
            mask_kwargs = {
                "config": self.config.get_text_config(),
                "inputs_embeds": inputs_embeds,
                "attention_mask": attention_mask,
                "past_key_values": past_key_values,
                "position_ids": position_ids,
            }

            causal_mask_mapping = {
                "full_attention": create_causal_mask(**mask_kwargs),
                "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
                "linear_attention": create_recurrent_attention_mask(**mask_kwargs),
            }

        outputs = self.language_model(
            attention_mask=causal_mask_mapping,
            position_ids=position_ids,
            past_key_values=past_key_values,
            inputs_embeds=inputs_embeds,
            use_cache=use_cache,
            **lm_kwargs,
        )

        return InklingModelOutputWithPast(
            last_hidden_state=outputs.last_hidden_state,
            past_key_values=outputs.past_key_values,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
            image_hidden_states=image_features if pixel_values is not None else None,
        )


@auto_docstring(
    custom_intro="""
    The Base Inkling model which consists of a vision backbone and a language model without language modeling head.,
    """
)
class InklingForConditionalGeneration(InklingPreTrainedModel, GenerationMixin):
    # `embed` and `unembed` are separate tensors in the checkpoints, never tied
    _tied_weights_keys = {}
    _tp_plan = {"lm_head": "rowwise_split_input"}
    # we are filtering the logits/labels so we shouldn't divide the loss based on num_items_in_batch
    # Fix: https://github.com/huggingface/transformers/issues/40564
    accepts_loss_kwargs = False

    def __init__(self, config: InklingConfig):
        super().__init__(config)
        self.model = InklingModel(config)
        # checkpoints store `unembed` padded to vocab_size; logits are sliced to
        # unpadded_vocab_size in forward, like sglang
        self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
        self.post_init()

    @auto_docstring
    def get_image_features(self, pixel_values: torch.FloatTensor, **kwargs: Unpack[TransformersKwargs]):
        return self.model.get_image_features(pixel_values, **kwargs)

    @can_return_tuple
    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        pixel_values: torch.FloatTensor | None = None,
        attention_mask: torch.Tensor | None = None,
        position_ids: torch.LongTensor | None = None,
        past_key_values: Cache | None = None,
        audio_input_ids: torch.LongTensor | None = None,
        audio_input_ids_mask: torch.Tensor | None = None,
        inputs_embeds: torch.FloatTensor | None = None,
        labels: torch.LongTensor | None = None,
        use_cache: bool | None = None,
        logits_to_keep: int | torch.Tensor = 0,
        **kwargs: Unpack[TransformersKwargs],
    ) -> tuple | InklingCausalLMOutputWithPast:
        r"""
        audio_input_ids (`torch.LongTensor` of shape `(num_audios, max_num_frames, n_mel_bins)`, *optional*):
            Batch of (padded) discretized dMel bin tokens produced by [`InklingProcessor`].
        audio_input_ids_mask (`torch.Tensor` of shape `(num_audios, max_num_frames)`, *optional*):
            Mask marking valid (non-padding) audio frames in `audio_input_ids`.
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.text_config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.text_config.vocab_size]`.

        Example:

        ```python
        >>> from PIL import Image
        >>> import httpx
        >>> from io import BytesIO
        >>> from transformers import AutoProcessor, InklingForConditionalGeneration

        >>> model = InklingForConditionalGeneration.from_pretrained("google/gemma-3-4b-it")
        >>> processor = AutoProcessor.from_pretrained("google/gemma-3-4b-it")

        >>> messages = [
        ...     {
        ...         "role": "system",
        ...         "content": [
        ...             {"type": "text", "text": "You are a helpful assistant."}
        ...         ]
        ...     },
        ...     {
        ...         "role": "user", "content": [
        ...             {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/pipeline-cat-chonk.jpeg"},
        ...             {"type": "text", "text": "Where is the cat standing?"},
        ...         ]
        ...     },
        ... ]

        >>> inputs = processor.apply_chat_template(
        ...     messages,
        ...     tokenize=True,
        ...     return_dict=True,
        ...     return_tensors="pt",
        ...     add_generation_prompt=True
        ... )
        >>> # Generate
        >>> generate_ids = model.generate(**inputs)
        >>> processor.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "user\nYou are a helpful assistant.\n\n\n\n\n\nWhere is the cat standing?\nmodel\nBased on the image, the cat is standing in a snowy area, likely outdoors. It appears to"
        ```
        """
        outputs = self.model(
            input_ids=input_ids,
            pixel_values=pixel_values,
            audio_input_ids=audio_input_ids,
            audio_input_ids_mask=audio_input_ids_mask,
            attention_mask=attention_mask,
            position_ids=position_ids,
            past_key_values=past_key_values,
            inputs_embeds=inputs_embeds,
            use_cache=use_cache,
            labels=labels,
            **kwargs,
        )

        hidden_states = outputs[0] / self.config.text_config.logits_mup_width_multiplier
        # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
        slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
        logits = self.lm_head(hidden_states[:, slice_indices, :])
        unpadded_vocab_size = self.config.text_config.unpadded_vocab_size
        if unpadded_vocab_size is not None and unpadded_vocab_size < logits.shape[-1]:
            logits = logits[..., :unpadded_vocab_size]

        loss = None
        if labels is not None:
            loss = self.loss_function(logits=logits, labels=labels, vocab_size=logits.shape[-1], **kwargs)

        return InklingCausalLMOutputWithPast(
            loss=loss,
            logits=logits,
            past_key_values=outputs.past_key_values,
            hidden_states=outputs.hidden_states,
            attentions=outputs.attentions,
            image_hidden_states=outputs.image_hidden_states,
        )

    def prepare_inputs_for_generation(
        self,
        input_ids,
        past_key_values=None,
        inputs_embeds=None,
        position_ids=None,
        pixel_values=None,
        attention_mask=None,
        audio_input_ids=None,
        audio_input_ids_mask=None,
        use_cache=True,
        logits_to_keep=None,
        labels=None,
        is_first_iteration=False,
        **kwargs,
    ):
        # Overwritten -- custom `pixel_values/audio_input_ids` handling
        model_inputs = super().prepare_inputs_for_generation(
            input_ids,
            past_key_values=past_key_values,
            inputs_embeds=inputs_embeds,
            attention_mask=attention_mask,
            position_ids=position_ids,
            use_cache=use_cache,
            logits_to_keep=logits_to_keep,
            is_first_iteration=is_first_iteration,
            **kwargs,
        )

        if is_first_iteration or not use_cache:
            model_inputs["pixel_values"] = pixel_values
            model_inputs["audio_input_ids"] = audio_input_ids
            model_inputs["audio_input_ids_mask"] = audio_input_ids_mask

        return model_inputs


__all__ = [
    "InklingPreTrainedModel",
    "InklingTextModel",
    "InklingForCausalLM",
    "InklingAudioModel",
    "InklingVisionModel",
    "InklingModel",
    "InklingForConditionalGeneration",
]
