# Copyright 2024 state-spaces/mamba org and HuggingFace Inc. team.
#
# 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.
"""PyTorch MAMBA model."""

import math
from dataclasses import dataclass

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

from ... import initialization as init
from ...activations import ACT2FN
from ...cache_utils import Cache, DynamicCache
from ...generation import GenerationMixin
from ...integrations import use_kernel_func_from_hub_with_fallback, use_kernelized_func
from ...integrations.accelerate import force_accelerate_hooks
from ...modeling_layers import GradientCheckpointingLayer
from ...modeling_utils import PreTrainedModel
from ...utils import (
    ModelOutput,
    auto_docstring,
    logging,
)
from ...utils.import_utils import (
    is_mambapy_available,
    is_torch_greater_or_equal,
    is_torchdynamo_compiling,
    is_torchdynamo_exporting,
)
from .configuration_mamba import MambaConfig


logger = logging.get_logger(__name__)


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:
        dtype = hidden_states.dtype
        hidden_states = (hidden_states * attention_mask[:, :, None]).to(dtype)

    return hidden_states


@use_kernel_func_from_hub_with_fallback("causal_conv1d_update", "causal_conv1d")
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_with_fallback("causal_conv1d_fn", "causal_conv1d")
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_kernel_func_from_hub_with_fallback("mamba_inner_fn", "mamba_ssm")
def mamba_inner_fn(
    xz: torch.Tensor,
    conv1d_weight: torch.Tensor,
    conv1d_bias: torch.Tensor | None,
    x_proj_weight: torch.Tensor,
    delta_proj_weight: torch.Tensor,
    out_proj_weight: torch.Tensor,
    out_proj_bias: torch.Tensor | None,
    A: torch.Tensor,
    B: torch.Tensor | None = None,
    C: torch.Tensor | None = None,
    D: torch.Tensor | None = None,
    delta_bias: torch.Tensor | None = None,
    delta_softplus: bool = True,
    b_rms_weight: torch.Tensor | None = None,
    c_rms_weight: torch.Tensor | None = None,
    dt_rms_weight: torch.Tensor | None = None,
    b_c_dt_rms_eps: float = 1e-6,
    **kwargs,
):
    return None


@use_kernel_func_from_hub_with_fallback("selective_state_update", "mamba_ssm")
def mamba_selective_state_update(
    state: torch.Tensor,
    hidden_states: torch.Tensor,
    dt: torch.Tensor,
    A: torch.Tensor,
    B: torch.Tensor,
    C: torch.Tensor,
    D: torch.Tensor | None = None,
    dt_bias: torch.Tensor | None = None,
    dt_softplus: bool = False,
    z: torch.Tensor | None = None,
    **kwargs,
):
    input_dtype = hidden_states.dtype

    if dt_bias is not None:
        dt = dt + dt_bias.to(dt.dtype)
    if dt_softplus:
        dt = F.softplus(dt)

    # Discretize A
    dA = torch.exp(dt.float()[..., None] * A.float()).to(device=state.device)

    # Discretize B
    dB = dt.float()[..., None] * B.float()[:, None, :]
    # Discretize x into dB
    dBx = dB * hidden_states.float()[..., None]

    # State calculation
    ssm_state = state.float() * dA + dBx
    state.copy_(ssm_state.to(state.dtype))

    # Subsequent output
    out = torch.matmul(ssm_state.to(C.dtype), C.unsqueeze(-1)).squeeze(-1)

    # D skip connection
    if D is not None:
        out = out + hidden_states * D

    if z is not None:
        out = out * F.silu(z)

    return out.to(input_dtype)


