"""Decision-2.0-Kai-0.6B as one traceable graph: Qwen3 backbone + shared decision head. Inputs are a packed "tree" request: a shared prefix plus every question's suffix in one row, an additive attention mask, explicit RoPE positions, and per-candidate gather indices (candidate endpoint, its question's query token). Output: one logit per candidate. """ import json import math import torch import torch.nn.functional as F from safetensors.torch import load_file from torch import nn ROOT = "kai" class RMSNorm(nn.Module): def __init__(self, dim, eps): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x): x32 = x.float() var = x32.pow(2).mean(-1, keepdim=True) return (x32 * torch.rsqrt(var + self.eps)).to(x.dtype) * self.weight def rotate_half(x): a, b = x.chunk(2, dim=-1) return torch.cat((-b, a), dim=-1) class Layer(nn.Module): def __init__(self, c): super().__init__() h, d = c["hidden_size"], c["head_dim"] self.nq, self.nkv, self.d = c["num_attention_heads"], c["num_key_value_heads"], d eps = c["rms_norm_eps"] self.input_layernorm = RMSNorm(h, eps) self.post_attention_layernorm = RMSNorm(h, eps) self.q_proj = nn.Linear(h, self.nq * d, bias=False) self.k_proj = nn.Linear(h, self.nkv * d, bias=False) self.v_proj = nn.Linear(h, self.nkv * d, bias=False) self.o_proj = nn.Linear(self.nq * d, h, bias=False) self.q_norm = RMSNorm(d, eps) self.k_norm = RMSNorm(d, eps) i = c["intermediate_size"] self.gate_proj = nn.Linear(h, i, bias=False) self.up_proj = nn.Linear(h, i, bias=False) self.down_proj = nn.Linear(i, h, bias=False) def forward(self, x, cos, sin, mask): L = x.shape[1] y = self.input_layernorm(x) q = self.q_norm(self.q_proj(y).view(1, L, self.nq, self.d)).transpose(1, 2) k = self.k_norm(self.k_proj(y).view(1, L, self.nkv, self.d)).transpose(1, 2) v = self.v_proj(y).view(1, L, self.nkv, self.d).transpose(1, 2) q = q * cos + rotate_half(q) * sin k = k * cos + rotate_half(k) * sin rep = self.nq // self.nkv k = k.repeat_interleave(rep, dim=1) v = v.repeat_interleave(rep, dim=1) att = torch.matmul(q, k.transpose(-1, -2)) * (1.0 / math.sqrt(self.d)) + mask att = torch.softmax(att.float(), dim=-1).to(q.dtype) o = torch.matmul(att, v).transpose(1, 2).reshape(1, L, self.nq * self.d) x = x + self.o_proj(o) y = self.post_attention_layernorm(x) return x + self.down_proj(F.silu(self.gate_proj(y)) * self.up_proj(y)) class KaiGraph(nn.Module): def __init__(self, root=ROOT): super().__init__() c = json.load(open(f"{root}/backbone/config.json")) self.c = c h = c["hidden_size"] self.embed_tokens = nn.Embedding(c["vocab_size"], h) self.layers = nn.ModuleList(Layer(c) for _ in range(c["num_hidden_layers"])) self.norm = RMSNorm(h, c["rms_norm_eps"]) theta = c["rope_parameters"]["rope_theta"] d = c["head_dim"] self.register_buffer( "inv_freq", 1.0 / theta ** (torch.arange(0, d, 2).float() / d), persistent=False ) # head self.candidate_norm = nn.LayerNorm(h) self.query_norm = nn.LayerNorm(h) self.key = nn.Linear(h, 256, bias=False) self.query = nn.Linear(h, 256, bias=False) self.candidate_mlp = nn.Linear(h, 256) self.query_mlp = nn.Linear(h, 256, bias=False) self.scalar = nn.Linear(256, 1, bias=False) self._load(root) def _load(self, root): sd = {} for k, v in load_file(f"{root}/backbone/model.safetensors").items(): k = k.replace("self_attn.", "").replace("mlp.", "") sd[k] = v.float() sd.update({k: v.float() for k, v in load_file(f"{root}/decision_head.safetensors").items()}) missing, unexpected = self.load_state_dict(sd, strict=False) assert not missing and not unexpected, (missing, unexpected) def forward(self, input_ids, position_ids, mask, cand_idx, query_idx): # input_ids/position_ids [1,L] int32; mask [1,1,L,L] additive; cand/query_idx [N] int32 x = self.embed_tokens(input_ids) freqs = position_ids.float()[0, :, None] * self.inv_freq[None, :] emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos()[None, None].to(x.dtype) sin = emb.sin()[None, None].to(x.dtype) for layer in self.layers: x = layer(x, cos, sin, mask) x = self.norm(x)[0] cand = self.candidate_norm(x[cand_idx].float()) qry = self.query_norm(x[query_idx].float()) bilinear = (self.key(cand) * self.query(qry)).sum(-1) / 16.0 nonlinear = self.scalar(F.gelu(self.candidate_mlp(cand) + self.query_mlp(qry))).squeeze(-1) return bilinear + nonlinear