Sol-Lite-2 / sol2_core.py
j0no12's picture
Upload final Sol Lite 2 with Sol2ForCausalLM and banner
53d1439 verified
Raw History Blame Contribute Delete
14.6 kB
"""Unified SOL Lite 2 decoder. SPAB affects the first physical block only.
TN-Gram follows the factorized causal lookup in Sol Nano, with the short
sequence bug fixed and unused parameters removed. This is an experimental
architecture; no component here is a measured winner yet.
"""
from __future__ import annotations
from dataclasses import asdict, dataclass, replace
import math
import torch
from torch import nn
from torch.nn import functional as F
@dataclass(frozen=True)
class Config:
vocab_size: int = 4096
width: int = 256
heads: int = 8
kv_heads: int = 4
blocks: int = 10
ffn_width: int = 1536
recurrent_start: int = 1
recurrent_blocks: int = 4
passes: int = 2
loop_conditioning: bool = True
rope_theta: float = 20000.0
max_context: int = 512
xsa: bool = False
memory: str = "none"
tn_rank: int = 16
tn_buckets: int = 4096
tn_orders: tuple[int, ...] = (2, 3, 4, 5)
tn_gate: str = "static"
spab: str = "none"
spab_buckets: int = 262144
spab_scale_init: float = 0.0
spab_scale_trainable: bool = True
residual_scale: float = 1.0
backend: str = "reference"
def __post_init__(self):
if self.width % self.heads or self.heads % self.kv_heads:
raise ValueError("width/heads and heads/KV heads must divide evenly")
if (self.width // self.heads) % 2:
raise ValueError("RoPE needs an even head dimension")
if not (0 <= self.recurrent_start <= self.blocks and
0 <= self.recurrent_blocks <= self.blocks-self.recurrent_start):
raise ValueError("invalid recurrent range")
if self.passes < 1 or self.memory not in ("none", "tn", "engram"):
raise ValueError("invalid passes/memory")
if self.spab not in ("none", "pmi", "shuffled", "zero"):
raise ValueError("invalid SPAB control")
if self.backend not in ("reference", "sdpa", "flex"):
raise ValueError("invalid attention backend")
if self.tn_gate not in ("static", "context") or not self.tn_orders or min(self.tn_orders)<2:
raise ValueError("invalid TN-Gram configuration")
if min(self.vocab_size,self.width,self.blocks,self.ffn_width,self.kv_heads,self.max_context)<1:
raise ValueError("dimensions must be positive")
def to_dict(self):
return asdict(self)
class RMSNorm(nn.Module):
def __init__(self, width):
super().__init__()
self.weight = nn.Parameter(torch.ones(width))
def forward(self, x):
return (x.float() * torch.rsqrt(x.float().square().mean(-1,keepdim=True)+1e-6)).to(x.dtype)*self.weight.to(x.dtype)
def rope(x, theta):
length, dim = x.shape[-2:]
angles = torch.arange(length,device=x.device,dtype=torch.float32)[:,None]*theta**(-torch.arange(0,dim,2,device=x.device,dtype=torch.float32)/dim)
cos, sin = angles.cos().to(x.dtype), angles.sin().to(x.dtype)
even, odd = x[...,::2], x[...,1::2]
return torch.stack((even*cos-odd*sin,even*sin+odd*cos),-1).flatten(-2)
def pair_hash(query_ids, key_ids, buckets):
# Ordered pairs. Identical function is used in the corpus builder.
return ((query_ids.to(torch.int64)*1000003) ^ (key_ids.to(torch.int64)*9176+97)) % buckets
class SPAB(nn.Module):
def __init__(self, cfg, table=None):
super().__init__()
if cfg.spab in ("pmi", "shuffled") and table is None:
raise ValueError("PMI SPAB requires a frozen training-only table")
values = torch.zeros(cfg.spab_buckets) if table is None else torch.as_tensor(table,dtype=torch.float32).clone()
if values.shape != (cfg.spab_buckets,) or (values.device.type != "meta" and not torch.isfinite(values).all()):
raise ValueError("invalid SPAB table")
if cfg.spab == "zero":
values.zero_()
if cfg.spab == "shuffled":
values = values[torch.randperm(len(values),generator=torch.Generator().manual_seed(20260930))]
self.register_buffer("table",values,persistent=True)
self.scale = nn.Parameter(torch.full((cfg.heads,),cfg.spab_scale_init),requires_grad=cfg.spab_scale_trainable)
def bias(self, ids):
values = self.table[pair_hash(ids[:,:,None],ids[:,None,:],len(self.table))]
return values[:,None]*self.scale[None,:,None,None]
class TNGram(nn.Module):
def __init__(self, cfg):
super().__init__()
self.orders, self.buckets = cfg.tn_orders, cfg.tn_buckets
self.token_factors = nn.Parameter(torch.empty(cfg.vocab_size,cfg.tn_rank))
self.hash_tables = nn.Parameter(torch.empty(len(self.orders),cfg.tn_buckets,cfg.tn_rank))
self.order_factors = nn.Parameter(torch.ones(len(self.orders),cfg.tn_rank))
self.projection = nn.Linear(cfg.tn_rank,cfg.width,bias=False)
self.gates = nn.Parameter(torch.full((len(self.orders),),-2.0))
self.context_gate = nn.Linear(cfg.width,len(self.orders)) if cfg.tn_gate=="context" else None
nn.init.normal_(self.token_factors,mean=1.0,std=0.02)
nn.init.normal_(self.hash_tables,std=0.02)
nn.init.normal_(self.projection.weight,std=0.02)
if self.context_gate is not None:
nn.init.zeros_(self.context_gate.weight)
nn.init.zeros_(self.context_gate.bias)
def forward(self, ids, hidden):
batch,length=ids.shape
positions=torch.arange(length,device=ids.device)
result=hidden
contextual=self.context_gate(hidden) if self.context_gate is not None else None
for index,order in enumerate(self.orders):
hashed=torch.zeros_like(ids)
factors=torch.ones(batch,length,self.token_factors.shape[1],device=ids.device,dtype=torch.float32)
for lag in range(order):
# Slicing after padding keeps shape even when length < lag.
shifted=F.pad(ids,(lag,0),value=0)[:,:length]
hashed=(hashed*(131+index*6)+shifted)%self.buckets
factors=factors*self.token_factors[shifted].float()
rank=factors*self.hash_tables[index,hashed].float()*self.order_factors[index].float()
contribution=F.linear(rank,self.projection.weight.float()).to(hidden.dtype)
gate=torch.sigmoid(self.gates[index]+(contextual[...,index:index+1] if contextual is not None else 0))
valid=(positions>=order-1)[None,:,None].to(hidden.dtype)
result=result+contribution*gate*valid
return result
class Engram(nn.Module):
def __init__(self,cfg):
super().__init__()
self.tables=nn.ModuleList([nn.Embedding(cfg.tn_buckets,cfg.width) for _ in range(2)])
self.gate=nn.Linear(cfg.width,2)
self.scale=nn.Parameter(torch.tensor(0.1))
self.buckets=cfg.tn_buckets
for table in self.tables:
nn.init.normal_(table.weight,std=0.02)
def forward(self,ids,hidden):
memory=torch.zeros_like(hidden)
gates=torch.sigmoid(self.gate(hidden))
for index,order in enumerate((2,3)):
hashed=torch.zeros_like(ids)
for lag in range(order):
shifted=F.pad(ids,(lag,0),value=0)[:,:ids.shape[1]]
hashed=(hashed*(10007+2*index)+shifted)%self.buckets
valid=(torch.arange(ids.shape[1],device=ids.device)>=order-1)[None,:,None]
memory=memory+self.tables[index](hashed)*gates[...,index:index+1]*valid
return hidden+self.scale.tanh()*memory
_flex_kernel = None
def flex_kernel():
global _flex_kernel
if _flex_kernel is None:
from torch.nn.attention.flex_attention import flex_attention
_flex_kernel=torch.compile(flex_attention,dynamic=False)
return _flex_kernel
class Attention(nn.Module):
def __init__(self,cfg,first=False,table=None):
super().__init__()
self.cfg=cfg
dim=cfg.width//cfg.heads
self.q=nn.Linear(cfg.width,cfg.width,bias=False)
self.k=nn.Linear(cfg.width,cfg.kv_heads*dim,bias=False)
self.v=nn.Linear(cfg.width,cfg.kv_heads*dim,bias=False)
self.o=nn.Linear(cfg.width,cfg.width,bias=False)
self.q_norm,self.k_norm=RMSNorm(dim),RMSNorm(dim)
self.spab=SPAB(cfg,table) if first and cfg.spab!="none" else None
def forward(self,x,ids):
batch,length,_=x.shape
cfg=self.cfg; dim=cfg.width//cfg.heads
q=rope(self.q_norm(self.q(x).view(batch,length,cfg.heads,dim).transpose(1,2)),cfg.rope_theta)
k=rope(self.k_norm(self.k(x).view(batch,length,cfg.kv_heads,dim).transpose(1,2)),cfg.rope_theta)
v=self.v(x).view(batch,length,cfg.kv_heads,dim).transpose(1,2)
if cfg.backend=="flex":
if x.device.type!="cuda":
raise RuntimeError("Flex training backend requires CUDA; choose reference or sdpa for CPU/MPS")
prior=self.spab
if prior is not None:
table,scale=prior.table,prior.scale
def score_mod(score,b,h,qpos,kpos):
hashed=pair_hash(ids[b,qpos],ids[b,kpos],table.shape[0])
return torch.where(qpos>=kpos,score+table[hashed]*scale[h],-float("inf"))
else:
def score_mod(score,b,h,qpos,kpos):
return torch.where(qpos>=kpos,score,-float("inf"))
attended=flex_kernel()(q,k,v,score_mod=score_mod,enable_gqa=True)
else:
repeated_k=k.repeat_interleave(cfg.heads//cfg.kv_heads,dim=1)
repeated_v=v.repeat_interleave(cfg.heads//cfg.kv_heads,dim=1)
mask=torch.ones(length,length,device=x.device,dtype=torch.bool).tril()
bias=self.spab.bias(ids).to(q.dtype) if self.spab is not None else None
if cfg.backend=="sdpa":
additive=torch.zeros(length,length,device=x.device,dtype=q.dtype).masked_fill(~mask,-float("inf"))
if bias is not None: additive=additive+bias
attended=F.scaled_dot_product_attention(q,repeated_k,repeated_v,attn_mask=additive)
else:
scores=(q.float()@repeated_k.float().transpose(-2,-1))/math.sqrt(dim)
if bias is not None: scores=scores+bias.float()
attended=scores.masked_fill(~mask,-float("inf")).softmax(-1).to(v.dtype)@repeated_v
if cfg.xsa:
unit=F.normalize(v.repeat_interleave(cfg.heads//cfg.kv_heads,dim=1).float(),dim=-1,eps=1e-6).to(attended.dtype)
attended=attended-(attended*unit).sum(-1,keepdim=True)*unit
return self.o(attended.transpose(1,2).reshape(batch,length,cfg.width))
class Block(nn.Module):
def __init__(self,cfg,first=False,table=None):
super().__init__()
self.attn_norm,self.ffn_norm=RMSNorm(cfg.width),RMSNorm(cfg.width)
self.attn=Attention(cfg,first,table)
self.gate=nn.Linear(cfg.width,cfg.ffn_width,bias=False)
self.up=nn.Linear(cfg.width,cfg.ffn_width,bias=False)
self.down=nn.Linear(cfg.ffn_width,cfg.width,bias=False)
self.scale=cfg.residual_scale
def forward(self,x,ids):
x=x+self.scale*self.attn(self.attn_norm(x),ids)
normalized=self.ffn_norm(x)
return x+self.scale*self.down(F.silu(self.gate(normalized))*self.up(normalized))
class SolLite2(nn.Module):
def __init__(self,cfg=Config(),spab_table=None):
super().__init__()
self.config=cfg
self.embedding=nn.Embedding(cfg.vocab_size,cfg.width)
self.memory=TNGram(cfg) if cfg.memory=="tn" else Engram(cfg) if cfg.memory=="engram" else None
self.blocks=nn.ModuleList([Block(cfg,index==0,spab_table) for index in range(cfg.blocks)])
self.norm=RMSNorm(cfg.width)
if cfg.loop_conditioning and cfg.recurrent_blocks:
self.loop_embeddings=nn.Parameter(torch.empty(cfg.passes,cfg.width))
self.loop_gates=nn.Parameter(torch.zeros(cfg.passes,cfg.recurrent_blocks,cfg.width))
nn.init.normal_(self.loop_embeddings,std=0.01)
else:
self.register_parameter("loop_embeddings",None)
self.register_parameter("loop_gates",None)
nn.init.normal_(self.embedding.weight,std=0.02)
for block in self.blocks:
for layer in (block.attn.q,block.attn.k,block.attn.v,block.gate,block.up):
nn.init.normal_(layer.weight,std=0.02)
for layer in (block.attn.o,block.down):
nn.init.normal_(layer.weight,std=0.02/math.sqrt(2*cfg.blocks))
def forward(self,ids):
if ids.ndim!=2 or ids.shape[1]>self.config.max_context:
raise ValueError("expected [batch,sequence] inside configured context")
x=self.embedding(ids)
if self.memory is not None: x=self.memory(ids,x)
start=self.config.recurrent_start; stop=start+self.config.recurrent_blocks
for block in self.blocks[:start]: x=block(x,ids)
for index in range(self.config.passes):
if self.loop_embeddings is not None: x=x+self.loop_embeddings[index]
for local,block in enumerate(self.blocks[start:stop]):
proposal=block(x,ids)
x=x+torch.sigmoid(self.loop_gates[index,local])*(proposal-x) if self.loop_gates is not None else proposal
for block in self.blocks[stop:]: x=block(x,ids)
return F.linear(self.norm(x),self.embedding.weight)
# Public unified architecture name requested by the user.
def parameter_counts(model):
return {"trainable":sum(p.numel() for p in model.parameters() if p.requires_grad),
"parameters":sum(p.numel() for p in model.parameters()),
"frozen_buffer_bytes":sum(b.numel()*b.element_size() for b in model.buffers())}
def fit_budget(cfg,target):
# Construct on meta to audit sizes without allocating large weights.
with torch.device("meta"):
neutral=replace(cfg,ffn_width=1,spab="zero" if cfg.spab!="none" else "none")
fixed=parameter_counts(SolLite2(neutral))["parameters"]-3*cfg.blocks*cfg.width
ffn=max(8,((target-fixed)//(3*cfg.blocks*cfg.width)//8)*8)
if fixed+3*cfg.blocks*cfg.width*ffn>target:
raise ValueError("target too small for fixed architecture")
return replace(cfg,ffn_width=ffn)
PRESETS={"nano":(128,4,2,2900000),"lite":(256,8,4,15000000),
"flash":(480,15,5,50000000),"pro":(768,12,4,138000000)}
def preset(name,memory="none",spab="none",backend="reference"):
width,heads,kv,target=PRESETS[name]
return fit_budget(Config(width=width,heads=heads,kv_heads=kv,memory=memory,spab=spab,backend=backend),target)