# Copyright Lightning AI. Licensed under the Apache License 2.0, see LICENSE file. """Full definition of a decoder-only transformer-based language model, all of it in this single file. Based on the nanoGPT implementation: https://github.com/karpathy/nanoGPT and https://github.com/EleutherAI/gpt-neox/tree/main/megatron/model. """ import math from typing import Any, Optional, Tuple, Union, List from functools import partial from transformers import AutoConfig, Qwen2_5OmniForConditionalGeneration import torch import torch.nn as nn import torch.nn.functional as F from typing_extensions import Self import whisper from transformers import Qwen2AudioEncoder, Qwen2AudioConfig from src.audiointeraction.config import Config def qkv_reassemble( param: torch.Tensor, config: Config ) -> torch.Tensor: """Reassemble from a normal to an interleaved placement in a QKV matrix. [Q, K, V, Q, K, V, ...] --> [Q, Q, ..., K, K, ..., V, V, ...] """ q_per_kv = config.n_head // config.n_query_groups qs = [] ks = [] vs = [] for chunk in torch.chunk(param, config.n_query_groups): split = torch.split(chunk, [config.head_size * q_per_kv, config.head_size, config.head_size]) qs.append(split[0]) ks.append(split[1]) vs.append(split[2]) q = torch.cat(qs) k = torch.cat(ks) v = torch.cat(vs) return torch.cat((q, k, v)) class GPT(nn.Module): def __init__(self, config: Config) -> None: super().__init__() assert config.padded_vocab_size is not None self.config = config self.lm_head = nn.Linear( config.n_embd, config.padded_vocab_size, bias=config.lm_head_bias ) self.transformer = nn.ModuleDict( dict( wte=nn.Embedding(config.padded_vocab_size, config.n_embd), h=nn.ModuleList( Block(config, block_idx) for block_idx in range(config.n_layer) ), ln_f=config.norm_class(config.n_embd, eps=config.norm_eps), ) ) self.mask_cache: Optional[torch.Tensor] = None self.max_seq_length = self.config.block_size @property def max_seq_length(self) -> int: return self._max_seq_length @max_seq_length.setter def max_seq_length(self, value: int) -> None: """ When doing inference, the sequences used might be shorter than the model's context length. This allows setting a smaller number to avoid allocating unused memory """ if value > self.config.block_size: raise ValueError( f"Cannot attend to {value}, block size is only {self.config.block_size}." " This is likely because the input text exceeds the supported context length of this model." ) self._max_seq_length = value if not hasattr(self, "cos"): # first call cos, sin = self.rope_cache() self.register_buffer("cos", cos, persistent=False) self.register_buffer("sin", sin, persistent=False) # override elif value != self.cos.size(0): self.cos, self.sin = self.rope_cache(device=self.cos.device) # the mask and kv cache size will get updated on `set_kv_cache`. we cannot update it here because we don't know # if the kv cache is expected if self.mask_cache is not None and self.mask_cache.shape[-1] < value: print(f"Warning: KV cache has length {self.mask_cache.shape[-1]} < {value} = max_seq_length. Call 'set_kv_cache' before doing any forwards!") def reset_parameters(self) -> None: # Trigger resetting the rope-cache self.cos, self.sin = self.rope_cache(device=self.cos.device) def _init_weights(self, module: nn.Module) -> None: """Meant to be used with `gpt.apply(gpt._init_weights)`.""" if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def fill_in_audio_feature(self, input_embeddings: torch.Tensor, batch_size: int, audio_feats_list, audio_pos, tasks) -> torch.Tensor: """Replace AUDIO_PAD positions in input_embeddings with precomputed audio features. Two fill modes, dispatched per-sample by `tasks[batch_idx]`: - "online": audio is streamed in fixed 10-frame chunks. `audio_pos[i]` is a list of (start, end) tuples (each end-start == 10); slice the feature tensor by 10 per chunk. - "offline": audio is one contiguous block. `audio_pos[i]` is a single (start, end) tuple covering `output_len` positions; place the entire feature tensor in one shot. """ _, _, emb_dim = input_embeddings.shape if not (batch_size == len(audio_feats_list) == len(audio_pos) == len(tasks)): raise ValueError( f"length mismatch: batch_size={batch_size}, " f"feats={len(audio_feats_list)}, pos={len(audio_pos)}, tasks={len(tasks)}" ) for batch_idx in range(batch_size): audio_feats = audio_feats_list[batch_idx] segments = audio_pos[batch_idx] if segments is None or segments == -1: continue task = tasks[batch_idx] if task == "offline": # Single big block: place the whole feature tensor at the one segment. start, end = segments[0] if start >= self.max_seq_length: continue end = min(end, self.max_seq_length) input_embeddings[batch_idx, start:end, :] = audio_feats[: end - start] continue # Online: per-10-frame chunk placement. for seg_idx, (start, end) in enumerate(segments): if start > self.max_seq_length: continue if end > self.max_seq_length: input_embeddings[batch_idx, start:self.max_seq_length, :] = torch.zeros( self.max_seq_length - start, emb_dim ) else: audio_feat = audio_feats[seg_idx * 10 : (seg_idx + 1) * 10] seg_len, feat_dim = audio_feat.shape expected_len = end - start if seg_len != expected_len or feat_dim != emb_dim: raise ValueError( f"Loaded feature shape {audio_feat.shape} does not match expected " f"({expected_len}, {emb_dim}) at batch {batch_idx}, segment {seg_idx}") # Overwrite the embedding segment input_embeddings[batch_idx, start:end, :] = audio_feat return input_embeddings def forward( self, idx: torch.Tensor, tasks: Optional[List[str]], batch_size: int, audio_info: Optional[Union[dict, torch.Tensor]] = None, input_pos: Optional[torch.Tensor] = None, input_pos_maxp1: Optional[torch.Tensor] = None, audio_tokens_per_chunk: int = 10, lm_head_chunk_size: int = 0, ) -> Union[torch.Tensor, List[torch.Tensor]]: """ If `input_pos` is provided, the KV cache uses K and V vectors for positions smaller than entries in `input_pos`. For efficiency, pass `input_pos_maxp1` as `max(input_pos) + 1` if already available from your forward algorithm. This slices the KV cache buffers and speeds up multi-head attention. Without `input_pos_maxp1`, the computation uses the full KV cache (`max_seq_length`) with masking applied. Note that inferring `input_pos_maxp1` from `input_pos` causes graph breaks and prevents compilation. Args: idx: Token indices of input sequences, shape `(B, T)`, where `B` is batch size. input_pos: Optional. Positions of input tokens. The default is `arange(T)`. Can have shape `(T,)` or `(B, T)` (batched index). input_pos_maxp1: Optional. See above. lm_head_chunk_size: Optional. If `lm_head_chunk_size > 0`, the final `lm_head` computation is done in chunks of this size. Returns: Logit outputs, shape `(B, T, config.padded_vocab_size)`. If `lm_head_chunk_size > 0`, this is a list of chunks of shape `(B, lm_head_chunk_size, config.padded_vocab_size)`, the final entry can be shorter. """ T = idx.size(1) if self.max_seq_length < T: raise ValueError(f"Cannot forward sequence of length {T}, max seq length is only {self.max_seq_length}.") if input_pos is not None: # use the kv cache if input_pos.dim() > 2: # otherwise, things go wrong in `apply_rope` raise ValueError(f"input_pos must have 1 or 2 dimensions, input_pos.shape = {input_pos.shape}") if input_pos.shape[-1] != T: raise ValueError(f"input_pos.shape[-1] = {input_pos.shape[-1]} != {T} = idx.shape[1], must be the same") cos = batched_index_select(self.cos, 0, input_pos) sin = batched_index_select(self.sin, 0, input_pos) if input_pos.dim() == 1: cos = cos.unsqueeze(0) sin = sin.unsqueeze(0) if self.mask_cache is None: raise TypeError("You need to call `gpt.set_kv_cache()`") mask = batched_index_select(self.mask_cache, 2, input_pos) if mask.dim() > 4: # the mask cache has a batch dim of 1 in addition to the one # we get if input_pos has a batch dimension mask = mask.view(*(mask.shape[0:1] + mask.shape[2:])) if input_pos_maxp1 is not None: # Shorten final dimension so it just covers all `input_pos` entries if input_pos_maxp1 > self.max_seq_length: raise ValueError(f"Positions in 'input_pos' must be in [0,{self.max_seq_length})") mask = mask[..., :input_pos_maxp1] else: # unsqueeze to have a batch dimension cos = self.cos[:T].unsqueeze(0) sin = self.sin[:T].unsqueeze(0) # `cos`, `sin` have shape (1, T, config.rope_n_elem) mask = None # defaults to causal mask input_pos_maxp1 = None x = self.transformer.wte(idx) # token embeddings of shape (B, T, n_embd) # Audio feature injection — dispatch on input type. Encoder features are # already n_embd-dim (projected by audio_tower.proj), so we place them # directly into the input embeddings. # - dict : training path, segment-based fill from precomputed features # - Tensor : inference path, streaming chunk replacement # - None : no audio (e.g. text-only data or inter-token decoding step) if isinstance(audio_info, dict): # T_T (text-only) samples have audio_pos == None — nothing to fill. if audio_info.get("audio_pos") is not None: x = self.fill_in_audio_feature( x, batch_size, audio_info["feats_paths"], audio_info["audio_pos"], tasks, ) elif torch.is_tensor(audio_info): if T > audio_tokens_per_chunk: if x.size(0) != 1: raise ValueError("inference mode, it is not supported for batch size > 1") x[0, T - (audio_tokens_per_chunk + 1): T - 1, :] = audio_info if self.config.scale_embeddings: x = x * torch.tensor(self.config.n_embd ** 0.5, dtype=x.dtype) for block in self.transformer.h: x = block(x, cos, sin, mask, input_pos, input_pos_maxp1) x = self.transformer.ln_f(x) clamp_head = ( partial(do_softcapping, thresh=self.config.final_logit_softcapping) if self.config.final_logit_softcapping is not None else nn.Identity() ) if lm_head_chunk_size > 0: # chunk the lm head logits to reduce the peak memory used by autograd return [ clamp_head(self.lm_head(x_i)) for x_i in x.split(lm_head_chunk_size, dim=1) ] else: return clamp_head(self.lm_head(x)) # (B, T, padded_vocab_size) def rope_cache(self, device: Optional[torch.device] = None) -> Tuple[torch.Tensor, torch.Tensor]: if self.config.rope_adjustments is None: extra_config = None else: adjusted_params_required = ["factor", "low_freq_factor", "high_freq_factor", "original_max_seq_len"] params_present = [param in self.config.rope_adjustments for param in adjusted_params_required] num_params_present = sum(params_present) if num_params_present == 0: extra_config = None # uses standard RoPE elif num_params_present == 4: # These parameters should always be used together so that we don't interfere with standard rope extra_config = { name: self.config.rope_adjustments[name] for name in adjusted_params_required } else: # Some but not all parameters are specified; raise an error missing_params = [ param for param, present in zip(adjusted_params_required, params_present) if not present ] raise ValueError( f"The following adjusted RoPE parameters are missing in rope_adjustments: {', '.join(missing_params)}. " "All adjusted RoPE parameters must be specified together." ) return build_rope_cache( seq_len=self.max_seq_length, n_elem=self.config.rope_n_elem, device=device, condense_ratio=self.config.rope_condense_ratio, base=self.config.rope_base, extra_config=extra_config, ) def set_kv_cache( self, batch_size: int, max_seq_length: Optional[int] = None, rope_cache_length: Optional[int] = None, device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, ) -> None: if rope_cache_length is None: rope_cache_length = self.cos.size(-1) if max_seq_length is None: max_seq_length = self.max_seq_length # initialize the kv cache for all blocks for block in self.transformer.h: block.attn.kv_cache = block.attn.build_kv_cache( batch_size, max_seq_length, rope_cache_length, device, dtype, ) if self.mask_cache is None or self.mask_cache.size(3) != max_seq_length: # passing `attn_mask` to SDPA disables the flash implementation. since we only need the mask # for the kv-cache support (only during inference), we only create it in that situation self.mask_cache = build_mask_cache(max_seq_length, device) def clear_kv_cache(self) -> None: self.mask_cache = None for block in self.transformer.h: block.attn.kv_cache = None class Block(nn.Module): def __init__( self, config: Config, block_idx: int, ) -> None: super().__init__() if not config.parallel_residual and config.shared_attention_norm: raise NotImplementedError( "No checkpoint amongst the ones we support uses this configuration" " (non-parallel residual and shared attention norm)." ) self.norm_1 = config.norm_class(config.n_embd, eps=config.norm_eps) self.attn = CausalSelfAttention(config, block_idx) self.post_attention_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_attention_norm else nn.Identity() ) self.norm_2 = None if config.shared_attention_norm else config.norm_class(config.n_embd, eps=config.norm_eps) self.mlp = config.mlp_class(config) self.post_mlp_norm = ( config.norm_class(config.n_embd, eps=config.norm_eps) if config.post_mlp_norm else nn.Identity() ) self.config = config def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: Optional[torch.Tensor] = None, input_pos: Optional[torch.Tensor] = None, input_pos_maxp1: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Non-parallel residual Parallel residual ┌─ x ┌─ x ──────────────────┐ Note: if `shared_attention_norm` is True, │ ↓ │ ↓ ↓ the output from `norm_1` is reused │ norm_1 │ norm_1 ───────► norm_2 │ ↓ │ ↓ ↓ │ attn │ attn MLP │ ↓ │ ↓ ↓ | post_attn_norm | post_attn_norm post_mlp_norm | ↓ | ↓ ↓ ┌─ └► + └► + ◄─────────────────┘ | ↓ │ norm_2 │ ↓ │ MLP │ ↓ | post_mlp_norm | ↓ └───► + """ x_normed = self.norm_1(x) attention_output = self.attn( x_normed, cos, sin, mask, input_pos, input_pos_maxp1 ) attention_output = self.post_attention_norm(attention_output) if self.config.parallel_residual: if not self.config.shared_attention_norm: x_normed = self.norm_2(x) x = attention_output + x else: x = attention_output + x x_normed = self.norm_2(x) return self.post_mlp_norm(self.mlp(x_normed)) + x class CausalSelfAttention(nn.Module): def __init__(self, config: Config, block_idx: int) -> None: super().__init__() # key, query and value projections for all heads, but in a batch self.qkv = nn.Linear( config.n_embd, (config.n_head + 2 * config.n_query_groups) * config.head_size, # support for grouped/multi queries bias=config.bias or config.attn_bias, ) # output projection self.proj = nn.Linear( config.head_size * config.n_head, config.n_embd, bias=config.bias ) # disabled by default self.kv_cache: Optional[KVCache] = None self.apply_sliding_window_attention = ( config.sliding_window_size is not None and block_idx % config.sliding_window_layer_stride == 0 ) if config.norm_qk: self.norm_q = config.norm_class(config.head_size * config.n_head, eps=config.norm_eps) self.norm_k = config.norm_class(config.head_size * config.n_query_groups, eps=config.norm_eps) else: self.norm_q = self.norm_k = None self.config = config self.block_idx = block_idx # Attention capture flags (for analysis/visualization, disabled by default) self.capture_attn: bool = False self.captured_attn_weights: Optional[torch.Tensor] = None # shape: (B, n_head, T_q, T_k) def forward( self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mask: Optional[torch.Tensor] = None, input_pos: Optional[torch.Tensor] = None, input_pos_maxp1: Optional[torch.Tensor] = None, ) -> torch.Tensor: # Notation: # - B | batch size # - T | time-step (sequence length) # - C | model's embeddings size (n_embd) # - C* | attentions's embeddings size # - nh_(q,k,v) | number of heads for query, key and value # - hs | head size head_size = self.config.head_size n_head = self.config.n_head n_query_groups = self.config.n_query_groups rope_n_elem = self.config.rope_n_elem B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd) # Perform a single multiplication operation using a combined QKV matrix to calculate `query`, `key`, and `value` # instead of individually multiplying the input `x` with the respective weight matrices. qkv = self.qkv(x) # (B, T, 3xC*) # Define query, key and value sizes. # If grouped/multi query is enabled, these sizes are not equal (see the diagram in `lit_gpt/config.py::Config`). query_size = n_head * head_size key_size = value_size = n_query_groups * head_size # Split qkv into query, key and value matrices. q, k, v = qkv.split((query_size, key_size, value_size), dim=-1) # 3x(B, T, C*) if self.config.norm_qk: q = self.norm_q(q) k = self.norm_k(k) # To place the num_heads (nh) dimension right after the batch (B) dimension, the first step is to decouple the # embedding size (C) into num_heads (nh) and head_size (hs). q = q.view(B, T, n_head, head_size) # (B, T, nh_q, hs) k = k.view(B, T, n_query_groups, head_size) # (B, T, nh_k, hs) v = v.view(B, T, n_query_groups, head_size) # (B, T, nh_v, hs) # The tensors `query`, `key`, and `value` are now accurately structured: within each batch element (B), there are # multiple heads (nh), and within each head, there is a sequence of elements (T), each represented by a vector # of size `hs`. q = q.transpose(1, 2) # (B, nh_q, T, hs) k = k.transpose(1, 2) # (B, nh_k, T, hs) v = v.transpose(1, 2) # (B, nh_v, T, hs) # Unlike standard positional embeddings rotary embeddings must be applied at every layer. q_roped = apply_rope(q[..., : rope_n_elem], cos, sin) k_roped = apply_rope(k[..., : rope_n_elem], cos, sin) q = torch.cat((q_roped, q[..., rope_n_elem :]), dim=-1) # (B, nh_q, T, hs) k = torch.cat((k_roped, k[..., rope_n_elem :]), dim=-1) # (B, nh_k, T, hs) # Apply kv-cache during inference. if input_pos is not None: if not isinstance(self.kv_cache, KVCache): raise TypeError("You need to call `gpt.set_kv_cache()`") k, v = self.kv_cache(input_pos, k, v) if input_pos_maxp1 is not None: # Subselect along sequence dimension k = k[..., :input_pos_maxp1, :] v = v[..., :input_pos_maxp1, :] # k, v: (B, nh_k, input_pos_maxp1, hs) # If input_pos_maxp1 is None -> max_seq_length use_flash = (getattr(self.config, "use_flash_attention", True) and mask is None and n_query_groups == n_head ) if use_flash: # FlashAttention: B H T D -> B T H D q = q.transpose(1, 2).contiguous() # (B, T, nh_q, hs) k = k.transpose(1, 2).contiguous() # (B, T, nh_k, hs) v = v.transpose(1, 2).contiguous() # (B, T, nh_v, hs) from flash_attn.flash_attn_interface import flash_attn_func y = flash_attn_func(q, k, v, dropout_p=0.0, causal=True) y = y.transpose(1, 2) # back to B H T D else: # Grouped queries: balance the number of heads across all three matrices. # NOTE: flash attention requires it in training mode. # Multi-query: this step can be skipped since there is only 1 head, allowing us to use broadcasting. if n_query_groups != n_head and (input_pos is None or n_query_groups != 1): q_per_kv = n_head // n_query_groups k = k.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) v = v.repeat_interleave(q_per_kv, dim=1) # (B, nh_q, T, hs) if self.apply_sliding_window_attention: """ Global Window Sliding window Sliding window attention mask + bias = attention mask ┌────────────────────────┐ ┌───────────────────────┐ ┌─────────────────────────┐ │ True False False False │ │ True True True True │ │ True False False False │ │ True True False False │ │ True True True True │ │ True True False False │ │ True True True False │ │ False True True True │ │ False True True False │ │ True True True True │ │ False False True True │ │ False False True True │ └────────────────────────┘ └───────────────────────┘ └─────────────────────────┘ """ if mask is None: mask = torch.ones(T, T, dtype=q.dtype, device=q.device).triu(diagonal=1) mask.masked_fill_(mask.bool(), float("-inf")) mask = mask.view(1, 1, *mask.shape) sliding_window_bias = torch.ones_like(mask).tril(diagonal=-self.config.sliding_window_size) sliding_window_bias.masked_fill_(sliding_window_bias.bool(), float("-inf")) mask += sliding_window_bias # Efficient attention using Flash Attention CUDA kernels. # NOTE: efficient implementation is disabled if `mask` is not None or softcapping is enabled. # ↓ (B, nh, T, hs) @ (B, nh, T, hs).mT --> (B, nh, T, T) @ (B, nh, T, hs) --> (B, nh, T, hs) y = self.scaled_dot_product_attention(q, k, v, mask) # Re-assemble all head outputs side by side. y = y.reshape(B, T, head_size * n_head) # Output projection. return self.proj(y) # (B, T, C) def scaled_dot_product_attention( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: Optional[torch.Tensor] = None ) -> torch.Tensor: scale = 1.0 / math.sqrt(self.config.attention_scores_scalar or self.config.head_size) # with softcapping we cannot use SDPA if self.config.attention_logit_softcapping is not None: scores = q @ k.mT * scale scores = do_softcapping(scores, self.config.attention_logit_softcapping) if mask is None: mask = torch.ones(q.size(2), q.size(2), dtype=q.dtype, device=q.device).triu(diagonal=1) mask.masked_fill_(mask.bool(), torch.finfo(q.dtype).min) scores = scores + mask scores = F.softmax(scores, dim=-1, dtype=torch.float).to(dtype=q.dtype) if self.capture_attn: self.captured_attn_weights = scores.detach() y = scores @ v elif self.capture_attn: # Manual attention computation to capture weights (bypasses fused SDPA kernel) # q: (B, n_head, T_q, hs), k: (B, n_head, T_k, hs) scores = torch.matmul(q.float(), k.float().transpose(-2, -1)) * scale if mask is not None: if mask.dtype == torch.bool: scores = scores.masked_fill(~mask, float('-inf')) else: scores = scores + mask.float() else: # Apply causal mask manually when no mask is provided (training mode) T_q, T_k = q.size(-2), k.size(-2) causal = torch.ones(T_q, T_k, device=q.device, dtype=torch.bool).tril(diagonal=T_k - T_q) scores = scores.masked_fill(~causal, float('-inf')) attn_weights = F.softmax(scores, dim=-1) self.captured_attn_weights = attn_weights.detach() y = torch.matmul(attn_weights.to(dtype=v.dtype), v) return y.transpose(1, 2) else: y = F.scaled_dot_product_attention( q, k, v, attn_mask=mask, dropout_p=0.0, scale=scale, is_causal=mask is None ) return y.transpose(1, 2) def build_kv_cache( self, batch_size: int, max_seq_length: int, rope_cache_length: Optional[int] = None, device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, ) -> "KVCache": v_shape = (batch_size, self.config.n_query_groups, max_seq_length, self.config.head_size) if rope_cache_length is None: if self.config.rotary_percentage != 1.0: raise TypeError("Please pass the `rope_cache_length=gpt.cos.size(-1)` value") k_shape = v_shape else: k_shape = ( batch_size, self.config.n_query_groups, max_seq_length, rope_cache_length + self.config.head_size - self.config.rope_n_elem, ) return KVCache(k_shape, v_shape, device=device, dtype=dtype) def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None: """For compatibility with legacy checkpoints.""" for attr in ("weight", "bias"): legacy_key = f"{prefix}attn.{attr}" current_key = f"{prefix}qkv.{attr}" if legacy_key in state_dict: state_dict[current_key] = qkv_reassemble(state_dict.pop(legacy_key), self.config) super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) class GptNeoxMLP(nn.Module): def __init__(self, config: Config) -> None: super().__init__() self.fc = nn.Linear( config.n_embd, config.intermediate_size, bias=config.bias ) self.proj = nn.Linear( config.intermediate_size, config.n_embd, bias=config.bias ) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.fc(x) x = F.gelu(x, approximate=self.config.gelu_approximate) return self.proj(x) class LLaMAMLP(nn.Module): def __init__(self, config: Config) -> None: super().__init__() self.fc_1 = nn.Linear( config.n_embd, config.intermediate_size, bias=config.bias ) self.fc_2 = nn.Linear( config.n_embd, config.intermediate_size, bias=config.bias ) self.proj = nn.Linear( config.intermediate_size, config.n_embd, bias=config.bias ) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: x_fc_1 = self.fc_1(x) x_fc_2 = self.fc_2(x) x = F.silu(x_fc_1) * x_fc_2 return self.proj(x) class GemmaMLP(LLaMAMLP): def forward(self, x: torch.Tensor) -> torch.Tensor: x_fc_1 = self.fc_1(x) x_fc_2 = self.fc_2(x) x = F.gelu(x_fc_1, approximate=self.config.gelu_approximate) * x_fc_2 return self.proj(x) class LLaMAMoE(nn.Module): def __init__(self, config: Config) -> None: super().__init__() self.gate = nn.Linear(config.n_embd, config.n_expert, bias=False) self.experts = nn.ModuleList(LLaMAMLP(config) for _ in range(config.n_expert)) self.config = config def forward(self, x: torch.Tensor) -> torch.Tensor: """ Derived from: https://github.com/mistralai/mistral-src/blob/b46d6/moe_one_file_ref.py#L203-L219 See also figure 1 in https://arxiv.org/abs/2211.15841 """ B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd) x = x.view(-1, C) # (B*T, C) router = self.gate(x) # (B*T, n_expert) probs, indices = torch.topk(router, self.config.n_expert_per_token) # (B*T, n_expert_per_token) probs = probs.softmax(dim=1, dtype=torch.float).to(dtype=x.dtype) masks = indices.unsqueeze(-1) == torch.arange(self.config.n_expert, device=x.device) masks = masks.permute(2, 0, 1) # (n_expert, B*T, n_expert_per_token) y = torch.zeros_like(x) # (B*T, C) for mask, expert in zip(masks, self.experts): token_idx, expert_idx = torch.where(mask) y[token_idx] += probs[token_idx, expert_idx, None] * expert(x[token_idx]) return y.view(B, T, C) def build_rope_cache( seq_len: int, n_elem: int, device: Optional[torch.device] = None, base: int = 10000, condense_ratio: int = 1, extra_config: Optional[dict] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Enhanced Transformer with Rotary Position Embedding. Args: seq_len (int): Sequence length. n_elem (int): Number of elements (head dimension). device (torch.device, optional): Device for tensor allocations. base (int, optional): Base for computing inverse frequencies. condense_ratio (int, optional): Ratio to condense the position indices. extra_config (dict, optional): Configuration parameters for frequency adjustments (used by Llama 3.1 and 3.2) Returns: Tuple[torch.Tensor, torch.Tensor]: Cosine and sine caches for RoPE. Shapes are `(seq_len, n_elem)`. """ # Compute the inverse frequencies theta theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, device=device).float() / n_elem)) if extra_config is not None: orig_context_len = extra_config["original_max_seq_len"] factor = extra_config["factor"] low_freq_factor = extra_config["low_freq_factor"] high_freq_factor = extra_config["high_freq_factor"] wavelen = 2 * torch.pi / theta ratio = orig_context_len / wavelen smooth_factor = (ratio - low_freq_factor) / (high_freq_factor - low_freq_factor) smooth_factor = torch.clamp(smooth_factor, min=0.0, max=1.0) # Compute adjusted_theta without masked indexing adjusted_theta = (1 - smooth_factor) * (theta / factor) + smooth_factor * theta theta = adjusted_theta # Create position indices `[0, 1, ..., seq_len - 1]` ### Zhifei fix bug 1:. seq_idx = torch.arange(seq_len, device=device, dtype=torch.float16) / float(condense_ratio) # seq_idx = torch.arange(seq_len, device=device) / condense_ratio # Calculate the product of position index and $\theta_i$ idx_theta = torch.outer(seq_idx, theta).repeat(1, 2) # If `n_elem` is odd, the final dimension of `idx_theta` has size # `n_elem + 1`, so need to cut something off. # Due to a current bug in Hugging Face, in the case `n_elem == 1`, we leave # `idx_theta`, `cos`, `sin` as is. Things work out in `apply_rope` due to # broadcasting. If we shorten `idx_theta`, unit tests comparing to # Hugging Face fail. # https://github.com/huggingface/transformers/issues/35233 if idx_theta.shape[-1] > n_elem > 1: idx_theta = idx_theta[..., :n_elem] return torch.cos(idx_theta), torch.sin(idx_theta) def batched_index_select(t, dim, idx): """index_select for batched index and unbatched t""" if idx.dim() == 1: return torch.index_select(t, dim, idx) *batch_shape, idx_size = idx.shape res = torch.index_select(t, dim, idx.reshape(-1)) # flat index # split out single batch idx res = res.view(*t.shape[:dim], -1, idx_size, *t.shape[dim + 1 :]) if dim > 0: # move batch dim to front, this is np.rollaxis(res, dim, 0) for tensors dims = [dim] + list(range(res.dim())) del dims[dim + 1] res = res.permute(dims) # unflatten batch dims res = res.view(*batch_shape, *res.shape[1:]) return res def batched_index_copy_(t, dim, idx, val): """Index copy for batched t, idx, val""" if t.device.type == "mps": # Normalize negative dimensions if dim < 0: dim = t.dim() + dim if idx.dim() == 1: idx_shape = [1] * val.dim() idx_shape[dim] = -1 idx_expanded = idx.view(*idx_shape) idx_expanded = idx_expanded.expand_as(val) t.scatter_(dim, idx_expanded, val) return t elif idx.dim() == 2: assert dim != 0, "Cannot index the batch dimension" batch_size = idx.size(0) idx_size = idx.size(1) assert batch_size == t.size(0) == val.size(0) idx_shape = [batch_size] + [1] * (val.dim() - 1) idx_shape[dim] = idx_size idx_expanded = idx.view(*idx_shape) idx_expanded = idx_expanded.expand_as(val) t.scatter_(dim, idx_expanded, val) return t else: raise NotImplementedError(f"idx.dim() == {idx.dim()} not supported") else: if idx.dim() == 1: return t.index_copy_(dim, idx, val) assert idx.dim() == 2, f"multiple batch dims not yet {idx.shape=}" assert dim != 0, f"cannot index batch dim {dim=}" batch_size, idx_size = idx.shape assert batch_size == t.size(0) assert batch_size == val.size(0) # if we can view the batch and indexed dimensions together, we could # do index trickery. This is, sadly, not the case for kvcache so we # fall back to for loop for i in range(batch_size): unbatched_dim = dim if dim < 0 else dim - 1 t[i].index_copy_(unbatched_dim, idx[i], val[i]) return t def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """ Applies RoPE transform to `x`. Note that `cos`, `sin` need to have a batch dimension. Args: x: Input tensor, `(B, ..., T, head_size)` cos: Cached cosines, `(B, T, head_size)` or `(1, T, head_size)` sin: Cached sines, `(B, T, head_size)` or `(1, T, head_size)` Returns: Encoded tensor, `(B, ..., T, head_size)` """ if cos.dim() != 3: raise ValueError(f"cos must be three-dimensional, but shape is {cos.shape}") if cos.shape != sin.shape: raise ValueError(f"cos, sin must have same shape, but cos.shape={cos.shape}, sin.shape={sin.shape}") head_size_half = x.size(-1) // 2 x1 = x[..., : head_size_half] # (B, ..., T, head_size/2) x2 = x[..., head_size_half :] # (B, ..., T, head_size/2) rotated = torch.cat((-x2, x1), dim=-1) # (B, ..., T, head_size) dims_diff = x.dim() - cos.dim() if dims_diff > 0: # Ensure that shapes of `x`, `cos`, `sin` align new_shape = cos.shape[0:1] + (1,) * dims_diff + cos.shape[1:] cos = cos.view(*new_shape) sin = sin.view(*new_shape) roped = (x * cos) + (rotated * sin) return roped.to(dtype=x.dtype) def do_softcapping(x: torch.Tensor, thresh: float) -> torch.Tensor: return torch.tanh(x / thresh) * thresh class KVCache(nn.Module): """ Buffers `k`, `v` have shape `(batch_size, n_query_groups, max_seq_length, head_size)`. """ def __init__( self, k_shape: Tuple[int, int, int, int], v_shape: Tuple[int, int, int, int], device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, ) -> None: super().__init__() self.register_buffer("k", torch.zeros(k_shape, device=device, dtype=dtype), persistent=False) self.register_buffer("v", torch.zeros(v_shape, device=device, dtype=dtype), persistent=False) def forward(self, input_pos: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: """ Writes new values `k` and `v` into the cache at the positions specified by `input_pos` along the sequence dimension (`max_seq_length`). The batch size of `k` and `v` (`bs`) must be smaller or equal to `KVCache` batch size. Returns the full buffers, adjusted to the batch size `bs`. Args: input_pos: Position index, `(bs, T)` or `(T,)` k: New values, `(bs, n_query_groups, T, head_size)` v: New values, `(bs, n_query_groups, T, head_size)` Returns: k_full, v_full, `(bs, n_query_groups, max_seq_length, head_size)` """ # move the buffer to the activation dtype for when AMP is used self.k = self.k.to(k.dtype) self.v = self.v.to(v.dtype) # update the cache bs = k.size(0) k = batched_index_copy_(self.k[:bs, ...], -2, input_pos, k) v = batched_index_copy_(self.v[:bs, ...], -2, input_pos, v) return k, v def reset_parameters(self) -> None: torch.nn.init.zeros_(self.k) torch.nn.init.zeros_(self.v) def build_mask_cache(max_seq_length: int, device: Optional[torch.device] = None) -> torch.Tensor: ones = torch.ones((max_seq_length, max_seq_length), device=device, dtype=torch.bool) return torch.tril(ones).unsqueeze(0).unsqueeze(0) class RMSNorm(torch.nn.Module): """Root Mean Square Layer Normalization. Derived from https://github.com/bzhangGo/rmsnorm/blob/master/rmsnorm_torch.py. BSD 3-Clause License: https://github.com/bzhangGo/rmsnorm/blob/master/LICENSE. """ def __init__(self, size: int, dim: int = -1, eps: float = 1e-6, add_unit_offset: bool = False) -> None: super().__init__() self.weight = torch.nn.Parameter(torch.ones(size)) self.eps = eps self.dim = dim self.add_unit_offset = add_unit_offset def forward(self, x: torch.Tensor) -> torch.Tensor: dtype = x.dtype x = x.float() # NOTE: the original RMSNorm paper implementation is not equivalent norm_x = torch.mean(x * x, dim=self.dim, keepdim=True) x_normed = x * torch.rsqrt(norm_x + self.eps) weight = (1 + self.weight) if self.add_unit_offset else self.weight return (x_normed * weight.float()).to(dtype=dtype) def reset_parameters(self) -> None: torch.nn.init.ones_(self.weight)