Download code/leanoar/model.py from Duoia/duotactic: direct link, hf CLI and curl.
- Browser
- Download file 9.79 kB
-
https://huggingface.co/Duoia/duotactic/resolve/main/code/leanoar/model.py
- Command line
-
hf download hf://Duoia/duotactic/code/leanoar/model.py
-
curl -L -o model.py https://huggingface.co/Duoia/duotactic/resolve/main/code/leanoar/model.py
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 | |
| 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) | |
| 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 | |
| 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) | |
| 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 | |