llm-slice/pico-decoder-medium
018
1"""2Pico Decoder: A Lightweight Causal Transformer Language Model3 4Pico Decoder uses a simple LLAMA-style transformer architecture, written for clarity and educational purposes.5 6Everything is written with a modular design for easy modification and experimentation.7 8Key features:9- RMSNorm for layer normalization10- Rotary Positional Embeddings (RoPE)11- Multi-head attention with KV-cache support12- SwiGLU activation function13- Residual connections throughout14 15- KV-cache for faster autoregressive generation16 17References:18 - RoPE: https://arxiv.org/abs/2104.0986419 - SwiGLU: https://arxiv.org/abs/2002.0520220 - LLAMA: https://arxiv.org/abs/2302.1397121 22Adapted from:23 - OLMO: https://github.com/allenai/OLMo24 - LLAMA: https://github.com/meta/llama25"""26 27from dataclasses import asdict28from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union29 30import torch31import torch.nn as nn32import torch.nn.functional as F33from torch.nn.attention import SDPBackend, sdpa_kernel34from transformers import PretrainedConfig, PreTrainedModel35from transformers.modeling_outputs import CausalLMOutput, CausalLMOutputWithPast36 37try:38 if TYPE_CHECKING:39 # We need to do this to avoid importing these when creating the HF-compatible models40 from src.config import ModelConfig41except ImportError:42 pass43 44########################################################45#46# Layer Normalization47#48########################################################49 50 51class RMSNorm(torch.nn.Module):52 """Root Mean Square Layer Normalization.53 54 A variant of Layer Normalization that uses RMS statistics instead of mean/variance,55 resulting in improved stability and performance.56 57 Args:58 config (Union[ModelConfig, PicoHFConfig]): Configuration object containing normalization parameters59 - config.norm_eps: Small constant for numerical stability60 - config.d_model: Model dimension for the weight parameter61 62 References:63 https://arxiv.org/abs/1910.0746764 """65 66 def __init__(self, config: Union["ModelConfig", "PicoDecoderHFConfig"]):67 super().__init__()68 self.eps = config.norm_eps69 self.weight = nn.Parameter(torch.ones(config.d_model))70 71 def _norm(self, x: torch.Tensor) -> torch.Tensor:72 """73 Normalizes the input tensor by its RMS value.74 """75 return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)76 77 def forward(self, x: torch.Tensor) -> torch.Tensor:78 """79 Applies RMS normalization to the input tensor and scales it by the weight parameter.80 """81 output = self._norm(x.float()).type_as(x)82 return output * self.weight83 84 85########################################################86#87# Positional Embedding88#89########################################################90 91 92class RoPE(nn.Module):93 """Rotary Positional Embeddings (RoPE).94 95 Implements position-dependent rotation of keys and queries in attention mechanism,96 allowing better modeling of relative positions in sequences. Uses complex number97 operations for efficient rotation.98 99 Args:100 config (Union[ModelConfig, PicoHFConfig]): Model configuration containing:101 - config.position_emb_theta: Base for frequency computation102 - config.d_model: Model dimension103 - config.attention_n_heads: Number of attention heads104 - config.max_seq_len: Maximum sequence length105 106 References:107 https://arxiv.org/abs/2104.09864108 """109 110 _freqs_cis_tensor: torch.Tensor | None = None111 112 def __init__(self, config: Union["ModelConfig", "PicoDecoderHFConfig"]):113 super().__init__()114 115 self.theta = config.position_emb_theta116 self.dim = config.d_model // config.attention_n_heads117 118 max_seq_len = config.max_seq_len119 120 # only gets set once, and then reused for all RoPE instances121 if RoPE._freqs_cis_tensor is None:122 RoPE._freqs_cis_tensor = self._setup_freqs_cis(123 max_seq_len, self.theta, self.dim124 )125 126 # register _freqs_cis buffer127 # can be easily recomputed so persistent=False128 self.register_buffer("_freqs_cis", self._freqs_cis_tensor, persistent=False)129 130 @classmethod131 def _setup_freqs_cis(cls, seq_len: int, theta: float, dim: int) -> torch.Tensor:132 """Setup Frequency Tensor for RoPE Embeddings133 134 Initializes the complex frequency tensor that is used to compute the RoPE embeddings.135 136 Note other implementations will use cos and sin directly, but using the complex137 number representation is (probably?) more efficient:138 139 e^(theta * i * t) = cos(theta * t) + i * sin(theta * t) [Euler's formula]140 """141 _freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))142 positions = torch.arange(seq_len)143 freqs = torch.outer(positions, _freqs)144 return torch.polar(torch.ones_like(freqs), freqs) # complex64145 146 def get_freqs_cis(147 self, input_shape: torch.Size, start_pos: int, end_pos: int148 ) -> torch.Tensor:149 """Reshape Frequency Tensor for RoPE Embeddings150 151 Makes the frequency tensor broadcastable with the input tensor.152 """153 _freqs_cis = self._freqs_cis[start_pos:end_pos]154 ndim = len(input_shape)155 assert 0 <= 1 < ndim156 assert _freqs_cis.shape == (input_shape[1], input_shape[-1])157 158 # TODO: Check whether this is correct (might be able to remove this)159 shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(input_shape)]160 return _freqs_cis.view(*shape)161 162 def forward(163 self,164 queries: torch.Tensor,165 keys: torch.Tensor,166 start_pos: int = 0,167 ) -> Tuple[torch.Tensor, torch.Tensor]:168 """Apply RoPE Embeddings to Queries and Keys169 170 Applies the rotary positional embeddings to the input tensors via complex num multiplication171 172 NOTE: The start_pos is used if we want to use the kv_cache in the attention mechanism.173 """174 queries_ = torch.view_as_complex(175 queries.float().reshape(*queries.shape[:-1], -1, 2)176 )177 keys_ = torch.view_as_complex(keys.float().reshape(*keys.shape[:-1], -1, 2))178 179 input_shape = (180 queries_.shape181 ) # same as keys: (batch_size, seq_len, n_heads, head_dim/2)182 freqs_start_pos = start_pos183 freqs_end_pos = freqs_start_pos + queries_.shape[1]184 185 freqs_cis = self.get_freqs_cis(input_shape, freqs_start_pos, freqs_end_pos)186 187 queries_rotated = torch.view_as_real(queries_ * freqs_cis).flatten(3)188 keys_rotated = torch.view_as_real(keys_ * freqs_cis).flatten(3)189 return queries_rotated.type_as(queries), keys_rotated.type_as(keys)190 191 192########################################################193#194# Attention195#196########################################################197 198 199class Attention(nn.Module):200 """Multi-head Attention with Group Query Attention support.201 202 Implements scaled dot-product attention and supports:203 - Grouped Query Attention (GQA)204 - Key-Value caching for efficient inference205 - RoPE integration206 207 Args:208 config (Union[ModelConfig, PretrainedConfig]): Configuration containing:209 - config.attention_n_heads: Number of attention heads210 - config.attention_n_kv_heads: Number of key/value heads211 - config.d_model: Model dimension212 - config.batch_size: Maximum batch size213 - config.max_seq_len: Maximum sequence length214 215 Shape:216 - Input: (batch_size, seq_len, d_model)217 - Output: (batch_size, seq_len, d_model)218 """219 220 def __init__(221 self,222 config: Union["ModelConfig", "PicoDecoderHFConfig"],223 ):224 super().__init__()225 226 self.n_heads = config.attention_n_heads227 self.n_kv_heads = config.attention_n_kv_heads228 229 self.batch_size = config.batch_size230 self.max_seq_len = config.max_seq_len231 232 d_model = config.d_model233 self.head_dim = d_model // self.n_heads234 235 self.n_rep = self.n_heads // self.n_kv_heads236 237 self.q_proj = nn.Linear(d_model, self.n_heads * self.head_dim, bias=False)238 self.k_proj = nn.Linear(d_model, self.n_kv_heads * self.head_dim, bias=False)239 self.v_proj = nn.Linear(d_model, self.n_kv_heads * self.head_dim, bias=False)240 self.o_proj = nn.Linear(self.n_heads * self.head_dim, d_model, bias=False)241 242 self.rope = RoPE(config)243 244 def forward(245 self,246 input: torch.Tensor,247 mask: Optional[torch.Tensor] = None,248 past_key_values: Optional[Tuple[torch.Tensor, ...]] = None,249 use_cache: bool = False,250 ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:251 """Forward pass for the attention mechanism.252 253 Computes queries, keys, and values for the attention mechanism. Applies rotary positional254 embeddings to the queries and keys, and then computes attention scores and outputs.255 256 For an introduction to the attention mechanism, see:257 https://arxiv.org/abs/1706.03762258 259 A few things to note:260 - The past_key_values is used to implement the KV cache, which is used to speed up261 generation by caching the KV pairs from previous forward passes. This is useful when doing262 tasks that require generating multiple tokens conditioned on previous tokens (e.g. language263 modeling, text generation, etc.). The way the KV cache is implemented is that each layer has264 its own KV cache - this KV cache is implemented as a tuple.265 """266 bsz, seq_len, _ = input.shape267 _queries, _keys, _values = (268 self.q_proj(input),269 self.k_proj(input),270 self.v_proj(input),271 )272 273 # Reshaping for multi-head attention274 queries = _queries.view(bsz, seq_len, self.n_heads, self.head_dim)275 keys = _keys.view(bsz, seq_len, self.n_kv_heads, self.head_dim)276 values = _values.view(bsz, seq_len, self.n_kv_heads, self.head_dim)277 278 # The start position is used to apply the RoPE embeddings to only the new tokens279 # when using the kv_cache in the attention mechanism.280 # We want to start from the last position in the cache.281 start_pos = past_key_values[0].shape[1] if past_key_values is not None else 0282 283 # apply rotary positional embeddings284 queries, keys = self.rope(queries, keys, start_pos)285 286 if past_key_values is not None:287 keys = torch.cat([past_key_values[0], keys], dim=1)288 values = torch.cat([past_key_values[1], values], dim=1)289 290 if use_cache:291 cached_keys = keys292 cached_values = values293 else:294 cached_keys = None295 cached_values = None296 297 queries = queries.transpose(1, 2)298 keys = keys.transpose(1, 2)299 values = values.transpose(1, 2)300 301 apply_gqa = self.n_rep > 1302 if apply_gqa and queries.device.type == "mps":303 # NOTE: MPS does not support GQA in the SDPA kernel, but we can repeat the keys and values304 # outside of the kernel to get the same effect.305 # See: https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html306 keys = keys.repeat_interleave(self.n_rep, dim=-3)307 values = values.repeat_interleave(self.n_rep, dim=-3)308 apply_gqa = False309 310 backends = [SDPBackend.CUDNN_ATTENTION, SDPBackend.MATH]311 312 with sdpa_kernel(backends=backends):313 attn_output = F.scaled_dot_product_attention(314 queries.contiguous(),315 keys.contiguous(),316 values.contiguous(),317 attn_mask=mask.to(queries.dtype),318 enable_gqa=apply_gqa,319 )320 321 attn_output = attn_output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)322 output = self.o_proj(attn_output)323 324 return output, (cached_keys, cached_values)325 326 327########################################################328#329# SwiGLU (Combines MLP and Activation)330#331########################################################332 333 334class SwiGLU(nn.Module):335 """SwiGLU Activation Function with Linear Projections.336 337 Implements the SwiGLU activation function combined with linear transformations,338 serving as the feed-forward network in transformer blocks.339 340 Args:341 config (Union[ModelConfig, PicoDecoderHFConfig]): Configuration containing:342 - config.d_model: Model dimension343 - config.activation_hidden_dim: Hidden dimension (typically 4 * d_model)344 345 References:346 https://arxiv.org/abs/2002.05202347 """348 349 def __init__(self, config: Union["ModelConfig", "PicoDecoderHFConfig"]):350 super().__init__()351 352 model_dim = config.d_model353 act_hidden_dim = config.activation_hidden_dim # usually 4 * d_model354 355 self.w_0 = nn.Linear(model_dim, act_hidden_dim, bias=False)356 self.w_1 = nn.Linear(model_dim, act_hidden_dim, bias=False)357 self.w_2 = nn.Linear(act_hidden_dim, model_dim, bias=False)358 359 def forward(self, x: torch.Tensor) -> torch.Tensor:360 return self.w_2(F.silu(self.w_0(x)) * self.w_1(x))361 362 363########################################################364#365# PicoDecoderBlock366#367########################################################368 369 370class PicoDecoderBlock(nn.Module):371 """Single Transformer Block with Attention and Feed-forward layers.372 373 Implements a standard transformer block with:374 - Multi-head attention with normalization and residual connection375 - SwiGLU feed-forward network with normalization and residual connection376 377 Args:378 config (Union[ModelConfig, PicoDecoderHFConfig]): Model configuration; either a dataclass or379 a HuggingFace PicoDecoderHFConfig380 """381 382 def __init__(383 self,384 config: Union["ModelConfig", "PicoDecoderHFConfig"],385 ):386 super().__init__()387 388 self.attention = Attention(config)389 self.swiglu = SwiGLU(config)390 self.attention_norm = RMSNorm(config)391 self.swiglu_norm = RMSNorm(config)392 393 def forward(394 self,395 input: torch.Tensor,396 mask: Optional[torch.Tensor] = None,397 past_key_values: Optional[Tuple[torch.Tensor]] = None,398 use_cache: bool = False,399 ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:400 attention_output, cached_key_values = self.attention(401 self.attention_norm(input),402 mask=mask,403 past_key_values=past_key_values,404 use_cache=use_cache,405 )406 # NOTE: cached_key_values is None if use_cache is False407 408 h = input + attention_output409 out = h + self.swiglu(self.swiglu_norm(h))410 return out, cached_key_values411 412 413########################################################414#415# Pico Decoder (Causal Transformer Model)416#417########################################################418 419 420class PicoDecoder(nn.Module):421 """422 Pico Decoder: combines the embedding, causal decoder blocks, and output projection into a423 single autoregressive model.424 425 For more information on the model, see the classes for the modules that make up the model.426 """427 428 def __init__(429 self,430 model_config: Union["ModelConfig", "PicoDecoderHFConfig"],431 ):432 super().__init__()433 self.config = model_config434 435 self.embedding_proj = nn.Embedding(self.config.vocab_size, self.config.d_model)436 self.layers = nn.ModuleList(437 [PicoDecoderBlock(self.config) for _ in range(self.config.n_layers)]438 )439 self.output_norm = RMSNorm(self.config)440 self.de_embedding_proj = nn.Linear(441 self.config.d_model, self.config.vocab_size, bias=False442 )443 444 def convert_to_hf_model(self) -> "PicoDecoderHF":445 """Convert the Lightning model to a HuggingFace model."""446 # Create HF config without fabric-specific settings447 hf_config = PicoDecoderHFConfig.from_dataclass(self.config)448 449 # Create new HF model450 hf_model = PicoDecoderHF(hf_config)451 452 # Copy state dict, excluding fabric-specific keys453 hf_model.load_state_dict(self.state_dict(prefix="pico_decoder."))454 455 return hf_model456 457 def forward(458 self,459 input_ids: torch.Tensor,460 past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None,461 use_cache: bool = False,462 ) -> Tuple[torch.Tensor, Optional[Tuple[Tuple[torch.Tensor, torch.Tensor]]]]:463 """464 This is the forward pass for the entire Pico model. It boils down to:465 - Embedding the input ids466 - Creating a causal mask467 - Processing through the pico layers468 - Projecting the output to logits469 470 NOTE: One feature that might be confusing is the KV cache. The KV cache is used to speed up471 generation by caching the KV pairs from previous forward passes. This is useful when doing472 tasks that require generating multiple tokens conditioned on previous tokens (e.g. language473 modeling, text generation, etc.). The way the KV cache is implemented is that each layer has474 its own KV cache which is stored as a tuple. The whole model then stores a tuple of these475 KV caches (so a tuple of tuples).476 """477 478 seq_len = input_ids.shape[-1]479 h = self.embedding_proj(input_ids)480 481 # Calculate start position from past cached KV pairs. Remember that each layer has its482 # own KV Cache. So when we index past_key_values, we need to index into the KV pairs for the483 # correct layer and then for either the keys or values.484 start_pos = 0 if past_key_values is None else past_key_values[0][0].shape[1]485 486 # Create causal mask for current sequence487 mask = None488 if seq_len > 1:489 mask = torch.full((seq_len, seq_len), float("-inf"))490 mask = torch.triu(mask, diagonal=1)491 492 # If using KV cache, extend mask to cover cached sequence length493 if past_key_values is not None:494 # Add zeros for cached tokens (we can attend to all of them)495 mask = torch.hstack([torch.zeros((seq_len, start_pos)), mask])496 497 mask = mask.to(h.device)498 499 # NOTE: If we are using the cache, we need to store the cached KV pairs for each layer500 # in a tuple. Each layer will have its own cached KV pair which we aggregate in a tuple.501 cached_key_values = () if use_cache else None502 503 # Process through transformer blocks504 for idx, layer in enumerate(self.layers):505 layer_past_key_values = (506 past_key_values[idx] if past_key_values is not None else None507 )508 509 h, layer_cached_key_values = layer(510 h, mask=mask, past_key_values=layer_past_key_values, use_cache=use_cache511 )512 513 if use_cache:514 cached_key_values += (layer_cached_key_values,)515 516 # Final norm and projection517 h = self.output_norm(h)518 logits = self.de_embedding_proj(h).float()519 520 return logits, cached_key_values521 522 523########################################################524#525# HuggingFace Wrapper for the Pico Decoder model.526#527########################################################528 529 530class PicoDecoderHFConfig(PretrainedConfig):531 """Config class for the Pico Decoder HuggingFace wrapper."""532 533 model_type = "pico_decoder"534 535 @classmethod536 def from_dict(cls, config_dict: Dict[str, Any], **kwargs) -> "PicoDecoderHFConfig":537 """538 Initialize config from a dictionary. Note that no kwargs are passed to the constructor --539 this is because with some kwargs special handling is required and can make this class540 brittle.541 """542 pico_config = cls(**config_dict)543 544 return_unused_kwargs = kwargs.pop("return_unused_kwargs", False)545 unused_kwargs = {546 key: value for key, value in kwargs.items() if not hasattr(pico_config, key)547 }548 549 if return_unused_kwargs:550 return pico_config, unused_kwargs551 return pico_config552 553 @classmethod554 def from_dataclass(cls, model_config: "ModelConfig"):555 """Initialise from our custom config dataclass."""556 return cls.from_dict(asdict(model_config))557 558 559class PicoDecoderHF(PreTrainedModel):560 """561 HuggingFace wrapper for the Pico model.562 563 Many evaluation frameworks require a model be setup as a HuggingFace model, so we provide a simple564 wrapper that does just that. When we save checkpoints of the Pico model, we save both the normal565 Pico model as well as the model wrapped in this HuggingFace class.566 567 This also lets you do cool things like:568 569 `model = AutoModelForCausalLM.from_pretrained("path/to/checkpoint")`570 """571 572 config_class = PicoDecoderHFConfig573 _no_split_modules = ["PicoBlock", "Attention", "SwiGLU", "RMSNorm"]574 575 def __init__(self, config: PicoDecoderHFConfig):576 super().__init__(config)577 self.pico_decoder = PicoDecoder(config)578 579 def forward(580 self,581 input_ids: torch.Tensor,582 past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None,583 use_cache: bool = False,584 **kwargs,585 ) -> Union[CausalLMOutput, CausalLMOutputWithPast]:586 """HuggingFace forward pass wrapper.587 588 Forwards pass for the HuggingFace version of the Pico Model. Basic wrapper around the589 Pico model's forward pass, and returns the output as a HuggingFace CausalLMOutput.590 """591 logits, past_key_values = self.pico_decoder(592 input_ids, past_key_values, use_cache593 )594 if use_cache:595 return CausalLMOutputWithPast(596 logits=logits,597 past_key_values=past_key_values,598 )599 else:600 return CausalLMOutput(601 logits=logits,602 )603 604 605# Register for auto classes606PicoDecoderHFConfig.register_for_auto_class()607PicoDecoderHF.register_for_auto_class("AutoModel")608PicoDecoderHF.register_for_auto_class("AutoModelForCausalLM")609 