@use_kernel_func_from_hub_with_fallback("selective_scan_fn", "mamba_ssm")
def mamba_selective_scan(
    hidden_states: torch.Tensor,
    dt: torch.Tensor,
    A: torch.Tensor,
    B: torch.Tensor,
    C: torch.Tensor,
    D: torch.Tensor | None = None,
    z: torch.Tensor | None = None,
    delta_bias: torch.Tensor | None = None,
    delta_softplus: bool = False,
    return_last_state: bool = False,
    use_mambapy: bool = False,
    use_associative_scan: bool = False,
    **kwargs,
):
    # Torch only alternatives to the recurrent path
    if is_torch_greater_or_equal("2.9.0"):
        from torch._higher_order_ops.associative_scan import associative_scan
    else:
        associative_scan = None

    if is_mambapy_available():
        from mambapy.pscan import pscan
    else:
        pscan = None

    batch_size, intermediate_size, seq_len = hidden_states.shape
    input_dtype = hidden_states.dtype

    if delta_bias is not None:
        dt = dt + delta_bias.to(dt.dtype)[..., None]
    if delta_softplus:
        dt = F.softplus(dt)

    # We need to transpose on the basis of the original kernel layout
    B = B.transpose(1, 2)
    C = C.transpose(1, 2)

    # Discretize A and B for the entire sequence
    discrete_A = torch.exp(A[None, :, None, :] * dt[:, :, :, None])
    discrete_B = dt[:, :, :, None] * B[:, None, :, :].float()
    deltaB_u = discrete_B * hidden_states[:, :, :, None].float()

    if use_mambapy and pscan is not None:
        all_states = pscan(discrete_A.transpose(1, 2), deltaB_u.transpose(1, 2))

        scan_output = (all_states @ C.unsqueeze(-1)).squeeze(3).transpose(1, 2)
        ssm_state = all_states[:, -1]

    elif (
        use_associative_scan
        and associative_scan is not None
        # There is no onnx translation for this op so we rely on the normal sequential path then
        and (is_torchdynamo_compiling() and not is_torchdynamo_exporting())
    ):

        def combine_fn(left, right):
            a_left, b_left = left
            a_right, b_right = right
            return a_left * a_right, a_right * b_left + b_right

        combine_mode = "pointwise" if discrete_A.device.type in ("cuda", "xpu") else "generic"
        _, all_states = associative_scan(
            combine_fn,
            (discrete_A, deltaB_u),
            dim=2,
            combine_mode=combine_mode,
        )

        scan_output = torch.matmul(all_states.transpose(1, 2).to(input_dtype), C.unsqueeze(-1))
        scan_output = scan_output.squeeze(-1).transpose(1, 2)
        ssm_state = all_states[:, :, -1]

    # Recurrent iteration
    else:
        # "Initial hidden state" is not supported by the kernel path, so use
        # the same zero initialization as the kernel
        ssm_state = torch.zeros(
            batch_size,
            intermediate_size,
            A.shape[-1],
            dtype=input_dtype,
            device=hidden_states.device,
        )

        scan_outputs = []
        for index in range(seq_len):
            # State calculation
            ssm_state = discrete_A[:, :, index] * ssm_state + deltaB_u[:, :, index]

            # Subsequent output
            scan_output = torch.matmul(ssm_state.to(input_dtype), C[:, index, :].unsqueeze(-1))
            scan_outputs.append(scan_output[:, :, 0])
        scan_output = torch.stack(scan_outputs, dim=-1)

    if D is not None:
        scan_output = scan_output + hidden_states * D[None, :, None]

    if z is not None:
        scan_output = scan_output * F.silu(z)

    if return_last_state:
        return scan_output, ssm_state

    return scan_output


