duotactic / code /leanoar /model.py
Duoia's picture
duotactic full package: checkpoints, tokenizer, config, code, docs
32c0c6c verified
Raw History Blame Contribute Delete
9.79 kB
"""Scheme C Final architecture: 40M-parameter Lean 4 tactic generator backbone.
Spec (see D:\\prover\\Lean Prover\\总方案.md):
vocab 4096 (byte-level BPE, tied), d_model 640, 8 layers, 10 heads (head_dim 64),
d_ff 1536 SwiGLU, pre-norm RMSNorm, RoPE, no biases,
low-rank policy head 640->128->640 then tied embedding,
stepped value MLP 640 -(stop_grad)-> 128 -> 32 -> 3.
"""
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
@dataclass
class ModelConfig:
vocab_size: int = 4096
block_size: int = 768
n_layer: int = 8
n_head: int = 10
n_embd: int = 640
intermediate_size: int = 1536
dropout: float = 0.0
rope_base: float = 10000.0
norm_eps: float = 1e-5
policy_rank: int = 128
value_hidden: int = 128
value_mid: int = 32
n_value_out: int = 3
tie_embeddings: bool = True
depth_rope: bool = False # spec §3.2; off for v1 (no AST depth in the data)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return (x.to(dtype)) * self.weight
def build_rope_cache(seq_len: int, head_dim: int, base: float, device, dtype):
inv = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim))
t = torch.arange(seq_len, device=device, dtype=torch.float32)
freqs = torch.outer(t, inv)
return torch.cos(freqs).to(dtype), torch.sin(freqs).to(dtype)
def apply_rope(x, cos, sin, depth=None):
"""x: [B, H, T, D]. cos/sin: [T, D/2]. depth: optional [B, T] int for depth-aware RoPE."""
B, H, T, D = x.shape
half = D // 2
x1, x2 = x[..., :half], x[..., half:]
if depth is None:
c, s = cos[:T].view(1, 1, T, half), sin[:T].view(1, 1, T, half)
else:
# first half of the rotary space -> sequence position, second half -> AST depth
q = half // 2
c1, s1 = cos[:T].view(1, 1, T, half)[..., :q], sin[:T].view(1, 1, T, half)[..., :q]
d = depth.clamp(max=cos.shape[0] - 1)
cd, sd = cos[d].unsqueeze(1), sin[d].unsqueeze(1) # [B,1,T,half]
c = torch.cat([c1, cd[..., :half - q]], dim=-1)
s = torch.cat([s1, sd[..., :half - q]], dim=-1)
out1 = x1 * c - x2 * s
out2 = x1 * s + x2 * c
return torch.cat([out1, out2], dim=-1)
class Attention(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.n_head = cfg.n_head
self.head_dim = cfg.n_embd // cfg.n_head
self.qkv = nn.Linear(cfg.n_embd, 3 * cfg.n_embd, bias=False)
self.proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=False)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x, cos, sin, depth=None):
B, T, C = x.shape
q, k, v = self.qkv(x).split(C, dim=2)
q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
q = apply_rope(q, cos, sin, depth)
k = apply_rope(k, cos, sin, depth)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.drop(self.proj(y))
class SwiGLU(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.gate = nn.Linear(cfg.n_embd, cfg.intermediate_size, bias=False)
self.up = nn.Linear(cfg.n_embd, cfg.intermediate_size, bias=False)
self.down = nn.Linear(cfg.intermediate_size, cfg.n_embd, bias=False)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x):
return self.drop(self.down(F.silu(self.gate(x)) * self.up(x)))
class Block(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.norm_1 = RMSNorm(cfg.n_embd, cfg.norm_eps)
self.attn = Attention(cfg)
self.norm_2 = RMSNorm(cfg.n_embd, cfg.norm_eps)
self.mlp = SwiGLU(cfg)
def forward(self, x, cos, sin, depth=None):
x = x + self.attn(self.norm_1(x), cos, sin, depth)
x = x + self.mlp(self.norm_2(x))
return x
class SchemeC(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.cfg = cfg
self.wte = nn.Embedding(cfg.vocab_size, cfg.n_embd)
self.drop = nn.Dropout(cfg.dropout)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)])
self.ln_f = RMSNorm(cfg.n_embd, cfg.norm_eps)
# low-rank policy projection (decouples input/output semantic spaces)
self.policy_down = nn.Linear(cfg.n_embd, cfg.policy_rank, bias=False)
self.policy_up = nn.Linear(cfg.policy_rank, cfg.n_embd, bias=False)
# stepped value MLP with a stop-gradient between backbone and the MLP
self.value_fc1 = nn.Linear(cfg.n_embd, cfg.value_hidden, bias=True)
self.value_fc2 = nn.Linear(cfg.value_hidden, cfg.value_mid, bias=True)
self.value_fc3 = nn.Linear(cfg.value_mid, cfg.n_value_out, bias=True)
self.apply(self._init)
self._tied = False
if cfg.tie_embeddings:
self.tie_weights()
cos, sin = build_rope_cache(cfg.block_size, cfg.n_embd // cfg.n_head, cfg.rope_base,
torch.device('cpu'), torch.float32)
self.register_buffer('_cos', cos, persistent=False)
self.register_buffer('_sin', sin, persistent=False)
@staticmethod
def _init(m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.02)
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, std=0.02)
def tie_weights(self):
"""Policy logits reuse the token embedding matrix (weight tying)."""
self._tied = True
return self
def num_params(self, trainable_only: bool = False):
seen, total = set(), 0
for p in self.parameters():
if trainable_only and not p.requires_grad:
continue
if id(p) in seen:
continue
seen.add(id(p))
total += p.numel()
return total
def forward(self, idx, targets=None, loss_mask_start=None, value_targets=None,
value_weight=0.0, policy_weight=1.0):
B, T = idx.shape
assert T <= self.cfg.block_size, f'seq {T} > block {self.cfg.block_size}'
cos = self._cos.to(idx.device)
sin = self._sin.to(idx.device)
x = self.drop(self.wte(idx))
for blk in self.blocks:
x = blk(x, cos, sin)
h = self.ln_f(x)
# ---- policy: low-rank projection -> tied embedding -> logits
p = self.policy_up(F.gelu(self.policy_down(h)))
logits = p @ self.wte.weight.t()
out = {'logits': logits}
if targets is not None:
if loss_mask_start is not None:
loss_mask_start = loss_mask_start.to(idx.device)
labels = targets.clone()
pos = torch.arange(T, device=idx.device).unsqueeze(0)
labels[pos < loss_mask_start.unsqueeze(1)] = -100
labels[targets == 0] = -100 # <|pad|> id = 0
else:
labels = targets
out['policy_loss'] = F.cross_entropy(
logits.reshape(-1, logits.size(-1)).float(), labels.reshape(-1), ignore_index=-100)
if value_targets is not None and value_weight > 0:
hv = h.detach() # stop-gradient: isolates the backbone
v = self.value_fc3(F.gelu(self.value_fc2(F.gelu(self.value_fc1(hv)))))
out['value_pred'] = v
out['value_loss'] = F.mse_loss(v, value_targets)
out['loss'] = policy_weight * out['policy_loss'] + value_weight * out['value_loss']
elif 'policy_loss' in out:
out['loss'] = policy_weight * out['policy_loss']
return out
@torch.no_grad()
def value_head(self, idx):
"""[B,T,3] = (win, steps_left, confidence) — backbone frozen by stop-gradient."""
cos, sin = self._cos.to(idx.device), self._sin.to(idx.device)
x = self.drop(self.wte(idx))
for blk in self.blocks:
x = blk(x, cos, sin)
h = self.ln_f(x).detach()
return self.value_fc3(F.gelu(self.value_fc2(F.gelu(self.value_fc1(h)))))
def hidden(self, idx):
"""[B,T,d] final-layer states (no detach) — for caching features for the value head."""
cos = self._cos.to(idx.device)
sin = self._sin.to(idx.device)
x = self.drop(self.wte(idx))
for blk in self.blocks:
x = blk(x, cos, sin)
return self.ln_f(x)
@torch.no_grad()
def generate(self, idx, max_new_tokens: int = 48, temperature: float = 1.0,
logit_mask_fn=None):
self.eval()
for _ in range(max_new_tokens):
idx_c = idx[:, -self.cfg.block_size:]
logits = self.forward(idx_c)['logits'][:, -1, :].float()
if logit_mask_fn is not None:
logits = logit_mask_fn(idx, logits)
if temperature <= 1e-6:
nxt = logits.argmax(-1, keepdim=True)
else:
probs = F.softmax(logits / temperature, dim=-1)
nxt = torch.multinomial(probs, 1)
idx = torch.cat([idx, nxt], dim=1)
return idx