MiniMaxAI/MiniMax-Text-01
6573k
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