@use_kernelized_func(
    [mamba_inner_fn, mamba_selective_scan, mamba_selective_state_update, causal_conv1d_fn, causal_conv1d_update]
)
class MambaMixer(nn.Module):
    """
    Compute ∆, A, B, C, and D the state space parameters and compute the `contextualized_states`.
    A, D are input independent (see Mamba paper [1] Section 3.5.2 "Interpretation of A" for why A isn't selective)
    ∆, B, C are input-dependent (this is a key difference between Mamba and the linear time invariant S4,
    and is why Mamba is called **selective** state spaces)
    """

    def __init__(self, config: MambaConfig, layer_idx: int, initialize_mixer_weights: bool = True):
        super().__init__()
        self.config = config
        self.hidden_size = config.hidden_size
        self.ssm_state_size = config.state_size
        self.conv_kernel_size = config.conv_kernel
        self.intermediate_size = config.intermediate_size
        self.time_step_rank = int(config.time_step_rank)
        self.layer_idx = layer_idx
        self.use_conv_bias = config.use_conv_bias
        self.conv1d = nn.Conv1d(
            in_channels=self.intermediate_size,
            out_channels=self.intermediate_size,
            bias=config.use_conv_bias,
            kernel_size=config.conv_kernel,
            groups=self.intermediate_size,
            padding=config.conv_kernel - 1,
        )

        self.activation = config.hidden_act
        self.act = ACT2FN[config.hidden_act]

        self.use_mambapy = config.use_mambapy
        self.use_associative_scan = config.use_associative_scan

        # projection of the input hidden states
        self.in_proj = nn.Linear(self.hidden_size, self.intermediate_size * 2, bias=config.use_bias)
        # selective projection used to make dt, B and C input dependent
        self.x_proj = nn.Linear(self.intermediate_size, self.time_step_rank + self.ssm_state_size * 2, bias=False)
        # time step projection (discretization)
        self.dt_proj = nn.Linear(self.time_step_rank, self.intermediate_size, bias=True)

        # S4D real initialization. These are not discretized!
        # The core is to load them, compute the discrete states, then write the updated state. Keeps the memory bounded
        self.A_log = nn.Parameter(torch.empty(self.intermediate_size, self.ssm_state_size))
        self.D = nn.Parameter(torch.empty(self.intermediate_size))
        if initialize_mixer_weights and self.dt_proj.weight.device.type != "meta":
            self.init_mamba_weights()
        self.out_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.use_bias)
        self.use_bias = config.use_bias

        self.layer_type = config.layer_types[layer_idx]

    @torch.no_grad()
    def init_mamba_weights(self):
        A = torch.arange(1, self.ssm_state_size + 1, dtype=torch.float32, device=self.A_log.device)[None, :]
        A = A.expand(self.intermediate_size, -1).contiguous()
        init.copy_(self.A_log, torch.log(A))
        init.ones_(self.D)

        dt_init_std = self.config.time_step_rank**-0.5 * self.config.time_step_scale
        if self.config.time_step_init_scheme == "constant":
            init.constant_(self.dt_proj.weight, dt_init_std)
        elif self.config.time_step_init_scheme == "random":
            init.uniform_(self.dt_proj.weight, -dt_init_std, dt_init_std)

        dt = torch.exp(
            torch.rand(self.intermediate_size, device=self.dt_proj.bias.device, dtype=torch.float32)
            * (math.log(self.config.time_step_max) - math.log(self.config.time_step_min))
            + math.log(self.config.time_step_min)
        ).clamp(min=self.config.time_step_floor)
        # Inverse of softplus: https://github.com/pytorch/pytorch/issues/72759
        inv_dt = dt + torch.log(-torch.expm1(-dt))
        init.copy_(self.dt_proj.bias, inv_dt)

    @force_accelerate_hooks(["conv1d", "dt_proj"])
    def forward(
        self,
        hidden_states: torch.Tensor,
        cache_params: Cache | None = None,
        attention_mask: torch.LongTensor | None = None,
        **kwargs,
    ):
        seq_len = hidden_states.shape[1]
        dtype = hidden_states.dtype
        use_precomputed_states = cache_params is not None and cache_params.has_previous_state(self.layer_idx)

        # 1. Gated MLP's linear projection
        hidden_states = apply_mask_to_padding_states(hidden_states, attention_mask)
        projected_states = self.in_proj(hidden_states).transpose(1, 2)

        A = -torch.exp(self.A_log.float())
        if self.training and cache_params is None:
            fused_output = mamba_inner_fn(
                projected_states,
                self.conv1d.weight,
                self.conv1d.bias if self.use_conv_bias else None,
                self.x_proj.weight,
                self.dt_proj.weight,
                self.out_proj.weight,
                self.out_proj.bias.float() if self.use_bias else None,
                A,
                None,  # input-dependent B
                None,  # input-dependent C
                self.D.float(),
                delta_bias=self.dt_proj.bias.float(),
                delta_softplus=True,
            )

            # Only kernels can use this shortcircuit, fallback to normal torch otherwise
            if fused_output is not None:
                return fused_output

        hidden_states_B_C, gate = projected_states.chunk(2, dim=1)

        if use_precomputed_states:
            conv_state = cache_params.layers[self.layer_idx].conv_states[0]
            recurrent_state = cache_params.layers[self.layer_idx].recurrent_states[0]

        # 2. Convolution sequence transformation
        if use_precomputed_states and seq_len == 1 and not cache_params.layers[self.layer_idx].record_past:
            hidden_states_B_C = causal_conv1d_update(
                hidden_states_B_C,
                conv_state,
                self.conv1d.weight.squeeze(1),
                self.conv1d.bias,
                activation=self.activation,
            )
        else:
            if cache_params is not None:
                hidden_states_B_C = cache_params.update_conv_state(
                    hidden_states_B_C,
                    self.layer_idx,
                    conv_kernel_size=self.conv_kernel_size,
                )

            hidden_states_B_C = causal_conv1d_fn(
                hidden_states_B_C,
                self.conv1d.weight.squeeze(1),
                self.conv1d.bias,
                activation=self.activation,
                **kwargs,
            )

            if cache_params is not None:
                hidden_states_B_C = hidden_states_B_C[:, :, -seq_len:]

        # 3. SSM transformation
        hidden_states_B_C = apply_mask_to_padding_states(hidden_states_B_C.transpose(1, 2), attention_mask)
        time_step, B, C = torch.split(
            self.x_proj(hidden_states_B_C),
            [self.time_step_rank, self.ssm_state_size, self.ssm_state_size],
            dim=-1,
        )

        time_step = self.dt_proj.weight @ time_step.transpose(1, 2)
        time_proj_bias = self.dt_proj.bias.float() if self.dt_proj.bias is not None else None

        # Recurrent form
        if use_precomputed_states and seq_len == 1:
            scan_output = mamba_selective_state_update(
                recurrent_state,
                hidden_states_B_C.transpose(1, 2)[..., 0],
                time_step[..., 0],
                A,
                B[:, 0],
                C[:, 0],
                self.D,
                z=gate[..., 0],
                dt_bias=time_proj_bias,
                dt_softplus=True,
            ).unsqueeze(-1)

        # Full sequence form
        else:
            output_final_state = cache_params is not None
            scan_result = mamba_selective_scan(
                hidden_states_B_C.transpose(1, 2),
                time_step,
                A,
                B.transpose(1, 2),
                C.transpose(1, 2),
                D=self.D.float(),
                z=gate,
                delta_bias=time_proj_bias,
                delta_softplus=True,
                return_last_state=output_final_state,
                use_mambapy=self.use_mambapy,
                use_associative_scan=self.use_associative_scan,
            )

            if output_final_state:
                scan_output, final_state = scan_result
                cache_params.update_recurrent_state(final_state, self.layer_idx)
            else:
                scan_output = scan_result

        # 4. Final linear projection
        contextualized_states = self.out_proj(scan_output.transpose(1, 2).to(dtype))
        return contextualized_states


