CoolFace
Modelpublic

PhysiQuanty/Binary-LLM-POC

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
12likes164downloads
modeling_binaryllm.py163 linesDownload Raw Back to root
1import math2from dataclasses import dataclass3from typing import Optional4 5import torch6import torch.nn as nn7import torch.nn.functional as F8 9from transformers import PreTrainedModel10from transformers.modeling_outputs import CausalLMOutput11 12from .configuration_binaryllm import BinaryLLMConfig13 14 15class PositionalEncoding(nn.Module):16    """17    Sinusoidal positional encoding, stocké en fp32,18    puis casté au dtype de x à chaque forward.19    """20 21    def __init__(self, d_model: int, max_len: int) -> None:22        super().__init__()23        pe = torch.zeros(max_len, d_model, dtype=torch.float32)24        position = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1)25        div_term = torch.exp(26            torch.arange(0, d_model, 2, dtype=torch.float32) * (-torch.log(torch.tensor(10000.0)) / d_model)27        )28        pe[:, 0::2] = torch.sin(position * div_term)29        pe[:, 1::2] = torch.cos(position * div_term)30        pe = pe.unsqueeze(0)  # (1, max_len, d_model)31        self.register_buffer("pe", pe, persistent=False)32 33    def forward(self, x: torch.Tensor) -> torch.Tensor:34        t = x.size(1)35        pe = self.pe[:, :t, :]36        pe = pe.to(device=x.device, dtype=x.dtype)37        return x + pe38 39 40@dataclass41class _InnerCfg:42    block_size: int43    embed_dim: int44    vocab_size: int45    num_heads: int46    num_layers: int47    ff_hidden_dim: int48    dropout: float49    layernorm_dim: Optional[int] = None50    head_dim: Optional[int] = None51 52 53class TinyTransformerLM(nn.Module):54    def __init__(self, cfg: _InnerCfg) -> None:55        super().__init__()56        self.cfg = cfg57 58        vocab_size = cfg.vocab_size59        self.tok_embed = nn.Embedding(vocab_size, cfg.embed_dim)60        self.pos_encoding = PositionalEncoding(cfg.embed_dim, cfg.block_size)61 62        encoder_layer = nn.TransformerEncoderLayer(63            d_model=cfg.embed_dim,64            nhead=cfg.num_heads,65            dim_feedforward=cfg.ff_hidden_dim,66            dropout=cfg.dropout,67            activation="gelu",68            batch_first=True,69        )70        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=cfg.num_layers)71 72        ln_dim = cfg.layernorm_dim or cfg.embed_dim73        head_dim = cfg.head_dim or ln_dim74 75        self.pre_ln_proj: Optional[nn.Linear] = None76        if ln_dim != cfg.embed_dim:77            self.pre_ln_proj = nn.Linear(cfg.embed_dim, ln_dim)78 79        self.ln = nn.LayerNorm(ln_dim)80 81        self.head_pre: Optional[nn.Linear] = None82        if head_dim != ln_dim:83            self.head_pre = nn.Linear(ln_dim, head_dim)84 85        self.head = nn.Linear(head_dim, vocab_size, bias=False)86 87        # weight tying seulement si parfait alignement88        if self.pre_ln_proj is None and self.head_pre is None and head_dim == cfg.embed_dim:89            self.head.weight = self.tok_embed.weight90 91        causal = torch.triu(torch.ones(cfg.block_size, cfg.block_size, dtype=torch.bool), diagonal=1)92        self.register_buffer("causal_mask", causal, persistent=False)93 94    def forward(self, tokens: torch.Tensor, padding_mask: Optional[torch.Tensor] = None) -> torch.Tensor:95        x = self.tok_embed(tokens)96        x = self.pos_encoding(x)97 98        seq_len = tokens.size(1)99        attn_mask = self.causal_mask[:seq_len, :seq_len].to(device=tokens.device)100 101        if padding_mask is not None:102            padding_mask = padding_mask[:, :seq_len].to(device=tokens.device, dtype=torch.bool)103 104        x = self.encoder(x, mask=attn_mask, src_key_padding_mask=padding_mask)105 106        if self.pre_ln_proj is not None:107            x = self.pre_ln_proj(x)108 109        x = self.ln(x)110 111        if self.head_pre is not None:112            x = self.head_pre(x)113 114        return self.head(x)115 116 117class BinaryLLMForCausalLM(PreTrainedModel):118    config_class = BinaryLLMConfig119    main_input_name = "input_ids"120 121    def __init__(self, config: BinaryLLMConfig):122        super().__init__(config)123 124        inner = _InnerCfg(125            block_size=int(config.max_position_embeddings),126            embed_dim=int(config.hidden_size),127            vocab_size=int(config.vocab_size),128            num_heads=int(config.num_attention_heads),129            num_layers=int(config.num_hidden_layers),130            ff_hidden_dim=int(config.intermediate_size),131            dropout=float(getattr(config, "dropout", 0.0)),132            layernorm_dim=None,133            head_dim=None,134        )135        self.model = TinyTransformerLM(inner)136 137        self.post_init()138 139    def forward(140        self,141        input_ids: torch.LongTensor,142        attention_mask: Optional[torch.Tensor] = None,143        labels: Optional[torch.LongTensor] = None,144        **kwargs,145    ) -> CausalLMOutput:146        padding_mask = None147        if attention_mask is not None:148            padding_mask = ~attention_mask.to(torch.bool)  # True = ignore149 150        logits = self.model(input_ids, padding_mask=padding_mask)151 152        loss = None153        if labels is not None:154            shift_logits = logits[:, :-1, :].contiguous()155            shift_labels = labels[:, 1:].contiguous()156            loss = F.cross_entropy(157                shift_logits.view(-1, self.config.vocab_size),158                shift_labels.view(-1),159                ignore_index=-100,160            )161 162        return CausalLMOutput(loss=loss, logits=logits)163