Download conversion/kai_graph.py from FluidInference/decision-2.0-kai-coreml: direct link, hf CLI and curl.
- Browser
- Download file 5.02 kB
-
https://huggingface.co/FluidInference/decision-2.0-kai-coreml/resolve/main/conversion/kai_graph.py
- Command line
-
hf download hf://FluidInference/decision-2.0-kai-coreml/conversion/kai_graph.py
-
curl -L -o kai_graph.py https://huggingface.co/FluidInference/decision-2.0-kai-coreml/resolve/main/conversion/kai_graph.py
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 | |