class MambaRMSNorm(nn.Module):
    def __init__(self, hidden_size, eps=1e-6):
        """
        MambaRMSNorm is equivalent to T5LayerNorm and LlamaRMSNorm
        """
        super().__init__()
        self.weight = nn.Parameter(torch.ones(hidden_size))
        self.variance_epsilon = eps

    def forward(self, hidden_states):
        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"{self.weight.shape[0]}, eps={self.variance_epsilon}"


class MambaBlock(GradientCheckpointingLayer):
    def __init__(self, config, layer_idx):
        super().__init__()
        self.config = config
        self.layer_idx = layer_idx
        self.residual_in_fp32 = config.residual_in_fp32
        self.norm = MambaRMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
        self.mixer = MambaMixer(config, layer_idx=layer_idx, initialize_mixer_weights=False)

    def forward(
        self,
        hidden_states,
        cache_params: Cache | None = None,
        attention_mask: torch.LongTensor | None = None,
        **kwargs,
    ):
        residual = hidden_states
        hidden_states = self.norm(hidden_states.to(dtype=self.norm.weight.dtype))
        if self.residual_in_fp32:
            residual = residual.to(torch.float32)

        hidden_states = self.mixer(hidden_states, cache_params=cache_params, attention_mask=attention_mask)
        hidden_states = residual + hidden_states
        return hidden_states


