LibraMindMini / model.py
alby13's picture
Upload LibraMind Mini weights, tokenizer, and inference code
bcda938 verified
Raw History Blame Contribute Delete
12 kB
"""Hybrid language model: Gated DeltaNet + gated attention (3:1), Block Attention Residuals.
Architecture follows Qwen3.5 (3 Gated DeltaNet layers per gated full-attention layer,
sigmoid output gate on attention, QK-norm) with Kimi's Block Attention Residuals replacing
the plain residual stream. arch="dense" swaps every DeltaNet layer for gated attention,
which gives a like-for-like Transformer baseline.
"""
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
@dataclass
class ModelConfig:
vocab_size: int = 32768
n_layer: int = 24
d_model: int = 1024
arch: str = "hybrid" # "hybrid": every `attn_every`-th layer is attention, rest DeltaNet; "dense": all attention
attn_every: int = 4
n_head: int = 16 # attention query heads
n_kv_head: int = 4
head_dim: int = 128
gdn_heads: int = 16
gdn_head_dim: int = 128
ffn_dim: int = 3584
attnres_blocks: int = 8 # Block AttnRes: the 2*n_layer sublayers are split into this many blocks
rope_theta: float = 10000.0
max_seq_len: int = 2048
softcap: float = 15.0
grad_ckpt: bool = False # recompute sublayers in backward to save VRAM
PRESETS = {
# smoke tests only
"tiny": dict(n_layer=4, d_model=256, n_head=4, n_kv_head=2, head_dim=64, gdn_heads=4, gdn_head_dim=64,
ffn_dim=768, attnres_blocks=4),
# ~140M non-embedding: architecture bake-off size
"small": dict(n_layer=12, d_model=768, n_head=12, n_kv_head=3, head_dim=128, gdn_heads=12, gdn_head_dim=128,
ffn_dim=2688, attnres_blocks=8),
# ~500M non-embedding: Qwen3.5-0.8B layout with a 32k vocab
"base": dict(n_layer=24, d_model=1024, n_head=16, n_kv_head=4, head_dim=128, gdn_heads=16, gdn_head_dim=128,
ffn_dim=3584, attnres_blocks=8),
}
def rms(x):
return F.rms_norm(x, (x.size(-1),))
def inv_rms(x, eps=1e-6):
return torch.rsqrt(x.float().pow(2).mean(-1) + eps)
class GatedAttention(nn.Module):
"""Causal GQA attention with QK-norm, RoPE and Qwen's elementwise sigmoid output gate."""
def __init__(self, c: ModelConfig):
super().__init__()
self.nh, self.nkv, self.hd = c.n_head, c.n_kv_head, c.head_dim
self.q_proj = nn.Linear(c.d_model, 2 * self.nh * self.hd, bias=False) # query and gate
self.kv_proj = nn.Linear(c.d_model, 2 * self.nkv * self.hd, bias=False)
self.o_proj = nn.Linear(self.nh * self.hd, c.d_model, bias=False)
inv_freq = 1.0 / (c.rope_theta ** (torch.arange(0, self.hd, 2).float() / self.hd))
freqs = torch.outer(torch.arange(c.max_seq_len).float(), inv_freq)
self.register_buffer("cos", freqs.cos()[None, :, None, :], persistent=False)
self.register_buffer("sin", freqs.sin()[None, :, None, :], persistent=False)
def rope(self, x):
T = x.size(1)
cos, sin = self.cos[:, :T].to(x.dtype), self.sin[:, :T].to(x.dtype)
x1, x2 = x.chunk(2, -1)
return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], -1)
def forward(self, x):
B, T, _ = x.shape
q, gate = self.q_proj(x).split(self.nh * self.hd, -1)
q = q.view(B, T, self.nh, self.hd)
k, v = self.kv_proj(x).view(B, T, 2, self.nkv, self.hd).unbind(2)
q, k = self.rope(rms(q)), self.rope(rms(k))
if self.nkv != self.nh:
k = k.repeat_interleave(self.nh // self.nkv, dim=2)
v = v.repeat_interleave(self.nh // self.nkv, dim=2)
y = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=True)
y = y.transpose(1, 2).reshape(B, T, self.nh * self.hd)
return self.o_proj(y * torch.sigmoid(gate))
class DeltaNet(nn.Module):
"""Gated DeltaNet (linear attention with a delta-rule memory) from flash-linear-attention."""
def __init__(self, c: ModelConfig):
super().__init__()
from fla.layers import GatedDeltaNet
self.gdn = GatedDeltaNet(hidden_size=c.d_model, head_dim=c.gdn_head_dim, num_heads=c.gdn_heads,
expand_v=1.0, mode="chunk", use_gate=True, use_short_conv=True, conv_size=4)
def forward(self, x):
return self.gdn(x)[0]
class SwiGLU(nn.Module):
def __init__(self, c: ModelConfig):
super().__init__()
self.w_gate = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.w_up = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.w_down = nn.Linear(c.ffn_dim, c.d_model, bias=False)
def forward(self, x):
return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
def attnres_mix(sources, source_inv_rms, query):
"""Block AttnRes: softmax over sources of <query, RMSNorm(source)>, then a weighted sum of sources."""
logits = torch.stack([(s @ query.to(s.dtype)).float() * r for s, r in zip(sources, source_inv_rms)])
weights = logits.softmax(0)
h = weights[0, ..., None] * sources[0]
for w, s in zip(weights[1:], sources[1:]):
h = h + w[..., None] * s
return h
class LM(nn.Module):
def __init__(self, c: ModelConfig):
super().__init__()
self.config = c
self.embed = nn.Embedding(c.vocab_size, c.d_model)
self.sublayers = nn.ModuleList()
for i in range(c.n_layer):
is_attn = c.arch == "dense" or (i % c.attn_every == c.attn_every - 1)
self.sublayers.append(GatedAttention(c) if is_attn else DeltaNet(c))
self.sublayers.append(SwiGLU(c))
n_sub = len(self.sublayers)
self.block_size = math.ceil(n_sub / c.attnres_blocks)
# one pseudo-query per sublayer plus one for the final read-out; zero init = uniform average at start
self.attnres_queries = nn.Parameter(torch.zeros(n_sub + 1, c.d_model))
self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)
self.init_weights()
@torch.no_grad()
def init_weights(self):
c = self.config
nn.init.normal_(self.embed.weight, std=1.0)
nn.init.normal_(self.lm_head.weight, std=0.001)
s = 3 ** 0.5 * c.d_model ** -0.5
for m in self.sublayers:
if isinstance(m, GatedAttention):
nn.init.uniform_(m.q_proj.weight, -s, s)
nn.init.uniform_(m.kv_proj.weight, -s, s)
nn.init.zeros_(m.o_proj.weight)
elif isinstance(m, SwiGLU):
nn.init.uniform_(m.w_gate.weight, -s, s)
nn.init.uniform_(m.w_up.weight, -s, s)
nn.init.zeros_(m.w_down.weight)
else: # DeltaNet keeps FLA's init for its gates/decay; match the rest to the attention layers
g = m.gdn
for lin in (g.q_proj, g.k_proj, g.v_proj, g.g_proj):
nn.init.uniform_(lin.weight, -s, s)
nn.init.zeros_(g.o_proj.weight)
def num_params(self, non_embedding=True):
n = sum(p.numel() for p in self.parameters())
return n - self.embed.weight.numel() - self.lm_head.weight.numel() if non_embedding else n
def _run(self, sub, x):
if self.config.grad_ckpt and self.training:
return checkpoint(sub, x, use_reentrant=False)
return sub(x)
def hidden(self, idx):
"""Final normalized hidden states [B, T, d_model]."""
x0 = rms(self.embed(idx))
blocks, blocks_inv = [x0], [inv_rms(x0)] # completed block sums (block 0 = token embedding)
partial = None # running sum of sublayer outputs in the current block
for j, sub in enumerate(self.sublayers):
srcs = blocks if partial is None else blocks + [partial]
inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
h = attnres_mix(srcs, inv, self.attnres_queries[j])
out = self._run(sub, rms(h)).float()
partial = out if partial is None else partial + out
if (j + 1) % self.block_size == 0:
blocks.append(partial)
blocks_inv.append(inv_rms(partial))
partial = None
srcs = blocks if partial is None else blocks + [partial]
inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
return rms(attnres_mix(srcs, inv, self.attnres_queries[-1]))
def forward(self, idx, targets=None):
h = self.hidden(idx)
if targets is None:
return self._logits(h)
# Loss in chunks, recomputing each chunk's logits in backward, so the full
# (tokens x vocab) fp32 logit tensor never has to sit in VRAM.
# Targets of -1 (prompts, padding during fine-tuning) are ignored.
h, targets = h.flatten(0, 1), targets.flatten()
loss = 0.0
for i in range(0, h.size(0), self.LOSS_CHUNK):
loss = loss + checkpoint(self._chunk_loss, h[i:i + self.LOSS_CHUNK], targets[i:i + self.LOSS_CHUNK],
use_reentrant=False)
return loss / (targets >= 0).sum().clamp(min=1)
LOSS_CHUNK = 4096
def _logits(self, h):
logits = self.lm_head(h).float()
if self.config.softcap:
logits = self.config.softcap * torch.tanh(logits / self.config.softcap)
return logits
def _chunk_loss(self, h, targets):
return F.cross_entropy(self._logits(h), targets, reduction="sum", ignore_index=-1)
def token_logprobs(self, idx, targets):
"""Log-probability of each target token, [B, T] (0 where the target is -1). Chunked like the loss."""
h, t = self.hidden(idx).flatten(0, 1), targets.flatten()
parts = [checkpoint(self._chunk_logprobs, h[i:i + self.LOSS_CHUNK], t[i:i + self.LOSS_CHUNK], use_reentrant=False)
for i in range(0, h.size(0), self.LOSS_CHUNK)]
return torch.cat(parts).view(targets.shape)
def _chunk_logprobs(self, h, targets):
lp = torch.log_softmax(self._logits(h), -1).gather(1, targets.clamp(min=0)[:, None])[:, 0]
return lp * (targets >= 0)
def param_groups(self):
"""Split parameters: hidden matrices go to Muon, everything else to AdamW."""
muon, other = [], []
for name, p in self.sublayers.named_parameters():
(muon if p.ndim == 2 and min(p.shape) >= 64 else other).append(p)
other.append(self.attnres_queries)
return dict(muon=muon, embed=[self.embed.weight], head=[self.lm_head.weight], other=other)
def resize_vocab(state_dict, new_rows):
"""Grow the embedding and output matrices for added tokens; new rows start at the mean of the old ones."""
for key in ("embed.weight", "lm_head.weight"):
w = state_dict[key]
if w.size(0) < new_rows:
extra = w.float().mean(0, keepdim=True).expand(new_rows - w.size(0), -1).to(w.dtype)
state_dict[key] = torch.cat([w, extra])
return state_dict
@torch.no_grad()
def generate(model, idx, max_new_tokens, temperature=0.8, top_k=50):
"""Simple sampling that re-runs the full context each step (no cache; fine for short samples)."""
for _ in range(max_new_tokens):
ctx = idx[:, -model.config.max_seq_len:]
with torch.autocast("cuda", dtype=torch.bfloat16):
logits = model(ctx)[:, -1, :].float()
if temperature <= 0:
nxt = logits.argmax(-1, keepdim=True)
else:
logits = logits / temperature
if top_k:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = -float("inf")
nxt = torch.multinomial(F.softmax(logits, -1), 1)
idx = torch.cat([idx, nxt], 1)
return idx