import math import os import re from dataclasses import dataclass from typing import Optional, Tuple, Any, Dict import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutput from huggingface_hub import hf_hub_download from safetensors.torch import safe_open from .configuration_binaryllm import BinaryLLMConfig # ============================================================ # Helpers: u16 <-> bytes # ============================================================ def split_u16_to_bytes(u16: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: hi = (u16 >> 8) & 0xFF lo = u16 & 0xFF return hi.long(), lo.long() def factorized_ce_u16( logits_hi: torch.Tensor, # (B,T,256) logits_lo: torch.Tensor, # (B,T,256) target_u16: torch.Tensor, # (B,T) long (0..65535) ou ignore_index ignore_index: int = -100, ) -> torch.Tensor: y = target_u16 y_safe = torch.clamp(y, min=0) y_hi, y_lo = split_u16_to_bytes(y_safe) y_hi[y == ignore_index] = ignore_index y_lo[y == ignore_index] = ignore_index B, T, V = logits_hi.shape l1 = F.cross_entropy(logits_hi.view(B * T, V), y_hi.view(B * T), ignore_index=ignore_index) l2 = F.cross_entropy(logits_lo.view(B * T, V), y_lo.view(B * T), ignore_index=ignore_index) return l1 + l2 # ============================================================ # Positional Encoding (dtype-safe) # ============================================================ class PositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int) -> None: super().__init__() pe = torch.zeros(max_len, d_model, dtype=torch.float32) position = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1) div_term = torch.exp( torch.arange(0, d_model, 2, dtype=torch.float32) * (-torch.log(torch.tensor(10000.0)) / d_model) ) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) pe = pe.unsqueeze(0) # (1, max_len, d_model) self.register_buffer("pe", pe, persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: t = x.size(1) pe = self.pe[:, :t, :].to(device=x.device, dtype=x.dtype) return x + pe # ============================================================ # Factorized head (2 x softmax 256) # ============================================================ class FactorizedU16Head(nn.Module): def __init__(self, d_model: int, byte_emb_dim: int = 64) -> None: super().__init__() self.proj_hi = nn.Linear(d_model, 256) self.hi_emb = nn.Embedding(256, byte_emb_dim) self.proj_lo = nn.Linear(d_model + byte_emb_dim, 256) def forward(self, h: torch.Tensor, hi_cond: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: logits_hi = self.proj_hi(h) # (B,T,256) cond = torch.cat([h, self.hi_emb(hi_cond)], dim=-1) logits_lo = self.proj_lo(cond) # (B,T,256) return logits_hi, logits_lo # ============================================================ # Inner config (kept minimal) # ============================================================ @dataclass class _InnerCfg: block_size: int embed_dim: int vocab_size: int num_heads: int num_layers: int ff_hidden_dim: int dropout: float ignore_index: int = -100 byte_emb_dim: int = 64 layernorm_dim: Optional[int] = None head_dim: Optional[int] = None # ============================================================ # TinyTransformerLM (factorized) # ============================================================ class TinyTransformerLM(nn.Module): def __init__(self, cfg: _InnerCfg) -> None: super().__init__() self.cfg = cfg self.vocab_size = int(cfg.vocab_size) self.ignore_index = int(cfg.ignore_index) self.tok_embed = nn.Embedding(self.vocab_size, cfg.embed_dim) self.pos_encoding = PositionalEncoding(cfg.embed_dim, cfg.block_size) encoder_layer = nn.TransformerEncoderLayer( d_model=cfg.embed_dim, nhead=cfg.num_heads, dim_feedforward=cfg.ff_hidden_dim, dropout=cfg.dropout, activation="gelu", batch_first=True, ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=cfg.num_layers) ln_dim = cfg.layernorm_dim or cfg.embed_dim head_dim = cfg.head_dim or ln_dim self.pre_ln_proj: Optional[nn.Linear] = None if ln_dim != cfg.embed_dim: self.pre_ln_proj = nn.Linear(cfg.embed_dim, ln_dim) self.ln = nn.LayerNorm(ln_dim) self.head_pre: Optional[nn.Linear] = None if head_dim != ln_dim: self.head_pre = nn.Linear(ln_dim, head_dim) self.head = FactorizedU16Head(head_dim, byte_emb_dim=int(cfg.byte_emb_dim)) causal = torch.triu(torch.ones(cfg.block_size, cfg.block_size, dtype=torch.bool), diagonal=1) self.register_buffer("causal_mask", causal, persistent=False) def forward( self, tokens: torch.Tensor, padding_mask: Optional[torch.Tensor] = None, labels: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: x = self.tok_embed(tokens) x = self.pos_encoding(x) seq_len = tokens.size(1) attn_mask = self.causal_mask[:seq_len, :seq_len].to(device=tokens.device) if padding_mask is not None: padding_mask = padding_mask[:, :seq_len].to(device=tokens.device, dtype=torch.bool) x = self.encoder(x, mask=attn_mask, src_key_padding_mask=padding_mask) if self.pre_ln_proj is not None: x = self.pre_ln_proj(x) x = self.ln(x) if self.head_pre is not None: x = self.head_pre(x) # logits_hi first logits_hi = self.head.proj_hi(x) # hi conditioning: # - training: teacher forcing from labels # - inference: use argmax(logits_hi) (deterministic) so forward works without labels if labels is not None: hi_cond, _ = split_u16_to_bytes(labels) else: hi_cond = torch.argmax(logits_hi, dim=-1).long() cond = torch.cat([x, self.head.hi_emb(hi_cond)], dim=-1) logits_lo = self.head.proj_lo(cond) return logits_hi, logits_lo def compute_loss( self, outputs: Tuple[torch.Tensor, torch.Tensor], targets: torch.Tensor, padding_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: if padding_mask is not None: t = targets.clone() t[padding_mask] = self.ignore_index else: t = targets logits_hi, logits_lo = outputs return factorized_ce_u16(logits_hi, logits_lo, t, ignore_index=self.ignore_index) # ============================================================ # Shape detection from safetensors (cache/local) # ============================================================ def _infer_arch_from_safetensors(path: str) -> Dict[str, int]: info: Dict[str, int] = {} with safe_open(path, framework="pt", device="cpu") as f: # vocab + hidden w = f.get_tensor("model.tok_embed.weight") info["vocab_size"] = int(w.shape[0]) info["hidden_size"] = int(w.shape[1]) # layers count layer_ids = [] rx = re.compile(r"^model\.encoder\.layers\.(\d+)\.") for k in f.keys(): m = rx.match(k) if m: layer_ids.append(int(m.group(1))) info["num_hidden_layers"] = (max(layer_ids) + 1) if layer_ids else 0 # intermediate size from first layer linear1.weight k_lin1 = "model.encoder.layers.0.linear1.weight" if k_lin1 in f.keys(): info["intermediate_size"] = int(f.get_tensor(k_lin1).shape[0]) # byte_emb_dim from hi_emb.weight k_hi = "model.head.hi_emb.weight" if k_hi in f.keys(): info["byte_emb_dim"] = int(f.get_tensor(k_hi).shape[1]) return info # ============================================================ # HF Wrapper model # ============================================================ class BinaryLLMForCausalLM(PreTrainedModel): config_class = BinaryLLMConfig main_input_name = "input_ids" @classmethod def from_pretrained(cls, pretrained_model_name_or_path: str, *model_args, **kwargs): # Load config first (then patch it using safetensors shapes) config = kwargs.get("config", None) if config is None: config = BinaryLLMConfig.from_pretrained(pretrained_model_name_or_path, **{k: v for k, v in kwargs.items() if k in ["cache_dir", "revision", "token"]}) kwargs["config"] = config # Locate safetensors file cache_dir = kwargs.get("cache_dir", None) revision = kwargs.get("revision", None) token = kwargs.get("token", None) try: st_path = hf_hub_download( repo_id=pretrained_model_name_or_path, filename="model.safetensors", revision=revision, token=token, cache_dir=cache_dir, ) except Exception: # local path fallback local = os.path.join(str(pretrained_model_name_or_path), "model.safetensors") st_path = local arch = _infer_arch_from_safetensors(st_path) # Patch config to MATCH checkpoint if "vocab_size" in arch: config.vocab_size = int(arch["vocab_size"]) if "hidden_size" in arch: config.hidden_size = int(arch["hidden_size"]) if "num_hidden_layers" in arch and int(arch["num_hidden_layers"]) > 0: config.num_hidden_layers = int(arch["num_hidden_layers"]) if "intermediate_size" in arch: config.intermediate_size = int(arch["intermediate_size"]) # custom field (safe even if config doesn't define it strictly) if "byte_emb_dim" in arch: setattr(config, "byte_emb_dim", int(arch["byte_emb_dim"])) kwargs["config"] = config return super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs) def __init__(self, config: BinaryLLMConfig): super().__init__(config) byte_emb_dim = int(getattr(config, "byte_emb_dim", 64)) inner = _InnerCfg( block_size=int(config.max_position_embeddings), embed_dim=int(config.hidden_size), vocab_size=int(config.vocab_size), num_heads=int(config.num_attention_heads), num_layers=int(config.num_hidden_layers), ff_hidden_dim=int(config.intermediate_size), dropout=float(getattr(config, "dropout", 0.0)), ignore_index=int(getattr(config, "ignore_index", -100)), byte_emb_dim=int(byte_emb_dim), layernorm_dim=None, head_dim=None, ) self.model = TinyTransformerLM(inner) self.post_init() def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.LongTensor] = None, **kwargs, ) -> CausalLMOutput: padding_mask = None if attention_mask is not None: padding_mask = ~attention_mask.to(torch.bool) logits_hi, logits_lo = self.model(input_ids, padding_mask=padding_mask, labels=labels) loss = None if labels is not None: loss = self.model.compute_loss((logits_hi, logits_lo), labels, padding_mask=padding_mask) out = CausalLMOutput(loss=loss, logits=logits_hi) # expose both for your factorized inference scripts out.logits_hi = logits_hi out.logits_lo = logits_lo return out