@auto_docstring
class MambaPreTrainedModel(PreTrainedModel):
    config: MambaConfig
    base_model_prefix = "backbone"
    _no_split_modules = ["MambaBlock", "MambaMixer"]
    supports_gradient_checkpointing = True
    _is_stateful = True

    @torch.no_grad()
    def _init_weights(self, module):
        """Initialize the weights."""
        super()._init_weights(module)
        if isinstance(module, MambaMixer):
            # S4D real initialization. These are not discretized!
            # The core is to load them, compute the discrete states, then write the updated state. Keeps the memory bounded
            module.init_mamba_weights()

            init.kaiming_uniform_(module.conv1d.weight, a=math.sqrt(5))
            if module.conv1d.bias is not None:
                init.zeros_(module.conv1d.bias)
            init.kaiming_uniform_(module.out_proj.weight, a=math.sqrt(5))

            if self.config.rescale_prenorm_residual:
                # Reinitialize selected weights subject to the OpenAI GPT-2 Paper Scheme:
                #   > A modified initialization which accounts for the accumulation on the residual path with model depth. Scale
                #   > the weights of residual layers at initialization by a factor of 1/√N where N is the # of residual layers.
                #   >   -- GPT-2 :: https://openai.com/blog/better-language-models/
                #
                # Reference (Megatron-LM): https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/model/gpt_model.py
                # Special Scaled Initialization --> There are 2 Layer Norms per Transformer Block
                # Following Pytorch init, except scale by 1/sqrt(2 * n_layer)
                # We need to reinit p since this code could be called multiple times
                # Having just p *= scale would repeatedly scale it down
                p = module.out_proj.weight
                p /= math.sqrt(self.config.num_hidden_layers)


@auto_docstring(
    custom_intro="""
    Class for the MAMBA model outputs.
    """
)
@dataclass
class MambaOutput(ModelOutput):
    r"""
    cache_params (`Cache`):
        The state of the model at the last time step. Can be used in a forward method with the next `input_ids` to
        avoid providing the old `input_ids`.

        Includes both the State space model state matrices after the selective scan, and the Convolutional states
    """

    last_hidden_state: torch.FloatTensor | None = None
    cache_params: Cache | None = None
    hidden_states: tuple[torch.FloatTensor] | None = None


