alexwengg's picture
Decision-2.0-Kai-0.6B Core ML: packed multi-question fp16 package, runtime, parity reports
6a545c7 verified
Raw History Blame Contribute Delete
5.02 kB
"""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