Team Ai
Modelpublic

MiniMaxAI/MiniMax-Text-01

sourceHugging Faceupdated 1y agoView on Hugging Face
657likes3kdownloads
modeling_minimax_text_01.py1702 linesDownload Raw Back to root
1""" PyTorch MiniMaxText01 model."""2import inspect3import math4import warnings5from typing import List, Optional, Tuple, Union6import os7import copy8import torch9import torch.nn.functional as F10import torch.utils.checkpoint11from torch import nn12from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss13from einops import rearrange, repeat14from transformers.activations import ACT2FN15from transformers.cache_utils import Cache, DynamicCache16from transformers.modeling_attn_mask_utils import (17    _prepare_4d_causal_attention_mask,18)19from transformers.modeling_outputs import (20    MoeCausalLMOutputWithPast,21    MoeModelOutputWithPast,22    SequenceClassifierOutputWithPast,23)24from transformers.modeling_utils import PreTrainedModel25from transformers.utils import (26    add_start_docstrings,27    add_start_docstrings_to_model_forward,28    is_flash_attn_2_available,29    is_flash_attn_greater_or_equal_2_10,30    logging,31    replace_return_docstrings,32)33from transformers.utils.import_utils import is_torch_fx_available34from .configuration_minimax_text_01 import MiniMaxText01Config35 36if is_flash_attn_2_available():37    from flash_attn import flash_attn_func, flash_attn_varlen_func38    from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input  # noqa39 40    _flash_supports_window_size = "window_size" in list(inspect.signature(flash_attn_func).parameters)41 42# This makes `_prepare_4d_causal_attention_mask` a leaf function in the FX graph.43# It means that the function will not be traced through and simply appear as a node in the graph.44if is_torch_fx_available():45    _prepare_4d_causal_attention_mask = torch.fx.wrap(_prepare_4d_causal_attention_mask)46    47use_triton = eval(os.environ.get("use_triton", default="False"))48debug = eval(os.environ.get("debug", default="False"))49do_eval = eval(os.environ.get("do_eval", default="False"))50eval_and_not_generate = eval(os.environ.get("eval_and_not_generate", default="False"))51BLOCK = 25652 53logger = logging.get_logger(__name__)54 55_CONFIG_FOR_DOC = "MiniMaxText01Config"56 57 58def get_activation_fn(activation):59    if debug:60        logger.info(f"activation: {activation}")61    if activation == "gelu":62        return F.gelu63    elif activation == "relu":64        return F.relu65    elif activation == "elu":66        return F.elu67    elif activation == "sigmoid":68        return F.sigmoid69    elif activation == "exp":70 71        def f(x):72            with torch.no_grad():73                x_max = torch.max(x, dim=-1, keepdims=True).values74            y = torch.exp(x - x_max)75 76            return y77 78        return f79    elif activation == "leak":80        return F.leaky_relu81    elif activation == "1+elu":82 83        def f(x):84            return 1 + F.elu(x)85 86        return f87    elif activation == "2+elu":88 89        def f(x):90            return 2 + F.elu(x)91 92        return f93    elif activation == "silu" or activation == "swish":94        return F.silu95    elif activation == "sine":96        return torch.sin97    else:98        logger.info(99            f"activation: does not support {activation}, use Identity!!!")100        return lambda x: x101 102 103def load_balancing_loss_func(104        gate_logits: torch.Tensor, num_experts: torch.Tensor = None, top_k=2,105        attention_mask: Optional[torch.Tensor] = None106) -> float:107    r"""108    Computes auxiliary load balancing loss as in Switch Transformer - implemented in Pytorch.109 110    See Switch Transformer (https://arxiv.org/abs/2101.03961) for more details. This function implements the loss111    function presented in equations (4) - (6) of the paper. It aims at penalizing cases where the routing between112    experts is too unbalanced.113 114    Args:115        gate_logits (Union[`torch.Tensor`, Tuple[torch.Tensor]):116            Logits from the `gate`, should be a tuple of model.config.num_hidden_layers tensors of117            shape [batch_size X sequence_length, num_experts].118        attention_mask (`torch.Tensor`, None):119            The attention_mask used in forward function120            shape [batch_size X sequence_length] if not None.121        num_experts (`int`, *optional*):122            Number of experts123 124    Returns:125        The auxiliary loss.126    """127    if gate_logits is None or not isinstance(gate_logits, tuple):128        return 0129 130    if isinstance(gate_logits, tuple):131        compute_device = gate_logits[0].device132        concatenated_gate_logits = torch.cat([layer_gate.to(compute_device) for layer_gate in gate_logits], dim=0)133 134    routing_weights = torch.nn.functional.softmax(concatenated_gate_logits, dim=-1)135 136    _, selected_experts = torch.topk(routing_weights, top_k, dim=-1)137 138    expert_mask = torch.nn.functional.one_hot(selected_experts, num_experts)139 140    if attention_mask is None:141        # Compute the percentage of tokens routed to each experts142        tokens_per_expert = torch.mean(expert_mask.float(), dim=0)143 144        # Compute the average probability of routing to these experts145        router_prob_per_expert = torch.mean(routing_weights, dim=0)146    else:147        batch_size, sequence_length = attention_mask.shape148        num_hidden_layers = concatenated_gate_logits.shape[0] // (batch_size * sequence_length)149 150        # Compute the mask that masks all padding tokens as 0 with the same shape of expert_mask151        expert_attention_mask = (152            attention_mask[None, :, :, None, None]153            .expand((num_hidden_layers, batch_size, sequence_length, top_k, num_experts))154            .reshape(-1, top_k, num_experts)155            .to(compute_device)156        )157 158        # Compute the percentage of tokens routed to each experts159        tokens_per_expert = torch.sum(expert_mask.float() * expert_attention_mask, dim=0) / torch.sum(160            expert_attention_mask, dim=0161        )162 163        # Compute the mask that masks all padding tokens as 0 with the same shape of tokens_per_expert164        router_per_expert_attention_mask = (165            attention_mask[None, :, :, None]166            .expand((num_hidden_layers, batch_size, sequence_length, num_experts))167            .reshape(-1, num_experts)168            .to(compute_device)169        )170 171        # Compute the average probability of routing to these experts172        router_prob_per_expert = torch.sum(routing_weights * router_per_expert_attention_mask, dim=0) / torch.sum(173            router_per_expert_attention_mask, dim=0174        )175 176    overall_loss = torch.sum(tokens_per_expert * router_prob_per_expert.unsqueeze(0))177    return overall_loss * num_experts178 179 180# Copied from transformers.models.llama.modeling_llama._get_unpad_data181def _get_unpad_data(attention_mask):182    seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)183    indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()184    max_seqlen_in_batch = seqlens_in_batch.max().item()185    cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))186    return (187        indices,188        cu_seqlens,189        max_seqlen_in_batch,190    )191 192 193class GLU(nn.Module):194 195    def __init__(self, d1, d2, bias=False):196        super().__init__()197 198        self.l1 = nn.Linear(d1, d2, bias=bias)199        self.l2 = nn.Linear(d1, d2, bias=bias)200        self.l3 = nn.Linear(d2, d1, bias=bias)201 202    def forward(self, x):203        o1 = self.l1(x)204        o2 = self.l2(x)205        output = o1 * o2206        output = self.l3(output)207        return output208 209 210class MiniMaxText01LightningAttention(nn.Module):211    def __init__(self, config: MiniMaxText01Config, layer_idx: Optional[int] = None):212        super().__init__()213        bias = False214        self.hidden_size = config.hidden_size215        self.num_heads = config.num_attention_heads216        self.head_dim = getattr(config, 'head_dim', self.hidden_size // self.num_heads)217 218        self.out_proj = nn.Linear(self.head_dim * self.num_heads, self.hidden_size, bias=bias)219        self.act = get_activation_fn(config.hidden_act)220        self.norm = MiniMaxText01RMSNorm(self.head_dim * self.num_heads)221 222        self.qkv_proj = nn.Linear(self.hidden_size, 3 * self.head_dim * self.num_heads, bias=bias)223        self.output_gate = nn.Linear(self.hidden_size, self.head_dim * self.num_heads, bias=bias)224 225        # for inference only226        self.offset = 0227        self.layer_idx = layer_idx228 229    def forward(230            self,231            hidden_states,232            attn_mask: Optional[torch.Tensor] = None,  # (b, h, n, m)233            output_attentions: bool = False,234            past_key_value: Optional[Tuple[torch.Tensor]] = None,235            use_cache: bool = False,236            slope_rate: Optional[torch.Tensor] = None,237            **kwargs238    ):239        if (not self.training) and (not do_eval):240            return self.inference(241                hidden_states,242                attn_mask,243                output_attentions,244                past_key_value,245                use_cache,246                slope_rate,247            )248 249    def inference(250            self,251            x,252            attn_mask: Optional[torch.Tensor] = None,  # (b, n)253            output_attentions: bool = False,254            past_key_value: Optional[Tuple[torch.Tensor]] = None,255            use_cache: bool = False,256            slope_rate: Optional[torch.Tensor] = None,  # (h, 1, 1)257    ):258        # x: b n d259        b, n, d = x.shape260        # linear map261        qkv = self.act(self.qkv_proj(x))262        new_shape = qkv.size()[:-1] + (self.num_heads, -1)263        qkv = qkv.view(*new_shape)264        q, k, v = torch.split(qkv, [self.head_dim] * 3, dim=3)265        q = q.transpose(1, 2)266        k = k.transpose(1, 2)267        v = v.transpose(1, 2)268 269        if past_key_value is None:270            self.offset = q.shape[-2]271        else:272            self.offset += 1273 274        # for align with metaseq275        ratio = torch.exp(-slope_rate)276 277        # only use for the first time278        if past_key_value is None:279            slope_rate = slope_rate.to(torch.float32)280            if attn_mask is not None:281                v = v.masked_fill((1 - attn_mask).unsqueeze(1).unsqueeze(-1).to(torch.bool), 0)282            NUM_BLOCK = (n + BLOCK - 1) // BLOCK283            b, h, n, d = q.shape284            e = v.shape[-1]285            # other286            array = torch.arange(BLOCK).to(q) + 1287            q_decay = torch.exp(-slope_rate * array.reshape(-1, 1))288            k_decay = torch.exp(-slope_rate * (BLOCK - array.reshape(-1, 1)))289            index = array[:, None] - array[None, :]290            s_index = slope_rate * index[291                None,292                None,293            ]294            s_index = torch.where(index >= 0, -s_index, float("-inf"))295            diag_decay = torch.exp(s_index)296 297            kv = torch.zeros(b, h, d, e).to(torch.float32).to(q.device)298            output = torch.empty((b, h, n, e), dtype=q.dtype, device=q.device)299            for i in range(NUM_BLOCK):300                si = i * BLOCK301                ei = min(si + BLOCK, n)302                m = ei - si303                qi = q[:, :, si:ei].contiguous()304                ki = k[:, :, si:ei].contiguous()305                vi = v[:, :, si:ei].contiguous()306                qkv_none_diag = torch.matmul(qi * q_decay[:, :m], kv).to(torch.float32)307 308                # diag309                qk = torch.matmul(qi, ki.transpose(-1, -2)).to(torch.float32) * diag_decay[:, :, :m, :m]310                qkv_diag = torch.matmul(qk, vi.to(torch.float32))311                block_decay = torch.exp(-slope_rate * m)312                output[:, :, si:ei] = qkv_none_diag + qkv_diag313                kv = block_decay * kv + torch.matmul((ki * k_decay[:, -m:]).transpose(-1, -2).to(vi.dtype), vi)314 315        else:316            kv = past_key_value317            output = []318            for i in range(n):319                kv = ratio * kv + torch.einsum(320                    "... n d, ... n e -> ... d e",321                    k[:, :, i:i + 1],322                    v[:, :, i:i + 1],323                )324                qkv = torch.einsum("... n e, ... e d -> ... n d", q[:, :, i:i + 1], kv.to(q.dtype))325                output.append(qkv)326            output = torch.concat(output, dim=-2)327        # reshape328        output = rearrange(output, "b h n d -> b n (h d)")329        # normalize330        output = self.norm(output)331        # gate332        output = F.sigmoid(self.output_gate(x)) * output333        # outproj334        output = self.out_proj(output)335 336        attn_weights = None337 338        return output, attn_weights, kv339 340 341# Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->MiniMaxText01342class MiniMaxText01RMSNorm(nn.Module):343    def __init__(self, hidden_size, eps=1e-6):344        """345        MiniMaxText01RMSNorm is equivalent to T5LayerNorm346        """347        super().__init__()348        self.weight = nn.Parameter(torch.ones(hidden_size))349        self.variance_epsilon = eps350 351    def forward(self, hidden_states):352        input_dtype = hidden_states.dtype353        hidden_states = hidden_states.to(torch.float32)354        variance = hidden_states.pow(2).mean(-1, keepdim=True)355        hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)356        return self.weight * hidden_states.to(input_dtype)357 358 359# Copied from transformers.models.mistral.modeling_mistral.MistralRotaryEmbedding with Mistral->MiniMaxText01360class MiniMaxText01RotaryEmbedding(nn.Module):361    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):362        super().__init__()363 364        self.dim = dim365        self.max_position_embeddings = max_position_embeddings366        self.base = base367        inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float().to(device) / self.dim))368        self.register_buffer("inv_freq", inv_freq, persistent=False)369 370        # Build here to make `torch.jit.trace` work.371        self._set_cos_sin_cache(372            seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.float32373        )374 375    def _set_cos_sin_cache(self, seq_len, device, dtype):376        self.max_seq_len_cached = seq_len377        t = torch.arange(self.max_seq_len_cached, device=device, dtype=torch.int64).type_as(self.inv_freq)378 379        freqs = torch.outer(t, self.inv_freq)380        # Different from paper, but it uses a different permutation in order to obtain the same calculation381        emb = torch.cat((freqs, freqs), dim=-1)382        self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)383        self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)384 385    def forward(self, x, seq_len=None):386        # x: [bs, num_attention_heads, seq_len, head_size]387        if seq_len > self.max_seq_len_cached:388            self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=torch.float32)389 390        return (391            self.cos_cached[:seq_len].to(dtype=torch.float32),392            self.sin_cached[:seq_len].to(dtype=torch.float32),393        )394 395 396# Copied from transformers.models.llama.modeling_llama.rotate_half397def rotate_half(x):398    """Rotates half the hidden dims of the input."""399    x1 = x[..., : x.shape[-1] // 2]400    x2 = x[..., x.shape[-1] // 2:]401    return torch.cat((-x2, x1), dim=-1)402 403 404# Copied from transformers.models.mistral.modeling_mistral.apply_rotary_pos_emb405def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):406    """Applies Rotary Position Embedding to the query and key tensors.407 408    Args:409        q (`torch.Tensor`): The query tensor.410        k (`torch.Tensor`): The key tensor.411        cos (`torch.Tensor`): The cosine part of the rotary embedding.412        sin (`torch.Tensor`): The sine part of the rotary embedding.413        position_ids (`torch.Tensor`):414            The position indices of the tokens corresponding to the query and key tensors. For example, this can be415            used to pass offsetted position ids when working with a KV-cache.416        unsqueeze_dim (`int`, *optional*, defaults to 1):417            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and418            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note419            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and420            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes421            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have422            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.423    Returns:424        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.425    """426    dtype = q.dtype427    rot_dim = cos.shape[-1]428    q_, q_pass = q[..., :rot_dim], q[..., rot_dim:]429    k_, k_pass = k[..., :rot_dim], k[..., rot_dim:]430    cos = cos[position_ids].unsqueeze(unsqueeze_dim)431    sin = sin[position_ids].unsqueeze(unsqueeze_dim)432    q_embed = (q_ * cos) + (rotate_half(q_) * sin)433    k_embed = (k_ * cos) + (rotate_half(k_) * sin)434    return torch.cat((q_embed, q_pass), dim=-1).to(dtype), torch.cat((k_embed, k_pass), dim=-1).to(dtype)435 436 437# Copied from transformers.models.llama.modeling_llama.repeat_kv438def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:439    """440    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,441    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)442    """443    batch, num_key_value_heads, slen, head_dim = hidden_states.shape444    if n_rep == 1:445        return hidden_states446    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)447    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)448 449 450# Copied from transformers.models.mistral.modeling_mistral.MistralAttention with Mistral->MiniMaxText01451class MiniMaxText01Attention(nn.Module):452    """453    Multi-headed attention from 'Attention Is All You Need' paper. Modified to use sliding window attention: Longformer454    and "Generating Long Sequences with Sparse Transformers".455    """456 457    def __init__(self, config: MiniMaxText01Config, layer_idx: Optional[int] = None):458        super().__init__()459        self.config = config460        self.layer_idx = layer_idx461        if layer_idx is None:462            logger.warning_once(463                f"Instantiating {self.__class__.__name__} without passing a `layer_idx` is not recommended and will "464                "lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` "465                "when creating this class."466            )467 468        self.hidden_size = config.hidden_size469        self.num_heads = config.num_attention_heads470        self.head_dim = getattr(config, 'head_dim', self.hidden_size // self.num_heads)471        self.num_key_value_heads = config.num_key_value_heads472        self.num_key_value_groups = self.num_heads // self.num_key_value_heads473        self.max_position_embeddings = config.max_position_embeddings474        self.rope_theta = config.rope_theta475        self.is_causal = True476        self.attention_dropout = config.attention_dropout477 478        self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False)479        self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)480        self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)481        self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)482        self.rotary_dim = getattr(config, 'rotary_dim', self.head_dim)483 484        self.rotary_emb = MiniMaxText01RotaryEmbedding(485            self.rotary_dim,486            max_position_embeddings=self.max_position_embeddings,487            base=self.rope_theta,488        )489 490    def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):491        return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()492 493    def forward(494            self,495            hidden_states: torch.Tensor,496            attention_mask: Optional[torch.Tensor] = None,497            position_ids: Optional[torch.LongTensor] = None,498            past_key_value: Optional[Cache] = None,499            output_attentions: bool = False,500            use_cache: bool = False,501            **kwargs,502    ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:503        if "padding_mask" in kwargs:504            warnings.warn(505                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"506            )507        bsz, q_len, _ = hidden_states.size()508 509        query_states = self.q_proj(hidden_states)510        key_states = self.k_proj(hidden_states)511        value_states = self.v_proj(hidden_states)512 513        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)514        key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)515        value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)516 517        kv_seq_len = key_states.shape[-2]518        if past_key_value is not None:519            if self.layer_idx is None:520                raise ValueError(521                    f"The cache structure has changed since version v4.36. If you are using {self.__class__.__name__} "522                    "for auto-regressive decoding with k/v caching, please make sure to initialize the attention class "523                    "with a layer index."524                )525            kv_seq_len += past_key_value.get_usable_length(kv_seq_len, self.layer_idx)526        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)527        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)528 529        if past_key_value is not None:530            cache_kwargs = {"sin": sin, "cos": cos}  # Specific to RoPE models531            key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)532 533        # repeat k/v heads if n_kv_heads < n_heads534        key_states = repeat_kv(key_states, self.num_key_value_groups)535        value_states = repeat_kv(value_states, self.num_key_value_groups)536 537        attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)538 539        if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):540            raise ValueError(541                f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"542                f" {attn_weights.size()}"543            )544 545        if attention_mask is not None:546            if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):547                raise ValueError(548                    f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"549                )550 551            attn_weights = attn_weights + attention_mask552 553        # upcast attention to fp32554        attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)555        attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)556        attn_output = torch.matmul(attn_weights, value_states)557 558        if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):559            raise ValueError(560                f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"561                f" {attn_output.size()}"562            )563 564        attn_output = attn_output.transpose(1, 2).contiguous()565        attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)566 567        attn_output = self.o_proj(attn_output)568 569        if not output_attentions:570            attn_weights = None571 572        return attn_output, attn_weights, past_key_value573 574 575# Copied from transformers.models.mistral.modeling_mistral.MistralFlashAttention2 with Mistral->MiniMaxText01576class MiniMaxText01FlashAttention2(MiniMaxText01Attention):577    """578    MiniMaxText01 flash attention module. This module inherits from `MiniMaxText01Attention` as the weights of the module stays579    untouched. The only required change would be on the forward pass where it needs to correctly call the public API of580    flash attention and deal with padding tokens in case the input contains any of them.581    """582 583    # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2.__init__584    def __init__(self, *args, **kwargs):585        super().__init__(*args, **kwargs)586 587        # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.588        # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignement, that was made default for flash_attn>=2.1. This attribute is used to handle this difference. Reference: https://github.com/Dao-AILab/flash-attention/releases/tag/v2.1.0.589        # Beware that with flash_attn<2.1, using q_seqlen != k_seqlen (except for the case q_seqlen == 1) produces a wrong mask (top-left).590        self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()591 592    def forward(593            self,594            hidden_states: torch.Tensor,595            attention_mask: Optional[torch.Tensor] = None,596            position_ids: Optional[torch.LongTensor] = None,597            past_key_value: Optional[Union[Cache, Tuple[torch.Tensor]]] = None,598            output_attentions: bool = False,599            use_cache: bool = False,600            **kwargs,601    ):602        if "padding_mask" in kwargs:603            warnings.warn(604                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"605            )606 607            # overwrite attention_mask with padding_mask608            attention_mask = kwargs.pop("padding_mask")609        bsz, q_len, _ = hidden_states.size()610 611        query_states = self.q_proj(hidden_states)612        key_states = self.k_proj(hidden_states)613        value_states = self.v_proj(hidden_states)614 615        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)616        key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)617        value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)618 619        kv_seq_len = key_states.shape[-2]620        if past_key_value is not None:621            kv_seq_len += past_key_value[0].shape[-3]622 623        # Because the input can be padded, the absolute sequence length depends on the max position id.624        rotary_seq_len = max(kv_seq_len, position_ids[:, -1].max().item()) + 1625        cos, sin = self.rotary_emb(value_states, seq_len=rotary_seq_len)626 627        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)628 629        use_sliding_windows = (630                _flash_supports_window_size631                and getattr(self.config, "sliding_window", None) is not None632                and kv_seq_len > self.config.sliding_window633        )634 635        if not _flash_supports_window_size:636            logger.warning_once(637                "The current flash attention version does not support sliding window attention, for a more memory efficient implementation"638                " make sure to upgrade flash-attn library."639            )640 641        dropout_rate = 0.0 if not self.training else self.attention_dropout642 643        # In PEFT, usually we cast the layer norms in float32 for training stability reasons644        # therefore the input hidden states gets silently casted in float32. Hence, we need645        # cast them back in float16 just to be sure everything works as expected.646        input_dtype = query_states.dtype647        if input_dtype == torch.float32:648            if torch.is_autocast_enabled():649                target_dtype = torch.get_autocast_gpu_dtype()650            # Handle the case where the model is quantized651            elif hasattr(self.config, "_pre_quantization_dtype"):652                target_dtype = self.config._pre_quantization_dtype653            else:654                target_dtype = self.q_proj.weight.dtype655 656            logger.warning_once(657                f"The input hidden states seems to be silently casted in float32, this might be related to"658                f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"659                f" {target_dtype}."660            )661 662            query_states = query_states.to(target_dtype)663            key_states = key_states.to(target_dtype)664            value_states = value_states.to(target_dtype)665 666        # Reshape to the expected shape for Flash Attention667        query_states = query_states.transpose(1, 2)668        key_states = key_states.transpose(1, 2)669        value_states = value_states.transpose(1, 2)670 671        if past_key_value is not None:672            # reuse k, v, for evaluation only673            key_states = torch.cat([past_key_value[0], key_states], dim=-3)674            value_states = torch.cat([past_key_value[1], value_states], dim=-3)675        676        past_key_value = (key_states, value_states) if use_cache else None677 678        attn_output = self._flash_attention_forward(679            query_states,680            key_states,681            value_states,682            attention_mask,683            q_len,684            dropout=dropout_rate,685            use_sliding_windows=use_sliding_windows,686        )687 688        attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()689        attn_output = self.o_proj(attn_output)690 691        if not output_attentions:692            attn_weights = None693 694        return attn_output, attn_weights, past_key_value695 696    def _flash_attention_forward(697            self,698            query_states,699            key_states,700            value_states,701            attention_mask,702            query_length,703            dropout=0.0,704            softmax_scale=None,705            use_sliding_windows=False,706    ):707        """708        Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token709        first unpad the input, then computes the attention scores and pad the final attention scores.710 711        Args:712            query_states (`torch.Tensor`):713                Input query states to be passed to Flash Attention API714            key_states (`torch.Tensor`):715                Input key states to be passed to Flash Attention API716            value_states (`torch.Tensor`):717                Input value states to be passed to Flash Attention API718            attention_mask (`torch.Tensor`):719                The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the720                position of padding tokens and 1 for the position of non-padding tokens.721            dropout (`float`):722                Attention dropout723            softmax_scale (`float`, *optional*):724                The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)725            use_sliding_windows (`bool`, *optional*):726                Whether to activate sliding window attention.727        """728        if not self._flash_attn_uses_top_left_mask:729            causal = self.is_causal730        else:731            # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1. For details, please see the comment in LlamaFlashAttention2 __init__.732            causal = self.is_causal and query_length != 1733 734        # Contains at least one padding token in the sequence735        if attention_mask is not None:736            batch_size = query_states.shape[0]737            query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(738                query_states, key_states, value_states, attention_mask, query_length739            )740 741            cu_seqlens_q, cu_seqlens_k = cu_seq_lens742            max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens743 744            if not use_sliding_windows:745                attn_output_unpad = flash_attn_varlen_func(746                    query_states,747                    key_states,748                    value_states,749                    cu_seqlens_q=cu_seqlens_q,750                    cu_seqlens_k=cu_seqlens_k,751                    max_seqlen_q=max_seqlen_in_batch_q,752                    max_seqlen_k=max_seqlen_in_batch_k,753                    dropout_p=dropout,754                    softmax_scale=softmax_scale,755                    causal=causal,756                )757            else:758                attn_output_unpad = flash_attn_varlen_func(759                    query_states,760                    key_states,761                    value_states,762                    cu_seqlens_q=cu_seqlens_q,763                    cu_seqlens_k=cu_seqlens_k,764                    max_seqlen_q=max_seqlen_in_batch_q,765                    max_seqlen_k=max_seqlen_in_batch_k,766                    dropout_p=dropout,767                    softmax_scale=softmax_scale,768                    causal=causal,769                    window_size=(self.config.sliding_window, self.config.sliding_window),770                )771 772            attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)773        else:774            if not use_sliding_windows:775                attn_output = flash_attn_func(776                    query_states,777                    key_states,778                    value_states,779                    dropout,780                    softmax_scale=softmax_scale,781                    causal=causal,782                )783            else:784                attn_output = flash_attn_func(785                    query_states,786                    key_states,787                    value_states,788                    dropout,789                    softmax_scale=softmax_scale,790                    causal=causal,791                    window_size=(self.config.sliding_window, self.config.sliding_window),792                )793 794        return attn_output795 796    def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):797        batch_size, kv_seq_len, num_heads, head_dim = key_layer.shape798 799        # On the first iteration we need to properly re-create the padding mask800        # by slicing it on the proper place801        if kv_seq_len != attention_mask.shape[-1]:802            attention_mask_num_tokens = attention_mask.shape[-1]803            attention_mask = attention_mask[:, attention_mask_num_tokens - kv_seq_len:]804 805        indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)806 807        key_layer = index_first_axis(key_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)808        value_layer = index_first_axis(value_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k)809 810        if query_length == kv_seq_len:811            query_layer = index_first_axis(812                query_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), indices_k813            )814            cu_seqlens_q = cu_seqlens_k815            max_seqlen_in_batch_q = max_seqlen_in_batch_k816            indices_q = indices_k817        elif query_length == 1:818            max_seqlen_in_batch_q = 1819            cu_seqlens_q = torch.arange(820                batch_size + 1, dtype=torch.int32, device=query_layer.device821            )  # There is a memcpy here, that is very bad.822            indices_q = cu_seqlens_q[:-1]823            query_layer = query_layer.squeeze(1)824        else:825            # The -q_len: slice assumes left padding.826            attention_mask = attention_mask[:, -query_length:]827            query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)828 829        return (830            query_layer,831            key_layer,832            value_layer,833            indices_q,834            (cu_seqlens_q, cu_seqlens_k),835            (max_seqlen_in_batch_q, max_seqlen_in_batch_k),836        )837 838 839class MiniMaxText01MLP(nn.Module):840    def __init__(self, config):841        super().__init__()842        self.config = config843        self.hidden_size = config.hidden_size844        self.intermediate_size = config.intermediate_size845        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)846        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)847        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)848        self.act_fn = ACT2FN[config.hidden_act]849 850    def forward(self, x):851        down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))852        return down_proj853 854 855class MiniMaxText01BlockSparseTop2MLP(nn.Module):856    def __init__(self, config: MiniMaxText01Config):857        super().__init__()858        self.ffn_dim = config.intermediate_size859        self.hidden_dim = config.hidden_size860 861        self.w1 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False)862        self.w2 = nn.Linear(self.ffn_dim, self.hidden_dim, bias=False)863        self.w3 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False)864 865        self.act_fn = ACT2FN[config.hidden_act]866 867    def forward(self, hidden_states):868        current_hidden_states = self.act_fn(self.w1(hidden_states)) * self.w3(hidden_states)869        current_hidden_states = self.w2(current_hidden_states)870        return current_hidden_states871 872 873class MiniMaxText01BLockSparseTop2MLP(MiniMaxText01BlockSparseTop2MLP):874    def __init__(self, *args, **kwargs):875        logger.warning_once(876            "MiniMaxText01BLockSparseTop2MLP is deprecated by MiniMaxText01BlockSparseTop2MLP and will be removed in v4.40."877        )878        super().__init__(*args, **kwargs)879 880 881class MiniMaxText01SparseMoeBlock(nn.Module):882    """883    This implementation is884    strictly equivalent to standard MoE with full capacity (no885    dropped tokens). It's faster since it formulates MoE operations886    in terms of block-sparse operations to accomodate imbalanced887    assignments of tokens to experts, whereas standard MoE either888    (1) drop tokens at the cost of reduced performance or (2) set889    capacity factor to number of experts and thus waste computation890    and memory on padding.891    """892 893    def __init__(self, config):894        super().__init__()895        self.hidden_dim = config.hidden_size896        self.ffn_dim = config.intermediate_size897        self.num_experts = config.num_local_experts898        self.top_k = config.num_experts_per_tok899 900        # gating901        self.gate = nn.Linear(self.hidden_dim, self.num_experts, bias=False)902 903        self.experts = nn.ModuleList([MiniMaxText01BlockSparseTop2MLP(config) for _ in range(self.num_experts)])904 905        # Jitter parameters906        self.jitter_noise = config.router_jitter_noise907 908    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:909        """ """910        batch_size, sequence_length, hidden_dim = hidden_states.shape911        if self.training and self.jitter_noise > 0:912            hidden_states *= torch.empty_like(hidden_states).uniform_(1.0 - self.jitter_noise, 1.0 + self.jitter_noise)913        hidden_states = hidden_states.view(-1, hidden_dim)914        # router_logits: (batch * sequence_length, n_experts)915        router_logits = self.gate(hidden_states)916 917        routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)918        routing_weights, selected_experts = torch.topk(routing_weights, self.top_k, dim=-1)919        routing_weights /= routing_weights.sum(dim=-1, keepdim=True)920        # we cast back to the input dtype921        routing_weights = routing_weights.to(hidden_states.dtype)922 923        final_hidden_states = torch.zeros(924            (batch_size * sequence_length, hidden_dim), dtype=hidden_states.dtype, device=hidden_states.device925        )926 927        # One hot encode the selected experts to create an expert mask928        # this will be used to easily index which expert is going to be sollicitated929        expert_mask = torch.nn.functional.one_hot(selected_experts, num_classes=self.num_experts).permute(2, 1, 0)930 931        # Loop over all available experts in the model and perform the computation on each expert932        for expert_idx in range(self.num_experts):933            expert_layer = self.experts[expert_idx]934            idx, top_x = torch.where(expert_mask[expert_idx])935 936            # Index the correct hidden states and compute the expert hidden state for937            # the current expert. We need to make sure to multiply the output hidden938            # states by `routing_weights` on the corresponding tokens (top-1 and top-2)939            current_state = hidden_states[None, top_x].reshape(-1, hidden_dim)940            current_hidden_states = expert_layer(current_state) * routing_weights[top_x, idx, None]941 942            # However `index_add_` only support torch tensors for indexing so we'll use943            # the `top_x` tensor here.944            final_hidden_states.index_add_(0, top_x, current_hidden_states.to(hidden_states.dtype))945        final_hidden_states = final_hidden_states.reshape(batch_size, sequence_length, hidden_dim)946        return final_hidden_states, router_logits947 948 949class MiniMaxText01DecoderLayer(nn.Module):950    def __init__(self, config: MiniMaxText01Config, layer_idx: int):951        super().__init__()952        self.config = config953        self.hidden_size = config.hidden_size954 955        self.self_attn = self.build_attn(config, layer_idx)956 957        self.layer_idx = layer_idx958 959        self.block_sparse_moe = MiniMaxText01SparseMoeBlock(config)960        self.input_layernorm = MiniMaxText01RMSNorm(config.hidden_size, eps=config.rms_norm_eps)961        self.post_attention_layernorm = MiniMaxText01RMSNorm(config.hidden_size, eps=config.rms_norm_eps)962 963        self.postnorm = getattr(config, 'postnorm', False)964        self.layernorm_attention_alpha = getattr(config, 'layernorm_linear_attention_alpha', 1) \965            if config.attention_type == 0 else getattr(config, 'layernorm_full_attention_alpha', 1)966        self.layernorm_attention_beta = getattr(config, 'layernorm_linear_attention_beta', 1) \967            if config.attention_type == 0 else getattr(config, 'layernorm_full_attention_beta', 1)968        self.layernorm_mlp_alpha = getattr(config, 'layernorm_mlp_alpha', 1)969        self.layernorm_mlp_beta = getattr(config, 'layernorm_mlp_beta', 1)970 971        shared_intermediate = getattr(config, 'shared_intermediate_size', 0)972        self.shared_moe = False973        if shared_intermediate > 0:974            self.shared_moe = True975            self.shared_mlp = MiniMaxText01MLP(config)976            self.coefficient = torch.nn.Linear(self.hidden_size, 1, bias=False)977 978    def build_attn(self, config, layer_idx):979        if config.attention_type == 0:980            Attention_module = MiniMaxText01LightningAttention981        else:982            Attention_module = MiniMaxText01FlashAttention2983 984        return Attention_module(985            config,986            layer_idx987        )988 989    def forward(990            self,991            hidden_states: torch.Tensor,992            attention_mask: Optional[torch.Tensor] = None,993            position_ids: Optional[torch.LongTensor] = None,994            past_key_value: Optional[Tuple[torch.Tensor]] = None,995            output_attentions: Optional[bool] = False,996            output_router_logits: Optional[bool] = False,997            use_cache: Optional[bool] = False,998            slope_rate: Optional[float] = None,999            **kwargs,1000    ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:1001        if "padding_mask" in kwargs:1002            warnings.warn(1003                "Passing `padding_mask` is deprecated and will be removed in v4.37. Please make sure use `attention_mask` instead.`"1004            )1005        """1006        Args:1007            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`1008            attention_mask (`torch.FloatTensor`, *optional*): attention mask of size1009                `(batch, sequence_length)` where padding elements are indicated by 0.1010            past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states1011            output_attentions (`bool`, *optional*):1012                Whether or not to return the attentions tensors of all attention layers. See `attentions` under1013                returned tensors for more detail.1014            output_router_logits (`bool`, *optional*):1015                Whether or not to return the logits of all the routers. They are useful for computing the router loss, and1016                should not be returned during inference.1017            use_cache (`bool`, *optional*):1018                If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding1019                (see `past_key_values`).1020        """1021 1022        residual = hidden_states1023 1024        hidden_states = self.input_layernorm(hidden_states)1025        if self.postnorm:1026            residual = hidden_states1027 1028        hidden_states, self_attn_weights, present_key_value = self.self_attn(1029            hidden_states=hidden_states,1030            position_ids=position_ids,1031            attn_mask=attention_mask,1032            past_key_value=past_key_value,1033            output_attentions=output_attentions,1034            use_cache=use_cache,1035            slope_rate=slope_rate,1036        )1037 1038        hidden_states = residual * self.layernorm_attention_alpha \1039                        + hidden_states * self.layernorm_attention_beta1040 1041        # Fully Connected1042        residual = hidden_states1043        hidden_states = self.post_attention_layernorm(hidden_states)1044        if self.postnorm:1045            residual = hidden_states1046 1047        moe_hidden_states, router_logits = self.block_sparse_moe(hidden_states)1048        if self.shared_moe:1049            output_mlp = self.shared_mlp(hidden_states)1050            weight_fp32 = self.coefficient.weight.float()1051            coef = hidden_states.to(torch.float32) @ weight_fp32.T1052            coef = torch.nn.functional.sigmoid(coef).to(hidden_states.dtype)1053            hidden_states = moe_hidden_states * (1 - coef) + output_mlp * coef1054        else:1055            hidden_states = moe_hidden_states1056 1057        hidden_states = residual * self.layernorm_mlp_alpha \1058                        + hidden_states * self.layernorm_mlp_beta1059 1060        outputs = (hidden_states,)1061 1062        if output_attentions:1063            outputs += (self_attn_weights,)1064 1065        if use_cache:1066            outputs += (present_key_value,)1067 1068        if output_router_logits:1069            outputs += (router_logits,)1070 1071        return outputs1072 1073 1074MIXTRAL_START_DOCSTRING = r"""1075    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the1076    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads1077    etc.)1078 1079    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.1080    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage1081    and behavior.1082 1083    Parameters:1084        config ([`MiniMaxText01Config`]):1085            Model configuration class with all the parameters of the model. Initializing with a config file does not1086            load the weights associated with the model, only the configuration. Check out the1087            [`~PreTrainedModel.from_pretrained`] method to load the model weights.1088"""1089 1090 1091@add_start_docstrings(1092    "The bare MiniMaxText01 Model outputting raw hidden-states without any specific head on top.",1093    MIXTRAL_START_DOCSTRING,1094)1095# Copied from transformers.models.mistral.modeling_mistral.MistralPreTrainedModel with Mistral->MiniMaxText011096class MiniMaxText01PreTrainedModel(PreTrainedModel):1097    config_class = MiniMaxText01Config1098    base_model_prefix = "model"1099    supports_gradient_checkpointing = True1100    _no_split_modules = ["MiniMaxText01DecoderLayer"]1101    _skip_keys_device_placement = "past_key_values"1102    _supports_flash_attn_2 = True1103    _supports_sdpa = True1104 1105    def _init_weights(self, module):1106        std = self.config.initializer_range1107        if isinstance(module, nn.Linear):1108            module.weight.data.normal_(mean=0.0, std=std)1109            if module.bias is not None:1110                module.bias.data.zero_()1111        elif isinstance(module, nn.Embedding):1112            module.weight.data.normal_(mean=0.0, std=std)1113            if module.padding_idx is not None:1114                module.weight.data[module.padding_idx].zero_()1115 1116 1117MIXTRAL_INPUTS_DOCSTRING = r"""1118    Args:1119        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):1120            Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide1121            it.1122 1123            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1124            [`PreTrainedTokenizer.__call__`] for details.1125 1126            [What are input IDs?](../glossary#input-ids)1127        attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):1128            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:1129 1130            - 1 for tokens that are **not masked**,1131            - 0 for tokens that are **masked**.1132 1133            [What are attention masks?](../glossary#attention-mask)1134 1135            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and1136            [`PreTrainedTokenizer.__call__`] for details.1137 1138            If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see1139            `past_key_values`).1140 1141            If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]1142            and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more1143            information on the default strategy.1144 1145            - 1 indicates the head is **not masked**,1146            - 0 indicates the head is **masked**.1147        position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1148            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,1149            config.n_positions - 1]`.1150 1151            [What are position IDs?](../glossary#position-ids)1152        past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):1153            Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape1154            `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape1155            `(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.1156 1157            Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention1158            blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.1159 1160            If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those that1161            don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of all1162            `decoder_input_ids` of shape `(batch_size, sequence_length)`.1163        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):1164            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This1165            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the1166            model's internal embedding lookup matrix.1167        use_cache (`bool`, *optional*):1168            If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see1169            `past_key_values`).1170        output_attentions (`bool`, *optional*):1171            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned1172            tensors for more detail.1173        output_hidden_states (`bool`, *optional*):1174            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for1175            more detail.1176        output_router_logits (`bool`, *optional*):1177            Whether or not to return the logits of all the routers. They are useful for computing the router loss, and1178            should not be returned during inference.1179        return_dict (`bool`, *optional*):1180            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.1181"""1182 1183 1184@add_start_docstrings(1185    "The bare MiniMaxText01 Model outputting raw hidden-states without any specific head on top.",1186    MIXTRAL_START_DOCSTRING,1187)1188# Copied from transformers.models.mistral.modeling_mistral.MistralModel with MISTRAL->MIXTRAL,Mistral->MiniMaxText011189class MiniMaxText01Model(MiniMaxText01PreTrainedModel):1190    """1191    Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`MiniMaxText01DecoderLayer`]1192 1193    Args:1194        config: MiniMaxText01Config1195    """1196 1197    def __init__(self, config: MiniMaxText01Config):1198        super().__init__(config)1199        self.padding_idx = config.pad_token_id1200        self.vocab_size = config.vocab_size

Showing the first 1,200 of 1702 lines. Download the file for the rest.