@auto_docstring(
    custom_intro="""
    Base class for causal language model (or autoregressive) outputs.
    """
)
@dataclass
class MambaCausalLMOutput(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.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    cache_params (`Cache`):
        The state of the model at the last time step. Can be used in a forward method with the next `input_ids` to
        avoid providing the old `input_ids`.

        Includes both the State space model state matrices after the selective scan, and the Convolutional states
    """

    loss: torch.FloatTensor | None = None
    logits: torch.FloatTensor | None = None
    cache_params: Cache | None = None
    hidden_states: tuple[torch.FloatTensor] | None = None


@auto_docstring
class MambaModel(MambaPreTrainedModel):
    def __init__(self, config):
        super().__init__(config)

        self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size)
        self.layers = nn.ModuleList([MambaBlock(config, layer_idx=idx) for idx in range(config.num_hidden_layers)])

        self.gradient_checkpointing = False
        self.norm_f = MambaRMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
        # Initialize weights and apply final processing
        self._register_load_state_dict_pre_hook(self.load_hook)
        self.post_init()

    def load_hook(self, state_dict, prefix, *args):
        for k in state_dict:
            if "embedding." in k:
                state_dict[k.replace("embedding.", "embeddings.")] = state_dict.pop(k)
                break

    def get_input_embeddings(self):
        return self.embeddings

    def set_input_embeddings(self, new_embeddings):
        self.embeddings = new_embeddings

    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        inputs_embeds: torch.LongTensor | None = None,
        cache_params: Cache | None = None,
        use_cache: bool | None = None,
        output_hidden_states: bool | None = None,
        return_dict: bool | None = None,
        attention_mask: torch.LongTensor | None = None,
        **kwargs,
    ) -> tuple | MambaOutput:
        r"""
        cache_params (`Cache`, *optional*):
            If passed along, the model uses the previous state in all the blocks (which will give the output for the
            `input_ids` provided as if the model add `state_input_ids + input_ids` as context).
        use_cache (`bool`, *optional*):
            If set to `True`, the `cache_params` is returned and can be used to quickly generate the next logits.
        """
        output_hidden_states = (
            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
        )
        use_cache = use_cache if use_cache is not None else (self.config.use_cache if not self.training else False)
        return_dict = return_dict if return_dict is not None else self.config.return_dict

        if (input_ids is None) ^ (inputs_embeds is not None):  # ^ is python for xor
            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")

        if inputs_embeds is None:
            inputs_embeds = self.embeddings(input_ids)

        if self.gradient_checkpointing and self.training and use_cache:
            use_cache = False

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

        hidden_states = inputs_embeds
        all_hidden_states = () if output_hidden_states else None
        for mixer_block in self.layers:
            hidden_states = mixer_block(
                hidden_states,
                cache_params=cache_params,
                attention_mask=attention_mask,
            )

            if output_hidden_states:
                all_hidden_states = all_hidden_states + (hidden_states,)

        hidden_states = self.norm_f(hidden_states)

        if output_hidden_states:
            all_hidden_states = all_hidden_states + (hidden_states,)

        if not return_dict:
            return tuple(v for v in [hidden_states, cache_params, all_hidden_states] if v is not None)

        return MambaOutput(
            last_hidden_state=hidden_states,
            cache_params=cache_params if use_cache else None,
            hidden_states=all_hidden_states,
        )


@auto_docstring(
    custom_intro="""
    The MAMBA Model transformer with a language modeling head on top (linear layer with weights tied to the input
    embeddings).
    """
)
class MambaForCausalLM(MambaPreTrainedModel, GenerationMixin):
    _tied_weights_keys = {"lm_head.weight": "backbone.embeddings.weight"}

    def __init__(self, config):
        super().__init__(config)
        self.backbone = MambaModel(config)
        self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
        # Initialize weights and apply final processing
        self.post_init()

    def get_input_embeddings(self):
        return self.backbone.get_input_embeddings()

    def set_input_embeddings(self, new_embeddings):
        return self.backbone.set_input_embeddings(new_embeddings)

    def prepare_inputs_for_generation(
        self,
        input_ids,
        inputs_embeds=None,
        use_cache=None,
        cache_params: Cache | None = None,
        attention_mask: torch.LongTensor | None = None,
        is_first_iteration: bool | None = False,
        **kwargs,
    ):
        model_inputs = super().prepare_inputs_for_generation(
            input_ids,
            inputs_embeds=inputs_embeds,
            use_cache=use_cache,
            cache_params=cache_params,
            attention_mask=attention_mask,
            is_first_iteration=is_first_iteration,
            **kwargs,
        )

        if use_cache and not is_first_iteration:
            model_inputs["attention_mask"] = None

        return model_inputs

    @auto_docstring
    def forward(
        self,
        input_ids: torch.LongTensor | None = None,
        attention_mask: torch.LongTensor | None = None,
        inputs_embeds: torch.FloatTensor | None = None,
        cache_params: Cache | None = None,
        labels: torch.LongTensor | None = None,
        output_hidden_states: bool | None = None,
        return_dict: bool | None = None,
        use_cache: bool | None = None,
        logits_to_keep: int | torch.Tensor = 0,
        **kwargs,  # for now we need this for generation
    ) -> tuple | MambaCausalLMOutput:
        r"""
        cache_params (`Cache`, *optional*):
            If passed along, the model uses the previous state in all the blocks (which will give the output for the
            `input_ids` provided as if the model add `state_input_ids + input_ids` as context).
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set
            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`
            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`
        use_cache (`bool`, *optional*):
            If set to `True`, the `cache_params` is returned and can be used to quickly generate the next logits.
        """
        return_dict = return_dict if return_dict is not None else self.config.return_dict

        mamba_outputs = self.backbone(
            input_ids,
            cache_params=cache_params,
            inputs_embeds=inputs_embeds,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict,
            use_cache=use_cache,
            attention_mask=attention_mask,
        )

        hidden_states = mamba_outputs[0]
        # Only compute necessary logits
        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, :].to(self.lm_head.weight.dtype)).float()

        loss = None
        if labels is not None:
            # move labels to correct device
            labels = labels.to(logits.device)
            # Shift so that tokens < n predict n
            shift_logits = logits[..., :-1, :].contiguous()
            shift_labels = labels[..., 1:].contiguous()
            # Flatten the tokens
            loss_fct = CrossEntropyLoss()
            loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))

        if not return_dict:
            output = (logits,) + mamba_outputs[1:]
            return ((loss,) + output) if loss is not None else output

        return MambaCausalLMOutput(
            loss=loss,
            logits=logits,
            cache_params=mamba_outputs.cache_params,
            hidden_states=mamba_outputs.hidden_states,
        )


__all__ = ["MambaForCausalLM", "MambaModel", "MambaPreTrainedModel"]
