cortex / bet_model.py
appvoid's picture
SparkBET-9M inference export at update 414181
20d009d verified
Raw
History Blame Contribute Delete
20.4 kB
"""SparkBET-9M: min-spark-style looped core with BET phase conditioning."""
import math, hashlib
from dataclasses import dataclass, asdict
import torch
from torch import nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
SEQ_LEN = 1024
MAX_LOOPS = 8
PAD_ID, BOS_ID, EOS_ID = 256, 257, 258
VOCAB_SIZE = 259
EXPECTED_PARAM_COUNT = 9_353_876
EXPECTED_ARCH_SHAPE_SHA256 = "c85415fc50d03a23f89ccd870ebf53a3a00c5a61017c13a28f6444fdd71b72b9"
_USE_GRAD_CHECKPOINTING = False
_STATE_NOISE_SIGMA = 0.0
def set_gradient_checkpointing(enabled):
global _USE_GRAD_CHECKPOINTING
_USE_GRAD_CHECKPOINTING = bool(enabled)
def set_state_noise_sigma(value):
global _STATE_NOISE_SIGMA
value = float(value)
if value < 0:
raise ValueError("state noise sigma must be >= 0")
_STATE_NOISE_SIGMA = value
@dataclass(frozen=True)
class BETConfig:
vocab_size: int = VOCAB_SIZE
hidden_size: int = 324
intermediate_size: int = 864
prelude_layers: int = 1
body_blocks: int = 6
coda_layers: int = 1
num_heads: int = 6
num_kv_heads: int = 2
head_dim: int = 54
lora_rank: int = 16
hyper_lanes: int = 2
max_seq_len: int = SEQ_LEN
max_loops: int = MAX_LOOPS
rope_theta: float = 10_000.0
rms_eps: float = 1e-6
ddl_beta_init: float = 1.0
ddl_k_eps: float = 1e-2
ddl_v_sigmoid_scale: float = 4.0
@property
def q_dim(self):
return self.num_heads * self.head_dim
@property
def kv_dim(self):
return self.num_kv_heads * self.head_dim
@property
def qkv_dim(self):
return self.q_dim + 2 * self.kv_dim
CFG = BETConfig()
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x):
dtype = x.dtype
y = x.float()
y = y * torch.rsqrt(y.pow(2).mean(-1, keepdim=True) + self.eps)
return (y * self.weight.float()).to(dtype)
def rope_cos_sin(position_ids, dim, theta, dtype):
inv = 1.0 / (theta ** (torch.arange(0, dim, 2, device=position_ids.device, dtype=torch.float32) / dim))
f = position_ids.float().unsqueeze(-1) * inv
return f.cos().unsqueeze(1).to(dtype), f.sin().unsqueeze(1).to(dtype)
def apply_rope(x, cos, sin):
# RoPE tables may be prepared before the first autocast linear, when the
# embedding stream is FP32. Cast them to the projected Q/K dtype here so
# attention remains FP16 on every CUDA profile instead of being promoted.
cos, sin = cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
xe, xo = x[..., 0::2], x[..., 1::2]
return torch.stack((xe * cos - xo * sin, xe * sin + xo * cos), dim=-1).flatten(-2)
def attention_mask_and_positions(input_ids, attention_mask=None):
b, t = input_ids.shape
if attention_mask is None:
pos = torch.arange(t, device=input_ids.device).view(1, t).expand(b, t)
return None, pos
mask = attention_mask.to(device=input_ids.device, dtype=torch.bool)
if mask.shape != input_ids.shape:
raise ValueError(f"attention_mask shape {tuple(mask.shape)} != input_ids {tuple(input_ids.shape)}")
pos = (mask.long().cumsum(-1) - 1).clamp_min(0)
if bool(mask.all()):
return None, pos
causal = torch.ones((t, t), device=input_ids.device, dtype=torch.bool).tril()[None, None]
allowed = causal & mask[:, None, None, :]
return allowed, pos
class Attention(nn.Module):
"""6Q/2KV GQA with per-head QK RMSNorm before RoPE."""
def __init__(self, c):
super().__init__()
if c.hidden_size != c.q_dim:
raise ValueError("hidden_size must equal num_heads * head_dim")
if c.num_heads % c.num_kv_heads:
raise ValueError("num_heads must be divisible by num_kv_heads")
self.qkv = nn.Linear(c.hidden_size, c.qkv_dim, bias=False)
self.out = nn.Linear(c.q_dim, c.hidden_size, bias=False)
self.qn = RMSNorm(c.head_dim, c.rms_eps)
self.kn = RMSNorm(c.head_dim, c.rms_eps)
self.nh, self.nkv, self.hd = c.num_heads, c.num_kv_heads, c.head_dim
self.q_dim, self.kv_dim = c.q_dim, c.kv_dim
def forward(self, x, cos, sin, qkv_delta=None, attn_mask=None):
b, t, _ = x.shape
qkv = self.qkv(x)
if qkv_delta is not None:
if qkv_delta.shape != qkv.shape:
raise RuntimeError("phase LoRA QKV delta shape mismatch")
qkv = qkv + qkv_delta
q, k, v = qkv.split([self.q_dim, self.kv_dim, self.kv_dim], dim=-1)
q = q.view(b, t, self.nh, self.hd).transpose(1, 2)
k = k.view(b, t, self.nkv, self.hd).transpose(1, 2)
v = v.view(b, t, self.nkv, self.hd).transpose(1, 2)
q = apply_rope(self.qn(q), cos, sin)
k = apply_rope(self.kn(k), cos, sin)
if self.nkv != self.nh:
repeat = self.nh // self.nkv
k = k.repeat_interleave(repeat, dim=1)
v = v.repeat_interleave(repeat, dim=1)
if attn_mask is None:
z = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=0.0)
else:
z = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0)
return self.out(z.transpose(1, 2).contiguous().view(b, t, self.q_dim))
class SwiGLU(nn.Module):
def __init__(self, c):
super().__init__()
self.gate_up = nn.Linear(c.hidden_size, 2 * c.intermediate_size, bias=False)
self.down = nn.Linear(c.intermediate_size, c.hidden_size, bias=False)
def forward(self, x):
gate, up = self.gate_up(x).chunk(2, dim=-1)
return self.down(F.silu(gate) * up)
class DeepDeltaResidual(nn.Module):
"""Scalar Deep-Delta residual update used only in the shared recurrent body."""
def __init__(self, c):
super().__init__()
self.k_eps = c.ddl_k_eps
self.v_sigmoid_scale = c.ddl_v_sigmoid_scale
self.beta_init = c.ddl_beta_init
self.beta = nn.Linear(c.hidden_size, 1, bias=True)
self.v_proj = nn.Linear(c.hidden_size, 1, bias=True)
@torch.no_grad()
def reset_beta_bias(self):
p = min(max(self.beta_init, 0.0), 2.0) / 2.0
p = min(max(p, 1e-6), 1.0 - 1e-6)
self.beta.bias.fill_(math.log(p) - math.log(1.0 - p))
def forward(self, x, *, k_in, context):
d = k_in.size(-1)
eps_rms = (self.k_eps * self.k_eps) / d
k_rms = F.rms_norm(k_in, [d], eps=eps_rms)
scale = 1.0 / math.sqrt(d)
beta = 2.0 * torch.sigmoid(self.beta(context).float())
proj = torch.sum(k_rms * x, dim=-1, keepdim=True, dtype=torch.float32) * scale
v = torch.sigmoid(self.v_proj(x).float()) * self.v_sigmoid_scale
delta = ((beta * (v - proj)) * scale).to(dtype=x.dtype)
return x + delta * k_rms
class PlainBlock(nn.Module):
def __init__(self, c):
super().__init__()
self.attn_norm = RMSNorm(c.hidden_size, c.rms_eps)
self.attn = Attention(c)
self.ffn_norm = RMSNorm(c.hidden_size, c.rms_eps)
self.ffn = SwiGLU(c)
def forward(self, x, cos, sin, attn_mask):
x = x + self.attn(self.attn_norm(x), cos, sin, attn_mask=attn_mask)
return x + self.ffn(self.ffn_norm(x))
class ContinuousLoopConditioner(nn.Module):
def __init__(self, d):
super().__init__()
self.net = nn.Sequential(nn.Linear(8, d), nn.SiLU(), nn.Linear(d, 2 * d))
@staticmethod
def features(t, dt, device, dtype):
return torch.tensor([
t, dt,
math.sin(math.pi * t), math.cos(math.pi * t),
math.sin(2 * math.pi * t), math.cos(2 * math.pi * t),
math.log(max(dt, 1e-6)), math.log(max(1.0 - t + dt, 1e-6)),
], device=device, dtype=dtype)
def parameters_for(self, t, dt, device, dtype):
f = self.features(t, dt, device, dtype)
scale, shift = self.net(f).chunk(2, dim=-1)
return f, scale, shift
@staticmethod
def modulate(x, scale, shift):
return x * (1.0 + 0.1 * scale) + 0.1 * shift
class PhaseLoRA(nn.Module):
def __init__(self, c):
super().__init__()
self.down = nn.Linear(c.hidden_size, c.lora_rank, bias=False)
self.up = nn.Linear(c.lora_rank, c.qkv_dim, bias=False)
self.gate = nn.Linear(8, c.lora_rank, bias=True)
def forward(self, x, phase_features):
gate = (2.0 * torch.sigmoid(self.gate(phase_features).float())).to(dtype=x.dtype)
return self.up(self.down(x) * gate)
class LoopedBlock(nn.Module):
def __init__(self, c):
super().__init__()
self.attn_norm = RMSNorm(c.hidden_size, c.rms_eps)
self.attn = Attention(c)
self.phase_lora = PhaseLoRA(c)
self.ddl_attn = DeepDeltaResidual(c)
self.ffn_norm = RMSNorm(c.hidden_size, c.rms_eps)
self.ffn = SwiGLU(c)
self.ddl_ffn = DeepDeltaResidual(c)
def forward(self, x, cos, sin, phase_scale, phase_shift, phase_features, attn_mask):
attn_context = ContinuousLoopConditioner.modulate(x, phase_scale, phase_shift)
qkv_delta = self.phase_lora(attn_context, phase_features)
x_norm = self.attn_norm(attn_context)
x = self.ddl_attn(
x,
k_in=self.attn(x_norm, cos, sin, qkv_delta=qkv_delta, attn_mask=attn_mask),
context=x_norm,
)
ffn_context = ContinuousLoopConditioner.modulate(x, phase_scale, phase_shift)
x_norm = self.ffn_norm(ffn_context)
return self.ddl_ffn(x, k_in=self.ffn(x_norm), context=x_norm)
class LoopHyperConnection(nn.Module):
"""Two persistent loop lanes with per-budget read/mix/write scalars."""
def __init__(self, c):
super().__init__()
self.k = int(c.hyper_lanes)
self.max_loops = int(c.max_loops)
if self.k != 2:
raise ValueError("SparkBET is defined for two Hyper-Connection lanes")
shape = (self.max_loops, self.max_loops)
self.alpha = nn.Parameter(torch.zeros(*shape, self.k))
self.mix = nn.Parameter(torch.zeros(*shape, self.k, self.k))
self.beta = nn.Parameter(torch.zeros(*shape, self.k))
self.reset_parameters()
@torch.no_grad()
def reset_parameters(self):
self.alpha.zero_(); self.mix.zero_(); self.beta.zero_()
eye = torch.eye(self.k, device=self.mix.device, dtype=self.mix.dtype)
for budget in range(1, self.max_loops + 1):
b = budget - 1
for i in range(budget):
self.alpha[b, i, i % self.k] = 1.0
self.mix[b, i].copy_(eye)
self.beta[b, i].fill_(1.0)
def init_lanes(self, prelude_state):
return prelude_state.unsqueeze(0).expand(self.k, *prelude_state.shape)
def read(self, lanes, loops, iteration):
a = self.alpha[loops - 1, iteration].to(dtype=lanes.dtype)
return torch.einsum("k,kbtd->btd", a, lanes)
def write(self, lanes, branch_delta, loops, iteration):
m = self.mix[loops - 1, iteration].to(dtype=lanes.dtype)
b = self.beta[loops - 1, iteration].to(dtype=lanes.dtype)
mixed = torch.einsum("kj,jbtd->kbtd", m, lanes)
return mixed + b[:, None, None, None] * branch_delta.unsqueeze(0)
@staticmethod
def pool(lanes):
return lanes.mean(dim=0)
class SparkBET(nn.Module):
"""Prelude -> six shared recurrent blocks -> coda, with exact loop budgets 1..8."""
def __init__(self, c=CFG):
super().__init__()
self.c = c
self.embed = nn.Embedding(c.vocab_size, c.hidden_size)
self.prelude = nn.ModuleList([PlainBlock(c) for _ in range(c.prelude_layers)])
self.body = nn.ModuleList([LoopedBlock(c) for _ in range(c.body_blocks)])
self.time_cond = ContinuousLoopConditioner(c.hidden_size)
self.loop_hyper = LoopHyperConnection(c)
self.coda = nn.ModuleList([PlainBlock(c) for _ in range(c.coda_layers)])
self.final_norm = RMSNorm(c.hidden_size, c.rms_eps)
self.apply(self._generic_init)
self._mechanism_init()
@staticmethod
def _generic_init(m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, 0.0, 0.02)
if m.bias is not None: nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, 0.0, 0.02)
@torch.no_grad()
def _mechanism_init(self):
nn.init.zeros_(self.time_cond.net[-1].weight)
nn.init.zeros_(self.time_cond.net[-1].bias)
self.loop_hyper.reset_parameters()
for block in self.body:
nn.init.zeros_(block.phase_lora.up.weight)
nn.init.zeros_(block.phase_lora.gate.weight)
nn.init.zeros_(block.phase_lora.gate.bias)
for block in [*self.prelude, *self.body, *self.coda]:
nn.init.zeros_(block.attn.out.weight)
nn.init.zeros_(block.ffn.down.weight)
for block in self.body:
block.ddl_attn.reset_beta_bias(); block.ddl_ffn.reset_beta_bias()
def _run_plain(self, block, x, cos, sin, attn_mask):
if _USE_GRAD_CHECKPOINTING and self.training:
return checkpoint(block, x, cos, sin, attn_mask, use_reentrant=False)
return block(x, cos, sin, attn_mask)
def _run_looped(self, block, x, cos, sin, scale, shift, features, attn_mask):
if _USE_GRAD_CHECKPOINTING and self.training:
return checkpoint(block, x, cos, sin, scale, shift, features, attn_mask, use_reentrant=False)
return block(x, cos, sin, scale, shift, features, attn_mask)
def _readout(self, x, cos, sin, attn_mask):
h = x
for block in self.coda:
h = self._run_plain(block, h, cos, sin, attn_mask)
h = self.final_norm(h)
return F.linear(h, self.embed.weight)
@staticmethod
def _schedule(step_sizes):
if isinstance(step_sizes, int):
n = int(step_sizes)
step_sizes = uniform_steps(n)
if not step_sizes:
raise ValueError("empty refinement schedule")
values = [float(v) for v in step_sizes]
if any(v <= 0 for v in values):
raise ValueError("refinement strides must be positive")
if abs(sum(values) - 1.0) > 1e-5:
raise ValueError("refinement strides must sum to 1")
return values
def _prepare(self, input_ids, step_sizes, attention_mask=None):
steps = self._schedule(step_sizes)
loops = len(steps)
if loops > self.c.max_loops:
raise ValueError(f"loops {loops} > configured max_loops {self.c.max_loops}")
if input_ids.shape[1] > self.c.max_seq_len:
raise ValueError("context exceeds max_seq_len")
x = self.embed(input_ids)
attn_mask, pos = attention_mask_and_positions(input_ids, attention_mask)
cos, sin = rope_cos_sin(pos, self.c.head_dim, self.c.rope_theta, x.dtype)
for block in self.prelude:
x = self._run_plain(block, x, cos, sin, attn_mask)
lanes = self.loop_hyper.init_lanes(x)
shared_noise = torch.randn_like(x) if self.training and _STATE_NOISE_SIGMA > 0 else None
return steps, lanes, shared_noise, cos, sin, attn_mask
def _advance(self, lanes, steps, iteration, elapsed, shared_noise, cos, sin, attn_mask):
dt = steps[iteration]
t_mid = elapsed + 0.5 * dt
features, scale, shift = self.time_cond.parameters_for(t_mid, dt, lanes.device, lanes.dtype)
branch_input = self.loop_hyper.read(lanes, len(steps), iteration)
if shared_noise is not None:
t_end = elapsed + dt
sigma = _STATE_NOISE_SIGMA * max(0.0, 1.0 - t_end)
if sigma:
branch_input = branch_input + sigma * shared_noise
h = branch_input
for block in self.body:
h = self._run_looped(block, h, cos, sin, scale, shift, features, attn_mask)
return self.loop_hyper.write(lanes, h - branch_input, len(steps), iteration)
def _run_recurrence(self, input_ids, step_sizes, attention_mask=None, collect_states=False):
steps, lanes, noise, cos, sin, attn_mask = self._prepare(input_ids, step_sizes, attention_mask)
states = [] if collect_states else None
elapsed = 0.0
for i, dt in enumerate(steps):
lanes = self._advance(lanes, steps, i, elapsed, noise, cos, sin, attn_mask)
elapsed += dt
if collect_states: states.append(self.loop_hyper.pool(lanes))
return self.loop_hyper.pool(lanes), states, cos, sin, attn_mask
def forward(self, input_ids, step_sizes=None, attention_mask=None):
if step_sizes is None: step_sizes = uniform_steps(self.c.max_loops)
x, _, cos, sin, attn_mask = self._run_recurrence(input_ids, step_sizes, attention_mask, False)
return self._readout(x, cos, sin, attn_mask)
def forward_loop_exits(self, input_ids, step_sizes=None, attention_mask=None):
if step_sizes is None: step_sizes = uniform_steps(self.c.max_loops)
_, states, cos, sin, attn_mask = self._run_recurrence(input_ids, step_sizes, attention_mask, True)
return [self._readout(h, cos, sin, attn_mask) for h in states]
def count_params(self):
return sum(p.numel() for p in self.parameters())
# Preserve the old trainer/import name while changing the implementation.
BETFog = SparkBET
def uniform_steps(n):
n = int(n)
if not 1 <= n <= MAX_LOOPS:
raise ValueError(f"loop budget must be in [1,{MAX_LOOPS}]")
return [1.0 / n] * n
def architecture_shape_signature(model):
lines = [f"{k}:{tuple(v.shape)}:{v.dtype}" for k, v in model.state_dict().items()]
return hashlib.sha256("\n".join(lines).encode()).hexdigest()
@torch.no_grad()
def verify_architecture(model, rank=0):
expected = dict(
vocab_size=259, hidden_size=324, intermediate_size=864,
prelude_layers=1, body_blocks=6, coda_layers=1,
num_heads=6, num_kv_heads=2, head_dim=54,
lora_rank=16, hyper_lanes=2, max_seq_len=1024, max_loops=8,
rope_theta=10_000.0, rms_eps=1e-6,
ddl_beta_init=1.0, ddl_k_eps=1e-2, ddl_v_sigmoid_scale=4.0,
)
actual = asdict(model.c)
for k, v in expected.items():
if actual[k] != v:
raise AssertionError(f"Architecture drift: {k}={actual[k]} expected {v}")
if model.count_params() != EXPECTED_PARAM_COUNT:
raise AssertionError(f"Parameter drift: {model.count_params():,} != {EXPECTED_PARAM_COUNT:,}")
sig = architecture_shape_signature(model)
if sig != EXPECTED_ARCH_SHAPE_SHA256:
raise AssertionError(f"Architecture SHA drift: {sig} != {EXPECTED_ARCH_SHAPE_SHA256}")
if not torch.allclose(model.time_cond.net[-1].weight, torch.zeros_like(model.time_cond.net[-1].weight)):
raise AssertionError("time conditioner must start as identity")
for i, block in enumerate(model.body):
if not torch.allclose(block.phase_lora.up.weight, torch.zeros_like(block.phase_lora.up.weight)):
raise AssertionError(f"phase LoRA {i} up projection must start zero")
for name, ddl in (("attn", block.ddl_attn), ("ffn", block.ddl_ffn)):
beta = (2.0 * torch.sigmoid(ddl.beta.bias.float())).item()
if abs(beta - 1.0) > 1e-6:
raise AssertionError(f"body {i} {name} DDL beta init={beta}")
hc = model.loop_hyper
eye = torch.eye(hc.k, device=hc.mix.device, dtype=hc.mix.dtype)
for budget in range(1, model.c.max_loops + 1):
for i in range(budget):
expected_alpha = torch.zeros_like(hc.alpha[budget - 1, i]); expected_alpha[i % hc.k] = 1
if not torch.allclose(hc.alpha[budget - 1, i], expected_alpha): raise AssertionError("Hyper alpha init drift")
if not torch.allclose(hc.mix[budget - 1, i], eye): raise AssertionError("Hyper mix init drift")
if not torch.allclose(hc.beta[budget - 1, i], torch.ones_like(hc.beta[budget - 1, i])): raise AssertionError("Hyper beta init drift")
if rank == 0:
print("[verify] SparkBET architecture PASSED")
print(f"[verify] params: {model.count_params():,}")
print(f"[verify] physical blocks: 1 prelude + 6 shared body + 1 coda; L8 applications=50")
print(f"[verify] architecture SHA256: {sig}")