MyeongHoJeong's picture
Standard One 8B SH: adapter, schema head and server code
c1264c3 verified
Raw History Blame Contribute Delete
15.5 kB
"""Joint schema head (DESIGN.md section 2.2). Own design; provenance in DESIGN.md section 0.
Reads tapped backbone hidden states for ONE request at a time (a micro-batch is a Python loop over requests; the
head is small) and returns one logit vector per question over that question's options in CANONICAL order.
Nodes: one per question and one per option. No positional encoding over nodes, attention pooling inside spans and
relation-type attention biases that depend only on structure make the head exactly equivariant to the order of
choice/noul options and of questions (tests/test_equivariance.py). Score levels carry an ordinal encoding and are
deliberately not permutation-equivariant.
"""
import math
from dataclasses import asdict, dataclass, field
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
QTYPE_ID = {"choice": 0, "noul": 1, "score": 2}
# relation types for node pairs (i attends to j)
REL_SELF, REL_SIBLING, REL_OPT_TO_OWN_Q, REL_Q_TO_OWN_OPT, REL_OPT_TO_OTHER_OPT, REL_OPT_TO_OTHER_Q, \
REL_Q_TO_OTHER_Q, REL_Q_TO_OTHER_OPT = range(8)
N_REL = 8
ORD_FEATS = 33
@dataclass
class HeadConfig:
hidden_size: int = 4096
n_taps: int = 2
d: int = 1024
heads: int = 16
ffn: int = 2816
route_blocks: int = 2
interact_blocks: int = 2
dropout: float = 0.1
n_roles: int = 5
ord_hidden: int = 256
max_nodes: int = 4096
grad_checkpoint: bool = True
taps: list = field(default_factory=lambda: ["final", 25])
template: str = "schema-v1"
def to_dict(self):
return asdict(self)
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):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight
class SwiGLU(nn.Module):
def __init__(self, d, hidden, dropout):
super().__init__()
self.w1 = nn.Linear(d, hidden, bias=False)
self.w3 = nn.Linear(d, hidden, bias=False)
self.w2 = nn.Linear(hidden, d, bias=False)
self.drop = nn.Dropout(dropout)
def forward(self, x):
return self.drop(self.w2(F.silu(self.w1(x)) * self.w3(x)))
class Attention(nn.Module):
"""Multi-head attention (queries from nodes, keys/values from `kv`), SDPA with optional additive bias."""
def __init__(self, d, heads, dropout):
super().__init__()
self.h, self.dh = heads, d // heads
self.q = nn.Linear(d, d, bias=False)
self.k = nn.Linear(d, d, bias=False)
self.v = nn.Linear(d, d, bias=False)
self.o = nn.Linear(d, d, bias=False)
self.p = dropout
def _chunk(self, q, part, p):
k = self.k(part).view(-1, self.h, self.dh).transpose(0, 1).unsqueeze(0)
v = self.v(part).view(-1, self.h, self.dh).transpose(0, 1).unsqueeze(0)
scores = (q @ k.transpose(-1, -2)) / math.sqrt(self.dh)
lse = torch.logsumexp(scores, -1, keepdim=True)
probs = F.dropout(torch.softmax(scores, -1), p, self.training)
return probs @ v, lse
def forward(self, x, kv, bias=None, kv_chunk=0):
n, m = x.shape[0], kv.shape[0]
q = self.q(x).view(n, self.h, self.dh).transpose(0, 1).unsqueeze(0)
p = self.p if self.training else 0.0
if kv_chunk and m > kv_chunk and bias is None:
# exact chunked attention over keys (log-sum-exp merge) to bound K/V memory at long contexts.
# Dropout is applied to each chunk's attention PROBABILITIES (as SDPA's dropout_p): the global weight of
# key j is softmax_chunk(j) * exp(lse_chunk - lse_all), so dropping the chunk-local probability with
# scale 1/(1-p) is the same random variable as dropping the global one. When training, every chunk is
# recomputed in backward (checkpoint), so its scores/probabilities are not kept for the whole sequence.
outs, lses = [], []
use_ckpt = self.training and torch.is_grad_enabled()
for s in range(0, m, kv_chunk):
part = kv[s:s + kv_chunk]
if use_ckpt:
o, lse = checkpoint(self._chunk, q, part, p, use_reentrant=False)
else:
o, lse = self._chunk(q, part, p)
outs.append(o)
lses.append(lse)
lse_all = torch.logsumexp(torch.cat(lses, -1), -1, keepdim=True)
out = sum(o * torch.exp(l - lse_all) for o, l in zip(outs, lses))
else:
k = self.k(kv).view(m, self.h, self.dh).transpose(0, 1).unsqueeze(0)
v = self.v(kv).view(m, self.h, self.dh).transpose(0, 1).unsqueeze(0)
mask = None if bias is None else bias.unsqueeze(0)
out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=p)
return self.o(out.squeeze(0).transpose(0, 1).reshape(n, -1))
class RouteBlock(nn.Module):
"""Stage 1 (evidence routing): nodes cross-attend over the state tokens, then FFN."""
def __init__(self, cfg):
super().__init__()
self.n1, self.n2, self.nkv = RMSNorm(cfg.d), RMSNorm(cfg.d), RMSNorm(cfg.d)
self.cross = Attention(cfg.d, cfg.heads, cfg.dropout)
self.ffn = SwiGLU(cfg.d, cfg.ffn, cfg.dropout)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x, state, kv_chunk=0):
x = x + self.drop(self.cross(self.n1(x), self.nkv(state), kv_chunk=kv_chunk))
return x + self.ffn(self.n2(x))
class InteractBlock(nn.Module):
"""Stage 2 (field interaction): node self-attention with relation-type biases, cross-attention back to the
state tokens, FFN."""
def __init__(self, cfg):
super().__init__()
self.n1, self.n2, self.n3, self.nkv = RMSNorm(cfg.d), RMSNorm(cfg.d), RMSNorm(cfg.d), RMSNorm(cfg.d)
self.self_attn = Attention(cfg.d, cfg.heads, cfg.dropout)
self.cross = Attention(cfg.d, cfg.heads, cfg.dropout)
self.rel_bias = nn.Parameter(torch.zeros(N_REL, cfg.heads))
self.ffn = SwiGLU(cfg.d, cfg.ffn, cfg.dropout)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x, state, rel, kv_chunk=0):
bias = self.rel_bias[rel].permute(2, 0, 1) # [heads, N, N]
y = self.n1(x)
x = x + self.drop(self.self_attn(y, y, bias=bias))
x = x + self.drop(self.cross(self.n2(x), self.nkv(state), kv_chunk=kv_chunk))
return x + self.ffn(self.n3(x))
def ordinal_features(k, K):
"""33 features of score level k of K (2 <= K <= 10): position, scale size, one-hots, sin/cos."""
pos = k / (K - 1)
feats = [pos, K / 10.0]
feats += [1.0 if i == k else 0.0 for i in range(10)]
feats += [1.0 if K == i else 0.0 for i in range(2, 11)]
for f in range(1, 7):
feats += [math.sin(math.pi * f * pos), math.cos(math.pi * f * pos)]
assert len(feats) == ORD_FEATS
return feats
class JointSchemaHead(nn.Module):
def __init__(self, cfg: HeadConfig):
super().__init__()
self.cfg = cfg
d = cfg.d
self.tap_norms = nn.ModuleList([RMSNorm(cfg.hidden_size) for _ in range(cfg.n_taps)])
self.tap_mix = nn.Parameter(torch.zeros(cfg.n_taps))
self.tap_gain = nn.Parameter(torch.ones(1))
self.in_proj = nn.Linear(cfg.hidden_size, d, bias=False)
self.role_emb = nn.Embedding(cfg.n_roles, d)
self.qtype_emb = nn.Embedding(3, d)
self.kind_emb = nn.Embedding(2, d) # 0 question node, 1 option node
self.pool_query = nn.Parameter(torch.randn(2, d) / math.sqrt(d))
self.pool_key = nn.Linear(d, d, bias=False)
self.pool_out = nn.Linear(2 * d, d, bias=False)
self.ord_mlp = nn.Sequential(nn.Linear(ORD_FEATS, cfg.ord_hidden), nn.GELU(), nn.Linear(cfg.ord_hidden, d))
self.route = nn.ModuleList([RouteBlock(cfg) for _ in range(cfg.route_blocks)])
self.interact = nn.ModuleList([InteractBlock(cfg) for _ in range(cfg.interact_blocks)])
self.final_norm = RMSNorm(d)
self.scorer = nn.Sequential(nn.Linear(3 * d, d), nn.GELU(), nn.Linear(d, 1))
self.log_temp = nn.Parameter(torch.zeros(3))
self.kv_chunk = 0
self._init()
def _init(self):
for m in self.modules():
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)
# ------------------------------------------------------------------ structure
@staticmethod
def plan(enc, device):
"""Index tensors for one encoded request (see render.Encoder.encode)."""
node_kind, node_q, node_qtype, node_ord = [], [], [], []
pool_node, pool_tok, last_tok = [], [], []
q_node, opt_nodes = [], []
nq = len(enc["keys"])
for i in range(nq):
qt = enc["qtype"][i]
idx = len(node_kind)
q_node.append(idx)
node_kind.append(0); node_q.append(i); node_qtype.append(qt); node_ord.append(None)
toks = list(range(*enc["q_span"][i]))
if enc["p_span"][i]:
toks = list(range(*enc["p_span"][i])) + toks
pool_node += [idx] * len(toks); pool_tok += toks
last_tok.append(enc["q_span"][i][1] - 1)
ids = []
K = enc["n_opt"][i]
for j in range(K):
idx = len(node_kind)
ids.append(idx)
node_kind.append(1); node_q.append(i); node_qtype.append(qt)
node_ord.append((j, K) if qt == "score" else None)
a, b = enc["o_span"][i][j]
pool_node += [idx] * (b - a); pool_tok += list(range(a, b))
last_tok.append(b - 1)
opt_nodes.append(ids)
n = len(node_kind)
kind = torch.tensor(node_kind)
qn = torch.tensor(node_q)
same_q = qn[:, None] == qn[None, :]
ki, kj = kind[:, None], kind[None, :]
rel = torch.empty(n, n, dtype=torch.long)
rel[(ki == 1) & (kj == 1) & same_q] = REL_SIBLING
rel[(ki == 1) & (kj == 0) & same_q] = REL_OPT_TO_OWN_Q
rel[(ki == 0) & (kj == 1) & same_q] = REL_Q_TO_OWN_OPT
rel[(ki == 1) & (kj == 1) & ~same_q] = REL_OPT_TO_OTHER_OPT
rel[(ki == 1) & (kj == 0) & ~same_q] = REL_OPT_TO_OTHER_Q
rel[(ki == 0) & (kj == 0)] = REL_Q_TO_OTHER_Q
rel[(ki == 0) & (kj == 1) & ~same_q] = REL_Q_TO_OTHER_OPT
rel.fill_diagonal_(REL_SELF)
ords = [ordinal_features(*o) if o else [0.0] * ORD_FEATS for o in node_ord]
state_tok = [t for a, b in enc["state_spans"] for t in range(a, b)]
p = {"n": n, "kind": kind.to(device), "qtype": torch.tensor([QTYPE_ID[t] for t in node_qtype]).to(device),
"is_score": torch.tensor([o is not None for o in node_ord]).to(device),
"ord": torch.tensor(ords, dtype=torch.float32).to(device), "rel": rel.to(device),
"pool_node": torch.tensor(pool_node).to(device), "pool_tok": torch.tensor(pool_tok).to(device),
"last_tok": torch.tensor(last_tok).to(device), "q_node": q_node, "opt_nodes": opt_nodes,
"state_tok": torch.tensor(state_tok).to(device), "qtypes": list(enc["qtype"])}
return p
# ------------------------------------------------------------------ forward
def mix_taps(self, taps):
"""taps: list of [T, hidden] tensors (any float dtype) -> [T, hidden] in the head's dtype (fp32)."""
w = torch.softmax(self.tap_mix, 0)
dtype = self.in_proj.weight.dtype
out = 0
for i, t in enumerate(taps):
out = out + w[i] * self.tap_norms[i](t.to(dtype))
return out * self.tap_gain
def _token_proj(self, roles, *taps):
return self.in_proj(self.mix_taps(list(taps))) + self.role_emb(roles)
def project_tokens(self, taps, roles, use_ckpt, chunk=8192):
"""Per-token tap mix + in-proj + role embedding, in token chunks (each checkpointed when training), so the
fp32 intermediates of only one chunk exist at a time (forward and recompute)."""
outs = []
for a in range(0, roles.shape[0], chunk):
args = (roles[a:a + chunk], *[t[a:a + chunk] for t in taps])
outs.append(checkpoint(self._token_proj, *args, use_reentrant=False) if use_ckpt else self._token_proj(*args))
return outs[0] if len(outs) == 1 else torch.cat(outs, 0)
def forward_one(self, taps, enc, roles):
"""taps: list of [T, hidden]; enc: encoded request; roles: LongTensor [T]. Returns list of logit vectors
(one per question in enc['keys'] order), each over CANONICAL option order, already divided by the
per-type temperature."""
cfg = self.cfg
dev = taps[0].device
p = self.plan(enc, dev)
require_nodes = p["n"] <= cfg.max_nodes
if not require_nodes:
raise ValueError(f"too many nodes: {p['n']} > {cfg.max_nodes}")
use_ckpt = cfg.grad_checkpoint and self.training and torch.is_grad_enabled()
# token-level projection: chunked + checkpointed so only the (already alive, bf16) tapped states are kept for
# backward, not ~4 fp32 [T, hidden] intermediates per tap (about 2 GB per tap at 32k tokens)
x_tok = self.project_tokens(taps, roles, use_ckpt)
state = x_tok.index_select(0, p["state_tok"])
# attention pooling inside spans (segment softmax via scatter)
toks = x_tok.index_select(0, p["pool_tok"])
node_of = p["pool_node"]
kind_of = p["kind"].index_select(0, node_of)
query = self.pool_query.index_select(0, kind_of)
scores = (self.pool_key(toks) * query).sum(-1) / math.sqrt(cfg.d)
mx = torch.full((p["n"],), -1e30, device=dev, dtype=scores.dtype)
mx = mx.scatter_reduce(0, node_of, scores, reduce="amax", include_self=True)
e = torch.exp(scores - mx.index_select(0, node_of))
den = torch.zeros(p["n"], device=dev, dtype=e.dtype).index_add(0, node_of, e)
w = e / den.index_select(0, node_of)
pooled = torch.zeros(p["n"], cfg.d, device=dev, dtype=toks.dtype).index_add(0, node_of, toks * w[:, None])
last = x_tok.index_select(0, p["last_tok"])
x = self.pool_out(torch.cat([pooled, last], -1))
x = x + self.qtype_emb(p["qtype"]) + self.kind_emb(p["kind"])
x = x + self.ord_mlp(p["ord"].to(x.dtype)) * p["is_score"][:, None].to(x.dtype)
for blk in self.route:
x = checkpoint(blk, x, state, self.kv_chunk, use_reentrant=False) if use_ckpt else blk(x, state, self.kv_chunk)
for blk in self.interact:
x = (checkpoint(blk, x, state, p["rel"], self.kv_chunk, use_reentrant=False) if use_ckpt
else blk(x, state, p["rel"], self.kv_chunk))
x = self.final_norm(x)
out = []
for i, qn in enumerate(p["q_node"]):
opts = torch.tensor(p["opt_nodes"][i], device=dev)
o = x.index_select(0, opts)
q = x[qn].expand_as(o)
z = self.scorer(torch.cat([o, q, o * q], -1)).squeeze(-1)
t = torch.exp(self.log_temp[QTYPE_ID[p["qtypes"][i]]])
out.append(z / t)
return out
def count_parameters(module):
return sum(p.numel() for p in module.parameters())