"""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