''' This code merges canon_helper.py, configuration_llama_canon.py and modeling_llama_canon.py into a single file to avoid relative imports.''' # configuration_llama_canon.py begins here # Copyright (c) Meta Platforms, Inc. and affiliates. # # This code is modified from the huggingface v4.47-release on the Llama model config # Namely: https://github.com/huggingface/transformers/blob/v4.47-release/src/transformers/models/llama/configuration_llama.py # # Zeyuan's edit note: added support for canon layers, see "Part 4.1, Architecture Design and the Magic of Canon Layers" (https://ssrn.com/abstract=5240330) # # Zeyuan's edit note: added support for qk_norm, see for instance "Scaling Vision Transformers to 22 Billion Parameters" (arxiv.org/abs/2302.05442) # # Zeyuan's edit note: added support for rope_dim, which means only rope_dim of head_dim will be used for rotary position embeddings, if None, then all head_dim will be used # PS: GPTNeoXModel on HF defaults this to 25% of the head_dim, while Llama model sets this to None # # Zeyuan's edit note: the lingua codebase has slightly different RoPE implementation (for which coordinates are real/imaginary), and it is not compatible with the huggingface implementation # so we added a field to specify the version of RoPE, which is a string that can be either 'huggingface' or 'lingua' # When loading a checkpoint trained using the lingua codebase, must set `rope_version='lingua'` # """LLaMA Canon model configuration""" from transformers.configuration_utils import PretrainedConfig from transformers.modeling_rope_utils import rope_config_validation class LlamaCanonConfig(PretrainedConfig): model_type = "LlamaCanon" keys_to_ignore_at_inference = ["past_key_values"] # Default tensor parallel plan for base model `LlamaModel` base_model_tp_plan = { "layers.*.self_attn.q_proj": "colwise", "layers.*.self_attn.k_proj": "colwise", "layers.*.self_attn.v_proj": "colwise", "layers.*.self_attn.o_proj": "rowwise", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } def __init__( self, vocab_size=32000, hidden_size=4096, intermediate_size=11008, num_hidden_layers=32, num_attention_heads=32, num_key_value_heads=None, hidden_act="silu", max_position_embeddings=2048, initializer_range=0.02, rms_norm_eps=1e-6, use_cache=True, pad_token_id=None, bos_token_id=1, eos_token_id=2, pretraining_tp=1, tie_word_embeddings=False, rope_theta=10000.0, rope_scaling=None, attention_bias=False, attention_dropout=0.0, mlp_bias=False, head_dim=None, rope_version='huggingface', **kwargs, ): self.vocab_size = vocab_size self.max_position_embeddings = max_position_embeddings self.hidden_size = hidden_size self.intermediate_size = intermediate_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads # for backward compatibility if num_key_value_heads is None: num_key_value_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads self.hidden_act = hidden_act self.initializer_range = initializer_range self.rms_norm_eps = rms_norm_eps self.pretraining_tp = pretraining_tp self.use_cache = use_cache self.rope_theta = rope_theta self.rope_scaling = rope_scaling self.attention_bias = attention_bias self.attention_dropout = attention_dropout self.mlp_bias = mlp_bias self.head_dim = head_dim if head_dim is not None else self.hidden_size // self.num_attention_heads # Validate the correctness of rotary position embeddings parameters # BC: if there is a 'type' field, copy it it to 'rope_type'. if self.rope_scaling is not None and "type" in self.rope_scaling: self.rope_scaling["rope_type"] = self.rope_scaling["type"] rope_config_validation(self) # Zeyuan's edit note: added support for canon layers, see "Part 4.1, Architecture Design and the Magic of Canon Layers" (https://ssrn.com/abstract=5240330) self.canon_set = kwargs.pop("canon_set", "") self.canon_bias = kwargs.pop("canon_bias", False) self.canon_activation = kwargs.pop("canon_activation", False) self.canon_kernel = kwargs.pop("canon_kernel", 4) self.canon_residual = kwargs.pop("canon_residual", True) # Zeyuan's edit note: added support for qk_norm, see for instance "Scaling Vision Transformers to 22 Billion Parameters" (arxiv.org/abs/2302.05442) self.qk_norm = kwargs.pop("qk_norm", False) # Zeyuan's edit note: added support for rope_dim, which means only rope_dim of head_dim will be used for rotary position embeddings, if None, then all head_dim will be used # PS: GPTNeoXModel on HF defaults this to 25% of the head_dim, while Llama model sets this to None self.rope_dim = kwargs.pop("rope_dim", None) if self.rope_dim is not None: self.partial_rotary_factor = self.rope_dim / self.head_dim # Zeyuan's edit note: the lingua codebase has slightly different RoPE implementation (for which coordinates are real/imaginary), and it is not compatible with the huggingface implementation # so we added a field to specify the version of RoPE, which is a string that can be either 'huggingface' or 'lingua' self.rope_version = rope_version super().__init__( pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs, ) # canon_helper.py begins here # Copyright (c) Meta Platforms, Inc. and affiliates. # from typing import Any, Dict, List, Optional, Tuple import torch import warnings from typing import Optional, Tuple import torch.nn as nn import torch.nn.functional as F from einops import rearrange from transformers.activations import ACT2FN try: from causal_conv1d import causal_conv1d_fn, causal_conv1d_update except ImportError: causal_conv1d_fn = None causal_conv1d_update = None import torch._dynamo @torch._dynamo.disable def causal_conv1d_fn_safe(*args, **kwargs): return causal_conv1d_fn(*args, **kwargs) ## This is an exact copy of `fla.modules.ShortConvolution` with no modification ## The purpose is to make sure you don't need to install fla-org, which is not a stable package yet. class ShortConvolution(nn.Conv1d): """ Simple wrapper around `nn.Conv1d` that accepts dimension last. """ def __init__( self, hidden_size: int, kernel_size: int, bias: bool = False, activation: Optional[str] = 'silu', use_fast_conv1d: Optional[bool] = True, device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, ): super().__init__( in_channels=hidden_size, out_channels=hidden_size, kernel_size=kernel_size, groups=hidden_size, bias=bias, padding=kernel_size - 1, device=device, dtype=dtype, ) self.hidden_size = hidden_size self.activation = None if activation is not None: assert activation in ['silu', 'swish'], f"Activation `{activation}` not supported yet." self.activation = activation if causal_conv1d_fn is None: if use_fast_conv1d: raise RuntimeError( "Please either install `causal-conv1d>=1.4.0` to enable fast causal short convolution CUDA kernel " "or set `use_fast_conv1d` to False" ) else: warnings.warn( "The naive Pytorch verison is very slow in practice, " "please run `pip install causal-conv1d>=1.4.0` to install fast causal short convolution CUDA kernel", category=ImportWarning ) self.use_fast_conv1d = use_fast_conv1d def __repr__(self): # THIS helps TorchDynamo avoid collisions return f"CanonLayerCustom(hidden_size={self.hidden_size})" def extra_repr(self): s = ('{in_channels}, {out_channels}, kernel_size={kernel_size}' ', stride={stride}') if self.padding != (0,) * len(self.padding): s += ', padding={padding}' if self.dilation != (1,) * len(self.dilation): s += ', dilation={dilation}' if self.output_padding != (0,) * len(self.output_padding): s += ', output_padding={output_padding}' if self.groups != 1: s += ', groups={groups}' if self.bias is None: s += ', bias=False' if self.padding_mode != 'zeros': s += ', padding_mode={padding_mode}' if self.activation is not None: s += ', activation={activation}' if not self.use_fast_conv1d: s += ', use_fast_conv1d={use_fast_conv1d}' return s.format(**self.__dict__) def forward( self, x: torch.Tensor, mask: Optional[torch.Tensor] = None, cache: Optional[torch.Tensor] = None, output_final_state: bool = False, seq_idx: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Args: x (`torch.Tensor`): Tensor of shape `[batch_size, seq_len, hidden_size]` mask (`Optional[torch.Tensor]`): Attention mask dealing with padded positions. cache (`Optional[torch.Tensor]`): Previous cache tensor of shape `[batch_size, hidden_size, kernel_size]`. If provided, the cache is updated **inplace**. output_final_state (Optional[bool]): Whether to output the final state of shape `[batch_size, hidden_size, kernel_size]`. Default: `False`. seq_idx (Optional[torch.Tensor]): Sequence index for each token. Used for varlen. Default: `None`. Shape: [batch_size, seq_len] Suppose a batch consists of two sequences with lengths 3 and 4, seq_idx=[0, 0, 0, 1, 1, 1, 1] for this batch. Returns: Tensor of shape `[batch_size, seq_len, hidden_size]`. """ batch_size, _, hidden_size = x.shape if mask is not None: x = x.mul_(mask.unsqueeze(-1)) if output_final_state and cache is None: cache = x.new_zeros(batch_size, hidden_size, self.kernel_size[0]) if cache is not None and x.shape[1] == 1: return self.step(x, cache) x = rearrange(x, "b t d -> b d t") # Update state (B D W) if cache is not None: cache.copy_(F.pad(x, (self.kernel_size[0] - x.shape[-1], 0))) if self.use_fast_conv1d: x = causal_conv1d_fn_safe( x=x, weight=rearrange(self.weight, "d 1 w -> d w"), bias=self.bias, activation=self.activation, seq_idx=seq_idx, ) else: x = self._conv_forward(x, self.weight, self.bias)[..., :x.shape[-1]] if self.activation is not None: x = ACT2FN[self.activation](x) # Note I'm using huggingface's ACT2FN here, not fla-org's original one, so that you don't need to install fla-org return rearrange(x, "b d t -> b t d"), cache def step( self, x: torch.Tensor, cache: torch.Tensor ): assert x.shape[1] == 1, "Only support decoding with 1 token at a time for now" x = x.squeeze(1) if self.use_fast_conv1d: x = causal_conv1d_update( x=x, conv_state=cache, weight=rearrange(self.weight, "d 1 w -> d w"), bias=self.bias, activation=self.activation, ) else: dtype = x.dtype cache.copy_(torch.roll(cache, shifts=-1, dims=-1)) cache[:, :, -1] = x x = torch.sum(cache * rearrange(self.weight, "d 1 w -> d w"), dim=-1) if self.bias is not None: x = x + self.bias if self.activation is not None: x = ACT2FN[self.activation](x).to(dtype=dtype) return x.unsqueeze(1), cache @property def state_size(self) -> int: return self.hidden_size * self.kernel_size def create_canon(dim, config): canon = ShortConvolution( hidden_size=dim, kernel_size=config.canon_kernel, bias=config.canon_bias, activation='silu' if config.canon_activation else None, use_fast_conv1d=causal_conv1d_fn is not None and config.canon_kernel in [2, 3, 4], ) if config.canon_bias: canon.bias.data = torch.zeros_like(canon.bias) assert False, 'must put this into reset_parameters, as the bias default value may be overwritten by the model initialization' canon._zeyuan_residual = config.canon_residual return canon # Note this attention_mask must be the 1/0 form (1 for not mask, and 0 for mask), 2D [batch_size, seq_len] # This is incompatible with the HF GPT2Model's attention_mask, which is -inf for masked positions def apply_canon(store_name, canon, hidden_states, cache, layer_idx, attention_mask): if cache is not None and not hasattr(cache, store_name): setattr(cache, store_name, [None] * 256) # if you train model deeper than 256 layers (which you shouldn't...), you need to change this number conv_state = None if cache is not None: conv_state = getattr(cache, store_name)[layer_idx] if attention_mask is not None: print("Inside apply_canon, attention_mask", attention_mask.shape, attention_mask) if attention_mask is None: conv_mask = None elif len(attention_mask.shape)==4: assert False, "currently disabled, assuming attention_mask is 2D of the form [batch_size, seq_len]' with 0 and 1's" else: assert len(attention_mask.shape)==2 conv_mask = attention_mask[:, -hidden_states.shape[1] :] if attention_mask is not None else None hidden_states2, conv_state = canon(x=hidden_states, mask=conv_mask, cache=conv_state, output_final_state=cache is not None) if cache is not None: getattr(cache, store_name)[layer_idx] = conv_state if canon._zeyuan_residual: return hidden_states + hidden_states2 else: return hidden_states2 # modeling_llama_canon.py begins here # Copyright (c) Meta Platforms, Inc. and affiliates. # # This code is modified from the huggingface v4.47-release on the Llama model # Namely: https://github.com/huggingface/transformers/blob/v4.47-release/src/transformers/models/llama/modeling_llama.py # # Zeyuan's edit note: added support for canon layers, see "Part 4.1, Architecture Design and the Magic of Canon Layers" (https://ssrn.com/abstract=5240330) # # Zeyuan's edit note: added support for qk_norm, see for instance "Scaling Vision Transformers to 22 Billion Parameters" (arxiv.org/abs/2302.05442) # # Zeyuan's edit note: added support for rope_dim, which means only rope_dim of head_dim will be used for rotary position embeddings, if None, then all head_dim will be used # PS: GPTNeoXModel on HF defaults this to 25% of the head_dim, while Llama model sets this to None # # Zeyuan's edit note: the lingua codebase has slightly different RoPE implementation (for which coordinates are real/imaginary), and it is not compatible with the huggingface implementation # so we added a field to specify the version of RoPE, which is a string that can be either 'huggingface' or 'lingua' # When loading a checkpoint trained using the lingua codebase, must set `rope_version='lingua'` # import math from typing import List, Optional, Tuple, Union import torch import torch.utils.checkpoint from torch import nn from transformers.activations import ACT2FN from transformers.cache_utils import Cache, DynamicCache, StaticCache from transformers.generation import GenerationMixin from transformers.modeling_attn_mask_utils import AttentionMaskConverter from transformers.modeling_flash_attention_utils import FlashAttentionKwargs, _flash_attention_forward from transformers.modeling_outputs import ( BaseModelOutputWithPast, CausalLMOutputWithPast, QuestionAnsweringModelOutput, SequenceClassifierOutputWithPast, TokenClassifierOutput, ) from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS from transformers.modeling_utils import PreTrainedModel from transformers.processing_utils import Unpack from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS from transformers.utils import ( LossKwargs, add_code_sample_docstrings, add_start_docstrings, add_start_docstrings_to_model_forward, is_flash_attn_greater_or_equal_2_10, logging, replace_return_docstrings, ) logger = logging.get_logger(__name__) class LlamaRMSNorm(nn.Module): def __init__(self, hidden_size, eps=1e-6): """ LlamaRMSNorm is equivalent to T5LayerNorm """ super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps def forward(self, hidden_states): input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) return self.weight * hidden_states.to(input_dtype) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" ALL_LAYERNORM_LAYERS.append(LlamaRMSNorm) class LlamaRotaryEmbedding(nn.Module): def __init__( self, dim=None, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0, rope_type="default", config: Optional[LlamaCanonConfig] = None, ): super().__init__() # TODO (joao): remove the `if` below, only used for BC self.rope_kwargs = {} if config is None: logger.warning_once( "`LlamaRotaryEmbedding` can now be fully parameterized by passing the model config through the " "`config` argument. All other arguments will be removed in v4.46" ) self.rope_kwargs = { "rope_type": rope_type, "factor": scaling_factor, "dim": dim, "base": base, "max_position_embeddings": max_position_embeddings, } self.rope_type = rope_type self.max_seq_len_cached = max_position_embeddings self.original_max_seq_len = max_position_embeddings else: # BC: "rope_type" was originally "type" if config.rope_scaling is not None: self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type")) else: self.rope_type = "default" self.max_seq_len_cached = config.max_position_embeddings self.original_max_seq_len = config.max_position_embeddings self.config = config self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device, **self.rope_kwargs) self.register_buffer("inv_freq", inv_freq, persistent=False) self.original_inv_freq = self.inv_freq def _dynamic_frequency_update(self, position_ids, device): """ dynamic RoPE layers should recompute `inv_freq` in the following situations: 1 - growing beyond the cached sequence length (allow scaling) 2 - the current sequence length is in the original scale (avoid losing precision with small sequences) """ seq_len = torch.max(position_ids) + 1 if seq_len > self.max_seq_len_cached: # growth inv_freq, self.attention_scaling = self.rope_init_fn( self.config, device, seq_len=seq_len, **self.rope_kwargs ) self.register_buffer("inv_freq", inv_freq, persistent=False) # TODO joao: may break with compilation self.max_seq_len_cached = seq_len if seq_len < self.original_max_seq_len and self.max_seq_len_cached > self.original_max_seq_len: # reset self.register_buffer("inv_freq", self.original_inv_freq, persistent=False) self.max_seq_len_cached = self.original_max_seq_len @torch.no_grad() def forward(self, x, position_ids): if "dynamic" in self.rope_type: self._dynamic_frequency_update(position_ids, device=x.device) # Core RoPE block inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) position_ids_expanded = position_ids[:, None, :].float() # Force float32 (see https://github.com/huggingface/transformers/pull/29285) device_type = x.device.type device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu" with torch.autocast(device_type=device_type, enabled=False): freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() sin = emb.sin() # Advanced RoPE types (e.g. yarn) apply a post-processing scaling factor, equivalent to scaling attention cos = cos * self.attention_scaling sin = sin * self.attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) class LlamaLinearScalingRotaryEmbedding(LlamaRotaryEmbedding): """LlamaRotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev""" def __init__(self, *args, **kwargs): logger.warning_once( "`LlamaLinearScalingRotaryEmbedding` is deprecated an will be removed in v4.46. Please use " "`LlamaRotaryEmbedding`, which now also does linear scaling (simply pass the model config to __init__)." ) kwargs["rope_type"] = "linear" super().__init__(*args, **kwargs) class LlamaDynamicNTKScalingRotaryEmbedding(LlamaRotaryEmbedding): """LlamaRotaryEmbedding extended with Dynamic NTK scaling. Credits to the Reddit users /u/bloc97 and /u/emozilla""" def __init__(self, *args, **kwargs): logger.warning_once( "`LlamaDynamicNTKScalingRotaryEmbedding` is deprecated an will be removed in v4.46. Please use " "`LlamaRotaryEmbedding`, which now also does dynamic ntk scaling (simply pass the model config to " "__init__)." ) kwargs["rope_type"] = "dynamic" super().__init__(*args, **kwargs) def rotate_half(x): """Rotates half the hidden dims of the input.""" x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1, rope_version='huggingface'): """Applies Rotary Position Embedding to the query and key tensors. Args: q (`torch.Tensor`): The query tensor. k (`torch.Tensor`): The key tensor. cos (`torch.Tensor`): The cosine part of the rotary embedding. sin (`torch.Tensor`): The sine part of the rotary embedding. position_ids (`torch.Tensor`, *optional*): Deprecated and unused. unsqueeze_dim (`int`, *optional*, defaults to 1): The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. Returns: `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. """ if rope_version == 'huggingface': cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) elif rope_version == 'lingua': B, H, S, D = q.shape assert D % 2 == 0, "head_dim must be even" half = D // 2 # 1) take just the first half of cos/sin cos_h = cos[..., :half] # (B, S, half) sin_h = sin[..., :half] # (B, S, half) # 2) broadcast over heads cos_h = cos_h.unsqueeze(unsqueeze_dim) # (B, 1, S, half) sin_h = sin_h.unsqueeze(unsqueeze_dim) # 3) group into (even,odd) pairs --- note q/k may have different number of heads, so -1 means the head dimension q2 = q.view(B, -1, S, half, 2) # (B, H, S, half, 2) k2 = k.view(B, -1, S, half, 2) q_even, q_odd = q2[..., 0], q2[..., 1] # each (B, H, S, half) k_even, k_odd = k2[..., 0], k2[..., 1] # 4) apply [cos -sin; sin cos] to each pair # out0 = x0*cos + x1*sin # out1 = -x0*sin + x1*cos q_rot_even = q_even * cos_h - q_odd * sin_h q_rot_odd = q_even * sin_h + q_odd * cos_h k_rot_even = k_even * cos_h - k_odd * sin_h k_rot_odd = k_even * sin_h + k_odd * cos_h # 5) re-interleave back to (B, H, S, D) q_embed = torch.stack([q_rot_even, q_rot_odd], dim=-1).reshape(B, -1, S, D) k_embed = torch.stack([k_rot_even, k_rot_odd], dim=-1).reshape(B, -1, S, D) else: assert False, f"Unknown rope version: {rope_version}. Supported versions are 'huggingface' and 'lingua'." return q_embed, k_embed class LlamaCanonMLP(nn.Module): def __init__(self, config: LlamaCanonConfig): super().__init__() self.config = config self.hidden_size = config.hidden_size if config.intermediate_size is None: # Use this default one so a d-dim hidden size will mean 8d^2 params for the GatedMLP self.intermediate_size = config.hidden_size * 8 // 3 else: self.intermediate_size = config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias) self.act_fn = ACT2FN[config.hidden_act] # optional canonD if "D" in config.canon_set: self.canonD = create_canon(self.intermediate_size * 2, config) else: self.canonD = None def forward(self, x: torch.Tensor, old_attention_mask: Optional[torch.Tensor] = None, past_key_value: Optional[Cache] = None, layer_idx: Optional[int] = None): x1 = self.gate_proj(x) x3 = self.up_proj(x) if self.canonD is not None: cat = torch.cat([x1, x3], dim=-1) x1, x3 = apply_canon("canonD", self.canonD, hidden_states=cat, cache=past_key_value, layer_idx=layer_idx, attention_mask=old_attention_mask).chunk(2, dim=-1) return self.down_proj(self.act_fn(x1) * x3) def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: """ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch, num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim) """ batch, num_key_value_heads, slen, head_dim = hidden_states.shape if n_rep == 1: return hidden_states hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) class LlamaCanonAttention(nn.Module): """Multi-headed attention with optional Q/K norm and canonB""" def __init__(self, config: LlamaCanonConfig, layer_idx: Optional[int] = None): super().__init__() self.config = config self.layer_idx = layer_idx if layer_idx is None: logger.warning_once( f"Instantiating {self.__class__.__name__} without passing a `layer_idx` is not recommended and will " "lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` " "when creating this class." ) self.attention_dropout = config.attention_dropout self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads self.head_dim = getattr(config, "head_dim", self.hidden_size // self.num_heads) self.num_key_value_heads = config.num_key_value_heads self.num_key_value_groups = self.num_heads // self.num_key_value_heads self.max_position_embeddings = config.max_position_embeddings self.rope_theta = config.rope_theta self.is_causal = True self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=config.attention_bias) self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias) self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias) self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.attention_bias) # TODO (joao): remove in v4.46 (RoPE is computed in the model, not in the decoder layers) self.rotary_emb = LlamaRotaryEmbedding(config=self.config) # optional Q/K normalization if config.qk_norm: self.q_norm = LlamaRMSNorm( config.num_attention_heads * self.head_dim, eps=config.rms_norm_eps ) self.k_norm = LlamaRMSNorm( config.num_key_value_heads * self.head_dim, eps=config.rms_norm_eps ) else: self.q_norm = None self.k_norm = None # optional canonB if "B" in config.canon_set: total_dim = ( config.num_attention_heads * self.head_dim + 2 * config.num_key_value_heads * self.head_dim ) self.canonB = create_canon(total_dim, config) else: self.canonB = None def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, old_attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 **kwargs, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: bsz, q_len, _ = hidden_states.size() query_states = self.q_proj(hidden_states) key_states = self.k_proj(hidden_states) value_states = self.v_proj(hidden_states) # apply Q/K norm if self.q_norm is not None: query_states = self.q_norm(query_states) if self.k_norm is not None: key_states = self.k_norm(key_states) # apply canonB if self.canonB is not None: qkv = apply_canon('canonB', self.canonB, hidden_states=torch.cat([query_states, key_states, value_states], dim=-1), cache=past_key_value, layer_idx=self.layer_idx, attention_mask=old_attention_mask) sizes = [ self.config.num_attention_heads * self.head_dim, self.config.num_key_value_heads * self.head_dim, self.config.num_key_value_heads * self.head_dim, ] query_states, key_states, value_states = qkv.split(sizes, dim=-1) # use -1 to infer num_heads and num_key_value_heads as they may vary if tensor parallel is used query_states = query_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) if position_embeddings is None: logger.warning_once( "The attention layers in this model are transitioning from computing the RoPE embeddings internally " "through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed " "`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be " "removed and `position_embeddings` will be mandatory." ) cos, sin = self.rotary_emb(value_states, position_ids) else: cos, sin = position_embeddings # apply rotary with partial rope_dim rope_dim = getattr(self.config, "rope_dim", None) or self.head_dim if rope_dim < self.head_dim: q_rope, q_pass = ( query_states[..., :rope_dim], query_states[..., rope_dim:] ) k_rope, k_pass = ( key_states[..., :rope_dim], key_states[..., rope_dim:] ) q_rope, k_rope = apply_rotary_pos_emb( q_rope, k_rope, cos[..., :rope_dim], sin[..., :rope_dim], rope_version=self.config.rope_version ) query_states = torch.cat([q_rope, q_pass], dim=-1) key_states = torch.cat([k_rope, k_pass], dim=-1) else: query_states, key_states = apply_rotary_pos_emb( query_states, key_states, cos, sin, rope_version=self.config.rope_version ) if past_key_value is not None: # sin and cos are specific to RoPE models; cache_position needed for the static cache cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs) key_states = repeat_kv(key_states, self.num_key_value_groups) value_states = repeat_kv(value_states, self.num_key_value_groups) attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim) if attention_mask is not None: # no matter the length, we just slice it causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] attn_weights = attn_weights + causal_mask # upcast attention to fp32 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype) attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training) attn_output = torch.matmul(attn_weights, value_states) if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim): raise ValueError( f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is" f" {attn_output.size()}" ) attn_output = attn_output.transpose(1, 2).contiguous() attn_output = attn_output.reshape(bsz, q_len, -1) attn_output = self.o_proj(attn_output) if not output_attentions: attn_weights = None return attn_output, attn_weights, past_key_value class LlamaCanonSdpaAttention(LlamaCanonAttention): # Adapted from LlamaCanonAttention.forward def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, old_attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 **kwargs, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: if output_attentions: # TODO: Improve this warning with e.g. `model.config.attn_implementation = "manual"` once this is implemented. logger.warning_once( "LlamaModel is using LlamaSdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, " 'but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.' ) return super().forward( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, ) bsz, q_len, _ = hidden_states.size() query_states = self.q_proj(hidden_states) key_states = self.k_proj(hidden_states) value_states = self.v_proj(hidden_states) # apply Q/K norm if self.q_norm is not None: query_states = self.q_norm(query_states) if self.k_norm is not None: key_states = self.k_norm(key_states) # apply canonB if self.canonB is not None: qkv = apply_canon('canonB', self.canonB, hidden_states=torch.cat([query_states, key_states, value_states], dim=-1), cache=past_key_value, layer_idx=self.layer_idx, attention_mask=old_attention_mask) sizes = [ self.config.num_attention_heads * self.head_dim, self.config.num_key_value_heads * self.head_dim, self.config.num_key_value_heads * self.head_dim, ] query_states, key_states, value_states = qkv.split(sizes, dim=-1) # use -1 to infer num_heads and num_key_value_heads as they may vary if tensor parallel is used query_states = query_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2) if position_embeddings is None: logger.warning_once( "The attention layers in this model are transitioning from computing the RoPE embeddings internally " "through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed " "`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be " "removed and `position_embeddings` will be mandatory." ) cos, sin = self.rotary_emb(value_states, position_ids) else: cos, sin = position_embeddings # apply rotary with partial rope_dim rope_dim = getattr(self.config, "rope_dim", None) or self.head_dim if rope_dim < self.head_dim: q_rope, q_pass = ( query_states[..., :rope_dim], query_states[..., rope_dim:] ) k_rope, k_pass = ( key_states[..., :rope_dim], key_states[..., rope_dim:] ) q_rope, k_rope = apply_rotary_pos_emb( q_rope, k_rope, cos[..., :rope_dim], sin[..., :rope_dim], rope_version=self.config.rope_version ) query_states = torch.cat([q_rope, q_pass], dim=-1) key_states = torch.cat([k_rope, k_pass], dim=-1) else: # apply rotary with full rope_dim query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, rope_version=self.config.rope_version) if past_key_value is not None: # sin and cos are specific to RoPE models; cache_position needed for the static cache cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs) key_states = repeat_kv(key_states, self.num_key_value_groups) value_states = repeat_kv(value_states, self.num_key_value_groups) causal_mask = attention_mask if attention_mask is not None: causal_mask = causal_mask[:, :, :, : key_states.shape[-2]] #print(f"Inside LlamaCanonSdpaAttention: attention_mask={causal_mask.shape if causal_mask is not None else None} and value = {causal_mask}") # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with custom attn_mask, # Reference: https://github.com/pytorch/pytorch/issues/112577. if query_states.device.type == "cuda" and causal_mask is not None: query_states = query_states.contiguous() key_states = key_states.contiguous() value_states = value_states.contiguous() # We dispatch to SDPA's Flash Attention or Efficient kernels via this `is_causal` if statement instead of an inline conditional assignment # in SDPA to support both torch.compile's dynamic shapes and full graph options. An inline conditional prevents dynamic shapes from compiling. is_causal = True if causal_mask is None and q_len > 1 else False attn_output = torch.nn.functional.scaled_dot_product_attention( query_states, key_states, value_states, attn_mask=causal_mask, dropout_p=self.attention_dropout if self.training else 0.0, is_causal=is_causal, ) attn_output = attn_output.transpose(1, 2).contiguous() attn_output = attn_output.view(bsz, q_len, -1) attn_output = self.o_proj(attn_output) return attn_output, None, past_key_value LLAMA_ATTENTION_CLASSES = { "eager": LlamaCanonAttention, "flash_attention_2": "too lazy to implement, sorry", "sdpa": LlamaCanonSdpaAttention, } class LlamaCanonDecoderLayer(nn.Module): def __init__(self, config: LlamaCanonConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size self.layer_idx = layer_idx self.self_attn = LLAMA_ATTENTION_CLASSES[config._attn_implementation](config=config, layer_idx=layer_idx) self.mlp = LlamaCanonMLP(config) self.input_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.post_attention_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) # optional canonA if "A" in config.canon_set: self.canonA = create_canon(config.hidden_size, config) else: self.canonA = None # optional canonC if "C" in config.canon_set: self.canonC = create_canon(config.hidden_size, config) else: self.canonC = None def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, old_attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: Optional[bool] = False, use_cache: Optional[bool] = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 **kwargs, ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]: if not use_cache: assert past_key_value is None, "past_key_value should be None when use_cache is False" residual = hidden_states hidden_states = self.input_layernorm(hidden_states) if self.canonA is not None: hidden_states = apply_canon('canonA', self.canonA, hidden_states, cache=past_key_value, layer_idx=self.layer_idx, attention_mask=old_attention_mask) # Self Attention hidden_states, self_attn_weights, present_key_value = self.self_attn( hidden_states=hidden_states, attention_mask=attention_mask, old_attention_mask=old_attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, **kwargs, ) hidden_states = residual + hidden_states # Fully Connected residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) if self.canonC is not None: hidden_states = apply_canon('canonC', self.canonC, hidden_states, cache=past_key_value, layer_idx=self.layer_idx, attention_mask=old_attention_mask) hidden_states = self.mlp(hidden_states, old_attention_mask=old_attention_mask, past_key_value=past_key_value, layer_idx=self.layer_idx) hidden_states = residual + hidden_states outputs = (hidden_states,) if output_attentions: outputs += (self_attn_weights,) if use_cache: outputs += (present_key_value,) return outputs class LlamaCanonPreTrainedModel(PreTrainedModel): config_class = LlamaCanonConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["LlamaCanonDecoderLayer"] _skip_keys_device_placement = ["past_key_values"] _supports_flash_attn_2 = True _supports_sdpa = True _supports_cache_class = True _supports_quantized_cache = True _supports_static_cache = True def _init_weights(self, module): std = self.config.initializer_range if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=std) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=std) if module.padding_idx is not None: module.weight.data[module.padding_idx].zero_() elif isinstance(module, ShortConvolution): module.reset_parameters() # Use Kaiming initialization class LlamaCanonModel(LlamaCanonPreTrainedModel): def __init__(self, config: LlamaCanonConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList( [LlamaCanonDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self.norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.rotary_emb = LlamaRotaryEmbedding(config=config) self.gradient_checkpointing = False if getattr(config, "pretraining_tp", 1) != 1: logger.warn("`pretraining_tp` is deprecated, please use `model.tensor_parallel` instead.") # Initialize weights and apply final processing self.post_init() def get_input_embeddings(self): return self.embed_tokens def set_input_embeddings(self, value): self.embed_tokens = value def forward( self, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, **flash_attn_kwargs: Unpack[FlashAttentionKwargs], ) -> Union[Tuple, BaseModelOutputWithPast]: output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) use_cache = use_cache if use_cache is not None else self.config.use_cache return_dict = return_dict if return_dict is not None else self.config.use_return_dict if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("You must specify exactly one of input_ids or inputs_embeds") if self.gradient_checkpointing and self.training and use_cache: logger.warning_once( "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`." ) use_cache = False if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) # kept for BC (non `Cache` `past_key_values` inputs) return_legacy_cache = False if use_cache and not isinstance(past_key_values, Cache): return_legacy_cache = True if past_key_values is None: past_key_values = DynamicCache() else: past_key_values = DynamicCache.from_legacy_cache(past_key_values) logger.warning_once( "We detected that you are passing `past_key_values` as a tuple of tuples. This is deprecated and " "will be removed in v4.47. Please convert your cache or use an appropriate `Cache` class " "(https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)" ) if cache_position is None: past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 cache_position = torch.arange( past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device ) if position_ids is None: position_ids = cache_position.unsqueeze(0) if attention_mask is not None: assert len(attention_mask.shape)==2 and (attention_mask>=0).all() and (attention_mask<=1).all(), f"attention_mask should be a 2D tensor with values in [0, 1], but got {attention_mask.shape} and values {attention_mask}" # Canon layers / causal_conv1d support more complex forms of attention masks but I'm too lazy to implement it. #print(f"Before _update_causal_mask: attention_mask={attention_mask.shape if attention_mask is not None else None} and value = {attention_mask}") causal_mask = self._update_causal_mask( attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions ) #print(f"After _update_causal_mask: attention_mask={causal_mask.shape if causal_mask is not None else None} and value = {causal_mask}") if attention_mask is not None and (attention_mask==1).all(): attention_mask = None hidden_states = inputs_embeds # create position embeddings to be shared across the decoder layers position_embeddings = self.rotary_emb(hidden_states, position_ids) # decoder layers all_hidden_states = () if output_hidden_states else None all_self_attns = () if output_attentions else None next_decoder_cache = None for decoder_layer in self.layers[: self.config.num_hidden_layers]: if output_hidden_states: all_hidden_states += (hidden_states,) if self.gradient_checkpointing and self.training: layer_outputs = self._gradient_checkpointing_func( decoder_layer.__call__, hidden_states, causal_mask, position_ids, past_key_values, output_attentions, use_cache, cache_position, position_embeddings, ) else: layer_outputs = decoder_layer( hidden_states, attention_mask=causal_mask, old_attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_values, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, **flash_attn_kwargs, ) hidden_states = layer_outputs[0] if use_cache: next_decoder_cache = layer_outputs[2 if output_attentions else 1] if output_attentions: all_self_attns += (layer_outputs[1],) hidden_states = self.norm(hidden_states) # add hidden states from the last decoder layer if output_hidden_states: all_hidden_states += (hidden_states,) next_cache = next_decoder_cache if use_cache else None if return_legacy_cache: next_cache = next_cache.to_legacy_cache() if not return_dict: return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None) return BaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=next_cache, hidden_states=all_hidden_states, attentions=all_self_attns, ) def _update_causal_mask( self, attention_mask: torch.Tensor, input_tensor: torch.Tensor, cache_position: torch.Tensor, past_key_values: Cache, output_attentions: bool, ): if self.config._attn_implementation == "flash_attention_2": if attention_mask is not None and 0.0 in attention_mask: return attention_mask return None # For SDPA, when possible, we will rely on its `is_causal` argument instead of its `attn_mask` argument, in # order to dispatch on Flash Attention 2. This feature is not compatible with static cache, as SDPA will fail # to infer the attention mask. past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 using_static_cache = isinstance(past_key_values, StaticCache) # When output attentions is True, sdpa implementation's forward method calls the eager implementation's forward if self.config._attn_implementation == "sdpa" and not using_static_cache and not output_attentions: if AttentionMaskConverter._ignore_causal_mask_sdpa( attention_mask, inputs_embeds=input_tensor, past_key_values_length=past_seen_tokens, is_training=self.training, ): return None dtype, device = input_tensor.dtype, input_tensor.device sequence_length = input_tensor.shape[1] if using_static_cache: target_length = past_key_values.get_max_cache_shape() else: target_length = ( attention_mask.shape[-1] if isinstance(attention_mask, torch.Tensor) else past_seen_tokens + sequence_length + 1 ) # In case the provided `attention` mask is 2D, we generate a causal mask here (4D). causal_mask = self._prepare_4d_causal_attention_mask_with_cache_position( attention_mask, sequence_length=sequence_length, target_length=target_length, dtype=dtype, device=device, cache_position=cache_position, batch_size=input_tensor.shape[0], ) if ( self.config._attn_implementation == "sdpa" and attention_mask is not None and attention_mask.device.type == "cuda" and not output_attentions ): # Attend to all tokens in fully masked rows in the causal_mask, for example the relevant first rows when # using left padding. This is required by F.scaled_dot_product_attention memory-efficient attention path. # Details: https://github.com/pytorch/pytorch/issues/110213 min_dtype = torch.finfo(dtype).min causal_mask = AttentionMaskConverter._unmask_unattended(causal_mask, min_dtype) return causal_mask @staticmethod def _prepare_4d_causal_attention_mask_with_cache_position( attention_mask: torch.Tensor, sequence_length: int, target_length: int, dtype: torch.dtype, device: torch.device, cache_position: torch.Tensor, batch_size: int, **kwargs, ): """ Creates a causal 4D mask of shape `(batch_size, 1, query_length, key_value_length)` from a 2D mask of shape `(batch_size, key_value_length)`, or if the input `attention_mask` is already 4D, do nothing. Args: attention_mask (`torch.Tensor`): A 2D attention mask of shape `(batch_size, key_value_length)` or a 4D attention mask of shape `(batch_size, 1, query_length, key_value_length)`. sequence_length (`int`): The sequence length being processed. target_length (`int`): The target length: when generating with static cache, the mask should be as long as the static cache, to account for the 0 padding, the part of the cache that is not filled yet. dtype (`torch.dtype`): The dtype to use for the 4D attention mask. device (`torch.device`): The device to plcae the 4D attention mask on. cache_position (`torch.Tensor`): Indices depicting the position of the input sequence tokens in the sequence. batch_size (`torch.Tensor`): Batch size. """ if attention_mask is not None and attention_mask.dim() == 4: # In this case we assume that the mask comes already in inverted form and requires no inversion or slicing. causal_mask = attention_mask else: min_dtype = torch.finfo(dtype).min causal_mask = torch.full( (sequence_length, target_length), fill_value=min_dtype, dtype=dtype, device=device ) if sequence_length != 1: causal_mask = torch.triu(causal_mask, diagonal=1) causal_mask *= torch.arange(target_length, device=device) > cache_position.reshape(-1, 1) causal_mask = causal_mask[None, None, :, :].expand(batch_size, 1, -1, -1) if attention_mask is not None: causal_mask = causal_mask.clone() # copy to contiguous memory for in-place edit mask_length = attention_mask.shape[-1] padding_mask = causal_mask[:, :, :, :mask_length] + attention_mask[:, None, None, :] padding_mask = padding_mask == 0 causal_mask[:, :, :, :mask_length] = causal_mask[:, :, :, :mask_length].masked_fill( padding_mask, min_dtype ) return causal_mask class KwargsForCausalLM(FlashAttentionKwargs, LossKwargs): ... class LlamaCanonForCausalLM(LlamaCanonPreTrainedModel, GenerationMixin): _tied_weights_keys = ["lm_head.weight"] _tp_plan = {"lm_head": "colwise_rep"} def __init__(self, config: LlamaCanonConfig): super().__init__(config) self.config = config self.model = LlamaCanonModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) # Initialize weights and apply final processing self.post_init() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def set_decoder(self, decoder): self.model = decoder def get_decoder(self): return self.model def forward( self, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, num_logits_to_keep: int = 0, **kwargs: Unpack[KwargsForCausalLM], ) -> Union[Tuple, CausalLMOutputWithPast]: output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) return_dict = return_dict if return_dict is not None else self.config.use_return_dict # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, cache_position=cache_position, **kwargs, ) hidden_states = outputs[0] # Only compute necessary logits, and do not upcast them to float if we are not computing the loss logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :]) loss = None if labels is not None: loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs) if not return_dict: output = (logits,) + outputs[1:] return (loss,) + output if loss is not None else output return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) def load_from_lingua_state(self, state_dict: dict, strict: bool = True): assert self.config.rope_version=='lingua', f"Lingua uses different rope indexing comparing to Huggingface default, must have initialized with `rope_version='lingua'` but got {self.config.rope_version}" mapped = {} for k, v in state_dict.items(): if k.startswith("layers."): parts = k.split('.') idx = parts[1] name = parts[2:] if name[0] == 'attention': sub = name[1] if sub in ['wq','wk','wv','wo']: proj_map = {'wq':'q_proj','wk':'k_proj','wv':'v_proj','wo':'o_proj'} newk = f"model.layers.{idx}.self_attn.{proj_map[sub]}.weight" elif sub in ['q_norm','k_norm','canonB']: newk = f"model.layers.{idx}.self_attn.{sub}.weight" else: continue elif name[0] == 'feed_forward': sub = name[1] mp = {'w1':'gate_proj','w3':'up_proj','w2':'down_proj','canonD':'canonD'} if sub in mp: newk = f"model.layers.{idx}.mlp.{mp[sub]}.weight" else: continue elif name[0] == 'canonA': newk = f"model.layers.{idx}.canonA.weight" elif name[0] == 'canonC': newk = f"model.layers.{idx}.canonC.weight" elif name[0] == 'attention_norm': newk = f"model.layers.{idx}.input_layernorm.weight" elif name[0] == 'ffn_norm': newk = f"model.layers.{idx}.post_attention_layernorm.weight" else: continue mapped[newk] = v elif k == 'tok_embeddings.weight': mapped['model.embed_tokens.weight'] = v elif k == 'norm.weight': mapped['model.norm.weight'] = v elif k == 'output.weight': mapped['lm_head.weight'] = v # print(f"Target has {len(self.state_dict())} keys: \n {list(self.state_dict().keys())}") # print(f"Mapped has {len(mapped)} keys: \n {list(mapped.keys())}") self.load_state_dict(mapped, strict=strict) @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *model_args, variant="default", **kwargs): """ Overrides HF default loader to use custom .pth and config from subfolder. """ def device_map_to_map_location(device_map): if device_map == "cpu": return "cpu" elif device_map == "auto": return None # Let torch figure it out elif isinstance(device_map, dict): # Could be a more complex mapping, may need custom handling return lambda storage, loc: loc # identity (as fallback) elif isinstance(device_map, str): return device_map # e.g., "cuda:0" else: return None device_map = kwargs.pop("device_map", None) map_location = device_map_to_map_location(device_map) from huggingface_hub import hf_hub_download import os, json if os.path.isfile(os.path.join(pretrained_model_name_or_path, variant, "params.json")): config_path = os.path.join(pretrained_model_name_or_path, variant, "params.json") else: config_path = hf_hub_download( repo_id=pretrained_model_name_or_path, filename=f"{variant}/params.json", ) with open(config_path, "r") as f: dd = json.load(f) cfg = LlamaCanonConfig(vocab_size=dd['model']['vocab_size'], hidden_size=dd['model']['dim'], intermediate_size=dd['model']['hidden_dim'], num_hidden_layers=dd['model']['n_layers'], num_attention_heads=dd['model']['n_heads'], num_key_value_heads=dd['model']['n_kv_heads'] if 'n_kv_heads' in dd['model'] else None, qk_norm = dd['model'].get('qk_norm', False), rope_dim = dd['model'].get('rope_dim', None), canon_set = dd['model'].get('canon_set', ''), rope_theta = dd['model'].get('rope_theta'), rms_norm_eps = dd['model'].get('norm_eps'), max_position_embeddings=dd['data']['seq_len'], rope_version = 'lingua', ) if dd['model']['hidden_dim'] is None: cfg.intermediate_size = 256 * ((dd['model']['dim'] * 8 + 3*256-1) // (3*256)) cfg._attn_implementation = 'sdpa' if 'rope_dim' in dd['model'] and dd['model']['rope_dim'] is not None and dd['model']['rope_dim'] < dd['model']['dim'] // dd['model']['n_heads']: cfg.partial_rotary_factor = dd['model']['rope_dim'] / (dd['model']['dim'] // dd['model']['n_heads']) logger.info(f"Converted lingua params.json to Huggingface config; next creating Huggingface model") model = LlamaCanonForCausalLM(cfg) if os.path.isfile(os.path.join(pretrained_model_name_or_path, variant, "consolidated.pth")): weights_path = os.path.join(pretrained_model_name_or_path, variant, "consolidated.pth") else: weights_path = hf_hub_download( repo_id=pretrained_model_name_or_path, filename=f"{variant}/consolidated.pth", ) logger.info(f"Loading lingua model weights from {weights_path}") state = torch.load(weights_path, map_location=map_location, weights_only=True) model.load_from_lingua_state(state['model']) logger.info(f"Successfully converted lingua state to Huggingface state") return model