CoolFace
Modelpublic

llm-slice/pico-decoder-medium

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes18downloads
pico_decoder.py609 linesDownload Raw Back to root
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