PhysiQuanty/Binary-LLM-POC
12164
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 