"""Tokle: decoder-only transformer (RMSNorm, RoPE, GQA, SwiGLU) plus SPAB. Math is unchanged from the training-time model.py -- same norms, same half-split RoPE, same hash -- so an exported checkpoint scores identically. Only the module names were changed to the usual Transformers layout. Works on Transformers 4.x and 5.x. The v5 differences are marked inline; they are not cosmetic, see the note above _M1. """ import math import os from typing import Optional, Tuple, Union import torch import torch.nn as nn import torch.nn.functional as F from transformers import __version__ as _transformers_version from transformers.generation import GenerationMixin from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput from transformers.modeling_utils import PreTrainedModel try: # as a Transformers dynamic module (trust_remote_code) from .configuration_tokle import TokleConfig except ImportError: # running the files straight out of a checkout from configuration_tokle import TokleConfig # ----------------------------------------------------------------------------- norms and rope class TokleRMSNorm(nn.Module): def __init__(self, d, eps=1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(d)) def forward(self, x): # Always reduce in fp32: in fp16 the mean of squares underflows on the # long tail and the norm comes back slightly wrong. dt = x.dtype x = x.float() x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return (x * self.weight.float()).to(dt) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.eps}" def build_rope(head_dim, max_seq_len, theta, device=None): inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim)) t = torch.arange(max_seq_len, dtype=torch.float32, device=device) f = torch.outer(t, inv) return torch.cos(f), torch.sin(f) def apply_rope(x, cos, sin): """Half-split rotation (first half against second), not the interleaved variant. Training used this one; swapping it silently degrades the model.""" b, h, t, d = x.shape x1, x2 = x.float().chunk(2, dim=-1) c = cos[:t].view(1, 1, t, d // 2).to(x1.device) s = sin[:t].view(1, 1, t, d // 2).to(x1.device) return torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], dim=-1).to(x.dtype) # ----------------------------------------------------------------------------- spab # The table lives outside model.safetensors. It is frozen data, not weights, and # keeping it separate stops it being counted as trained parameters. Cost of that # choice: it has to be fetched and installed by hand after the weights load, so # every load path below goes through _install_spab_table. SPAB_TABLE_FILE = "spab_table.safetensors" # Mixing constants for the pair hash. Plain ints, not buffers, and that matters: # buffers here would be non-persistent, and Transformers 5 builds the model on # the meta device and fills only what the checkpoint contains. Non-persistent # buffers come back as uninitialised memory -- in practice zeros, which sends # every token pair to slot 0 and quietly flattens the whole bias. Loads fine, # generates garbage. Same reason the rope tables below are not buffers either. _M1 = -7046029254386353131 _M2 = -4417276706812531889 _M3 = -4658895280553007687 class SPABBias(nn.Module): """Static Pairwise Attention Bias, keyed by (source id, target id). Hashing the pair into one flat table costs a single lookup and keeps storage fixed -- a real vocab^2 bias matrix would be 25M entries here and would grow with the vocabulary. The table holds windowed PMI from the training corpus and never moves; `scale` (one per head) is the only trained tensor. Position plays no part in the lookup, which is what sets SPAB apart from ALiBi or T5-style biases. """ is_spab = True def __init__(self, config: TokleConfig): super().__init__() self.enabled = config.spab_enabled self.table_size = config.spab_table_size self.n_head = config.n_head init = config.spab_init_scale if config.spab_enabled else 0.0 self.scale = nn.Parameter(torch.full((config.n_head,), init)) if not config.spab_enabled: self.scale.requires_grad_(False) # persistent=False keeps it out of state_dict and therefore out of # model.safetensors; `loaded` guards against using the zeros. self.register_buffer("table", torch.zeros(config.spab_table_size, dtype=torch.float32), persistent=False) self.loaded = not config.spab_enabled def hash(self, i, j): # Murmur-style mix. Relies on int64 overflow wrapping, so keep it in # int64 throughout -- the masks after each shift are what make the # result reproducible across devices. h = i.to(torch.int64) * _M1 + j.to(torch.int64) * _M2 h = h ^ ((h >> 29) & 0x7FFFFFFFF) h = h * _M3 h = h ^ ((h >> 32) & 0xFFFFFFFF) return (h & 0x7FFFFFFFFFFFFFFF) % self.table_size def forward(self, idx): if not self.enabled: return None if not self.loaded: raise RuntimeError( f"the SPAB table was never loaded -- {SPAB_TABLE_FILE} is missing " f"or was not installed. Load the model with from_pretrained(), " f"which fetches it; an all-zero table silently ruins every output.") # Materialises b*t*t int64 pairs, so this is the memory high-water mark # of a forward pass. Fine at t=512; watch it if max_seq_len ever grows. b, t = idx.shape src = idx.unsqueeze(1).expand(b, t, t) dst = idx.unsqueeze(2).expand(b, t, t) vals = self.table[self.hash(src, dst)] return self.scale.view(1, -1, 1, 1) * vals.unsqueeze(1) # ----------------------------------------------------------------------------- blocks class TokleAttention(nn.Module): def __init__(self, config: TokleConfig, layer_idx: int): super().__init__() self.layer_idx = layer_idx self.n_head = config.n_head self.n_kv_head = config.n_kv_head self.head_dim = config.head_dim self.rep = config.n_head // config.n_kv_head self.q_proj = nn.Linear(config.d_model, config.n_head * config.head_dim, bias=False) self.k_proj = nn.Linear(config.d_model, config.n_kv_head * config.head_dim, bias=False) self.v_proj = nn.Linear(config.d_model, config.n_kv_head * config.head_dim, bias=False) self.o_proj = nn.Linear(config.n_head * config.head_dim, config.d_model, bias=False) self.q_norm = TokleRMSNorm(config.head_dim, config.norm_eps) self.k_norm = TokleRMSNorm(config.head_dim, config.norm_eps) def forward(self, x, cos, sin, attn_mask, spab_bias): b, t, _ = x.shape q = self.q_proj(x).view(b, t, self.n_head, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2) # QK-norm before rope, as trained. q = apply_rope(self.q_norm(q), cos, sin) k = apply_rope(self.k_norm(k), cos, sin) k = k.repeat_interleave(self.rep, dim=1) v = v.repeat_interleave(self.rep, dim=1) if spab_bias is None and attn_mask is None: # Nothing to add: let SDPA build the causal mask itself, which is # the fast kernel path. This is what layers 1..8 normally hit. o = F.scaled_dot_product_attention(q, k, v, is_causal=True) else: mask = attn_mask if spab_bias is not None: causal = torch.ones(t, t, dtype=torch.bool, device=x.device).triu(1) bias = spab_bias.masked_fill(causal, float("-inf")) mask = bias if mask is None else bias + mask o = F.scaled_dot_product_attention(q, k, v, attn_mask=mask.to(q.dtype)) return self.o_proj(o.transpose(1, 2).reshape(b, t, -1)) class TokleMLP(nn.Module): def __init__(self, config: TokleConfig): super().__init__() self.gate_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=False) self.up_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=False) self.down_proj = nn.Linear(config.ffn_hidden, config.d_model, bias=False) def forward(self, x): return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class TokleDecoderLayer(nn.Module): def __init__(self, config: TokleConfig, layer_idx: int): super().__init__() self.input_layernorm = TokleRMSNorm(config.d_model, config.norm_eps) self.self_attn = TokleAttention(config, layer_idx) self.post_attention_layernorm = TokleRMSNorm(config.d_model, config.norm_eps) self.mlp = TokleMLP(config) def forward(self, x, cos, sin, attn_mask, spab_bias): x = x + self.self_attn(self.input_layernorm(x), cos, sin, attn_mask, spab_bias) return x + self.mlp(self.post_attention_layernorm(x)) # ----------------------------------------------------------------------------- model class ToklePreTrainedModel(PreTrainedModel): config_class = TokleConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["TokleDecoderLayer"] _supports_sdpa = True def _spab(self): for m in self.modules(): if getattr(m, "is_spab", False): return m return None def install_spab_table(self, table): """Copy the frozen table in. Accepts a tensor or a .safetensors path.""" spab = self._spab() if spab is None: return self if isinstance(table, (str, os.PathLike)): from safetensors.torch import load_file table = load_file(str(table))["table"] table = table.reshape(-1) if table.numel() != spab.table_size: raise ValueError(f"SPAB table has {table.numel()} entries, " f"expected {spab.table_size}") # Match the existing buffer's device and dtype so a model already moved # to cuda/fp16 keeps behaving the way it did before the split. ref = spab.table spab.table = table.to(device=ref.device, dtype=ref.dtype) spab.loaded = True return self @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs): model = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs) if getattr(model.config, "spab_enabled", False): model.install_spab_table( _resolve_table(pretrained_model_name_or_path, **kwargs)) return model def save_pretrained(self, save_directory, *args, **kwargs): out = super().save_pretrained(save_directory, *args, **kwargs) spab = self._spab() if spab is not None and spab.enabled and spab.loaded: from safetensors.torch import save_file save_file({"table": spab.table.detach().cpu().to(torch.float32).contiguous()}, os.path.join(save_directory, SPAB_TABLE_FILE)) return out def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, 0.0, 0.02) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, 0.0, 0.02) elif isinstance(module, TokleRMSNorm): nn.init.ones_(module.weight) elif isinstance(module, SPABBias): nn.init.constant_( module.scale, self.config.spab_init_scale if self.config.spab_enabled else 0.0) # Residual-output projections get scaled down by depth so the residual # stream does not blow up at init (GPT-2 trick). if isinstance(module, (TokleAttention, TokleMLP)): out = module.o_proj if isinstance(module, TokleAttention) else module.down_proj nn.init.normal_(out.weight, 0.0, 0.02 / math.sqrt(2 * self.config.n_layer)) def _resolve_table(repo, **kwargs): """Find spab_table.safetensors in a local dir, the HF cache, or the Hub.""" local = os.path.join(str(repo), SPAB_TABLE_FILE) if os.path.isfile(local): return local try: from transformers.utils import cached_file return cached_file(str(repo), SPAB_TABLE_FILE, **{k: kwargs[k] for k in ("revision", "token", "cache_dir", "local_files_only", "subfolder") if k in kwargs}) except Exception: pass try: from huggingface_hub import hf_hub_download return hf_hub_download(str(repo), SPAB_TABLE_FILE, revision=kwargs.get("revision"), token=kwargs.get("token")) except Exception as e: # Worth a clear message: without the table the model loads happily and # then produces nonsense, which is a miserable thing to debug. raise FileNotFoundError( f"{SPAB_TABLE_FILE} was not found in '{repo}'. Tokle ships its frozen " f"SPAB table beside model.safetensors and cannot run without it; copy " f"the whole repo, not just the weights. ({type(e).__name__}: {e})" ) from e class TokleModel(ToklePreTrainedModel): def __init__(self, config: TokleConfig): super().__init__(config) self.embed_tokens = nn.Embedding(config.vocab_size, config.d_model) self.spab = SPABBias(config) self.layers = nn.ModuleList( [TokleDecoderLayer(config, i) for i in range(config.n_layer)]) self.norm = TokleRMSNorm(config.d_model, config.norm_eps) # Plain dict, not a buffer -- see the _M1 note. Keyed by device so # .to()/.cuda() need not move anything. self._rope_cache = {} self.gradient_checkpointing = False self.post_init() def get_input_embeddings(self): return self.embed_tokens def set_input_embeddings(self, value): self.embed_tokens = value def rope(self, device): key = str(device) if key not in self._rope_cache: self._rope_cache[key] = build_rope( self.config.head_dim, self.config.max_seq_len, self.config.rope_theta, device=device) return self._rope_cache[key] def _padding_mask(self, attention_mask, dtype): """2D mask -> additive (b, 1, 1, t), or None when nothing is padded so the caller can take the is_causal fast path.""" if attention_mask is None: return None attention_mask = attention_mask.to(torch.bool) if bool(attention_mask.all()): return None m = torch.zeros(attention_mask.shape, dtype=dtype, device=attention_mask.device) m = m.masked_fill(~attention_mask, torch.finfo(dtype).min) return m[:, None, None, :] def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[Tuple, BaseModelOutput]: # getattr, not config.use_return_dict: that property is deprecated in # Transformers 5 and warns on every call. return_dict = (return_dict if return_dict is not None else getattr(self.config, "return_dict", True)) output_hidden_states = (output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states) if (input_ids is None) == (inputs_embeds is None): raise ValueError("Pass exactly one of input_ids or inputs_embeds.") if inputs_embeds is not None and self.config.spab_enabled: raise ValueError( "inputs_embeds is not supported while SPAB is enabled: the bias is " "indexed by token id, so input_ids are required.") x = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds t = x.shape[1] if t > self.config.max_seq_len: raise ValueError( f"Sequence length {t} exceeds max_seq_len {self.config.max_seq_len}.") cos, sin = self.rope(x.device) bias = self.spab(input_ids) if input_ids is not None else None # Fold the causal mask in here once rather than per layer. Under causal # attention with right padding, position 0 always sees itself, so no row # ends up fully masked (which would give NaNs out of softmax). pad = self._padding_mask(attention_mask, x.dtype) if pad is not None: causal = torch.ones(t, t, dtype=torch.bool, device=x.device).triu(1) pad = pad.masked_fill(causal, torch.finfo(x.dtype).min) hidden_states = () if output_hidden_states else None for i, layer in enumerate(self.layers): if output_hidden_states: hidden_states += (x,) layer_bias = bias if i == 0 else None # SPAB is layer 0 only if self.gradient_checkpointing and self.training: x = self._gradient_checkpointing_func( layer.__call__, x, cos, sin, pad, layer_bias) else: x = layer(x, cos, sin, pad, layer_bias) x = self.norm(x) if output_hidden_states: hidden_states += (x,) if not return_dict: return tuple(v for v in (x, hidden_states) if v is not None) return BaseModelOutput(last_hidden_state=x, hidden_states=hidden_states) # Transformers 5 turned _tied_weights_keys from a list of names into a # {tied: source} mapping and calls .keys() on it, so the old form raises. _TRANSFORMERS_V5 = int(_transformers_version.split(".")[0]) >= 5 class TokleForCausalLM(ToklePreTrainedModel, GenerationMixin): _tied_weights_keys = ({"lm_head.weight": "model.embed_tokens.weight"} if _TRANSFORMERS_V5 else ["lm_head.weight"]) def __init__(self, config: TokleConfig): super().__init__(config) self.model = TokleModel(config) self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def get_decoder(self): return self.model def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[Tuple, CausalLMOutput]: return_dict = (return_dict if return_dict is not None else getattr(self.config, "return_dict", True)) out = self.model( input_ids=input_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds, output_hidden_states=output_hidden_states, return_dict=True, ) logits = self.lm_head(out.last_hidden_state) # Rows past real_vocab_size only pad the matrix to a multiple of 64. # They never saw a gradient and no target ever points at them, so keep # them out of sampling and out of any log-likelihood the harness sums. rv = self.config.real_vocab_size if self.config.mask_padded_vocab_logits and rv is not None and rv < self.config.vocab_size: logits = logits.clone() logits[..., rv:] = float("-inf") loss = None if labels is not None: loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs) if not return_dict: return (loss, logits) if loss is not None else (logits,) return CausalLMOutput(loss=loss, logits=logits, hidden_states=out.hidden_states)