alexwengg's picture
Decision-2.0-Eos-0.8B Core ML: shared prefix + chunked packed questions, fp16 package, runtime, parity reports
7bf8323 verified
Raw History Blame Contribute Delete
18.7 kB
"""Decision-2.0-Eos-0.8B (Qwen3.5 hybrid) as one fused Core ML graph: shared prefix + packed questions.
Built on the Kev-0.8B FusedPass (kev_stages.py): the shared prefix is the "state" (S tokens, right-padded; padded
positions are exact no-ops in the delta rule and masked in attention), the questions' suffixes are packed in P
tokens. Each suffix restarts from the prefix's recurrent state and conv tail (segment mask), and attends to the
prefix and itself. The readout is Decision 2.0's shared bilinear + MLP head at each option's last token and the
question's final token (gather indices into the packed region).
"""
import json
import numpy as np
import torch
import torch.nn.functional as F
from safetensors.torch import load_file
from torch import nn
import kev_stages
from kev_stages import MASK_VALUE, FusedPass
from qwen35_export import TextConfig, rope_cos_sin
LAGS = 3
class EosFused(FusedPass):
def __init__(self, cfg: TextConfig, state_len: int, packed_len: int, lane: int | None = None, chunk: int = 64):
# lane = None: questions may span the whole packed region (inverse over all P)
super().__init__(cfg, state_len, packed_len, batch=1, max_options=2, chunk_size=chunk,
lane=lane or packed_len)
h = cfg.hidden_size
del self.q, self.k, self.pointer_scale, self.inverse_temperature
self.embed_tokens = nn.Embedding(1, h) # resized on load
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)
def forward(self, input_ids, cos, sin, valid, tail_onehot, segment, lag_keep, lag_tail, cand_idx, query_idx):
"""input_ids [1, S+P]; cos/sin [S+P, R]; valid [S]; tail_onehot [3, S]; segment [P, P]; lag_keep [3, P];
lag_tail [3, P, 3]; cand_idx/query_idx [N] (indices into the packed region) -> logits [N]."""
S, P = self.S, self.P
bias = torch.cat([
torch.cat([self.state_causal, self.state_block], dim=1),
torch.cat([((1.0 - valid) * MASK_VALUE)[None, :].expand(P, S), (1.0 - segment) * MASK_VALUE], dim=1),
], dim=0)
x = self.embed_tokens(input_ids)
for layer in self.decoder.layers:
h = layer.input_layernorm(x)
if layer.is_linear:
out = self._delta(layer.linear_attn, h, valid, tail_onehot, segment, lag_keep, lag_tail)
else:
out = self._attention(layer.self_attn, h, cos, sin, bias)
x = x + out
x = x + layer.mlp(layer.post_attention_layernorm(x))
xq = self.norm(x[0, S:])
cand = self.candidate_norm(xq[cand_idx].float())
qry = self.query_norm(xq[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
def load_eos(root: str, state_len: int, packed_len: int, lane: int | None = None) -> EosFused:
cfg = TextConfig(json.load(open(f"{root}/backbone/config.json")))
state = {k: v.float() for k, v in load_file(f"{root}/backbone/model.safetensors").items()}
model = EosFused(cfg, state_len, packed_len, lane)
model.decoder.load_merged(state)
model.norm.load_state_dict({"weight": state["norm.weight"]})
model.embed_tokens = nn.Embedding.from_pretrained(state["embed_tokens.weight"], freeze=True)
head = {k: v.float() for k, v in load_file(f"{root}/decision_head.safetensors").items()}
missing, unexpected = model.load_state_dict(head, strict=False)
assert not unexpected and not [m for m in missing if m.split(".")[0] in
("candidate_norm", "query_norm", "key", "query", "candidate_mlp",
"query_mlp", "scalar")], (missing, unexpected)
return model.eval()
def prefix_len(encoded):
seqs = [e["ids"] for e in encoded]
limit = min(min(e["candidate_positions"][0] for e in encoded), min(len(s) for s in seqs) - 1)
for i in range(limit):
if any(s[i] != seqs[0][i] for s in seqs):
return i
return limit
def eos_inputs(cfg: TextConfig, encoded, S: int, P: int, pad_id: int, lane: int | None = None,
row_len: int | None = None):
"""Fused inputs for a group of encoded rows (vendored `encode` output) sharing their common token prefix.
With `lane`, a question that would cross a lane boundary starts at the next lane; with `row_len`, question j
starts at j * row_len (EosRows)."""
pre = prefix_len(encoded)
if not 1 <= pre <= S:
raise ValueError(f"prefix {pre} tokens does not fit S{S}")
ids = np.full(S + P, pad_id, dtype=np.int64)
pos = np.zeros(S + P, dtype=np.int64)
ids[:pre] = encoded[0]["ids"][:pre]
pos[:pre] = np.arange(pre)
pos[pre:S] = pre - 1
valid = np.zeros(S, dtype=np.float32)
valid[:pre] = 1
tail = np.zeros((LAGS, S), dtype=np.float32)
for i in range(LAGS): # tail row i = pre-conv input at position pre - LAGS + i (zero before the start)
p = pre - LAGS + i
if p >= 0:
tail[i, p] = 1
segment = np.eye(P, dtype=np.float32)
keep = np.zeros((LAGS, P), dtype=np.float32)
lag_tail = np.zeros((LAGS, P, LAGS), dtype=np.float32)
pos[S:] = pre
cand, qry, owner = [], [], []
start = 0
for j, e in enumerate(encoded):
suffix = e["ids"][pre:]
n = len(suffix)
if row_len:
start = j * row_len
if n > row_len:
raise ValueError(f"question of {n} tokens > row {row_len}")
elif lane and start // lane != (start + n - 1) // lane:
start = (start // lane + 1) * lane
if start + n > P:
raise ValueError(f"questions exceed P{P}")
ids[S + start : S + start + n] = suffix
pos[S + start : S + start + n] = np.arange(pre, pre + n)
segment[start : start + n, start : start + n] = np.tril(np.ones((n, n), dtype=np.float32))
for p in range(n):
for s in range(1, LAGS + 1):
if p >= s:
keep[s - 1, start + p] = 1
else:
lag_tail[s - 1, start + p, LAGS + p - s] = 1
for c in e["candidate_positions"]:
cand.append(start + c - pre)
qry.append(start + e["query_position"] - pre)
owner.append(j)
start += n
cos, sin = rope_cos_sin(cfg, torch.from_numpy(pos)[None].expand(3, -1))
return {
"input_ids": ids[None].astype(np.int32), "cos": cos.numpy(), "sin": sin.numpy(), "valid": valid,
"tail_onehot": tail, "segment": segment, "lag_keep": keep, "lag_tail": lag_tail,
}, cand, qry, owner, start
def with_slots(x, cand, qry, N):
if len(cand) > N:
raise ValueError(f"{len(cand)} options > N{N}")
pad = N - len(cand)
return {**x, "cand_idx": np.array(cand + [0] * pad, dtype=np.int32),
"query_idx": np.array(qry + [0] * pad, dtype=np.int32)}
class EosRows(EosFused):
"""Questions in B rows of Q tokens (one question per row, right-padded): the question delta rule runs chunked
(chunk 64) from the prefix's final state, batched over rows, instead of one P x P segment inverse."""
def __init__(self, cfg: TextConfig, state_len: int, rows: int, row_len: int, chunk: int = 64):
assert state_len % chunk == 0 and row_len % chunk == 0
super().__init__(cfg, state_len, rows * row_len, lane=row_len, chunk=chunk)
self.B, self.Q = rows, row_len
def _delta(self, net, x, valid, tail_onehot, segment, lag_keep, lag_tail):
cfg, S, P, B, Q = self.cfg, self.S, self.P, self.B, self.Q
T = S + P
H, Dk, Dv, kh = cfg.lin_v_heads, cfg.lin_k_dim, cfg.lin_v_dim, cfg.lin_k_heads
rep = H // kh
pre = net.in_proj_qkv(x)[0]
pre_s, pre_p = pre[:S], pre[S:]
tail = torch.matmul(tail_onehot, pre_s)
weight = net.conv1d.weight[:, 0]
lags = weight.shape[1] - 1
conv_s, conv_p = pre_s * weight[:, lags], pre_p * weight[:, lags]
for s in range(1, lags + 1):
conv_s = conv_s + F.pad(pre_s[: S - s], (0, 0, s, 0)) * weight[:, lags - s]
shifted = F.pad(pre_p[: P - s], (0, 0, s, 0)) * lag_keep[s - 1][:, None] + torch.matmul(lag_tail[s - 1], tail)
conv_p = conv_p + shifted * weight[:, lags - s]
qkv = F.silu(torch.cat([conv_s, conv_p], dim=0))
q = qkv[:, : net.key_dim].reshape(T, kh, Dk)
k = qkv[:, net.key_dim : 2 * net.key_dim].reshape(T, kh, Dk)
v = qkv[:, 2 * net.key_dim :].reshape(T, H, Dv).transpose(0, 1)
z = net.in_proj_z(x).reshape(1, T, H, Dv)
beta = torch.sigmoid(net.in_proj_b(x))[0].transpose(0, 1)
g = (-torch.exp(net.A_log) * F.softplus(net.in_proj_a(x) + net.dt_bias))[0].transpose(0, 1)
q = q * torch.rsqrt((q * q).sum(-1, keepdim=True) + 1e-6)
k = k * torch.rsqrt((k * k).sum(-1, keepdim=True) + 1e-6)
q = (q[:, :, None].expand(T, kh, rep, Dk).reshape(T, H, Dk) * (Dk**-0.5)).transpose(0, 1)
k = k[:, :, None].expand(T, kh, rep, Dk).reshape(T, H, Dk).transpose(0, 1)
C, N = net.chunk, net.num_chunks
new_v, k_cumdecay, q_dec, k_dec, attn_intra, chunk_decay = kev_stages._delta_chunk_terms(
net, q[:, :S].reshape(H, N, C, Dk), k[:, :S].reshape(H, N, C, Dk), v[:, :S].reshape(H, N, C, Dv),
(beta[:, :S] * valid).reshape(H, N, C, 1), (g[:, :S] * valid).reshape(H, N, C))
state = torch.zeros(H, Dk, Dv, dtype=x.dtype)
outs = []
for i in range(N):
v_new = new_v[:, i] - torch.matmul(k_cumdecay[:, i], state)
outs.append(torch.matmul(q_dec[:, i], state) + torch.matmul(attn_intra[:, i], v_new))
state = state * chunk_decay[:, i] + torch.matmul(k_dec[:, i].transpose(-1, -2), v_new)
G, M = H * B, Q // C
new_v, k_cumdecay, q_dec, k_dec, attn_intra, chunk_decay = kev_stages._delta_chunk_terms(
net, q[:, S:].reshape(G, M, C, Dk), k[:, S:].reshape(G, M, C, Dk), v[:, S:].reshape(G, M, C, Dv),
beta[:, S:].reshape(G, M, C, 1), g[:, S:].reshape(G, M, C))
st = state[:, None].expand(H, B, Dk, Dv).reshape(G, Dk, Dv)
rows = []
for i in range(M):
v_new = new_v[:, i] - torch.matmul(k_cumdecay[:, i], st)
rows.append(torch.matmul(q_dec[:, i], st) + torch.matmul(attn_intra[:, i], v_new))
if i + 1 < M:
st = st * chunk_decay[:, i] + torch.matmul(k_dec[:, i].transpose(-1, -2), v_new)
outs.append(torch.stack(rows, dim=1).reshape(H, P, Dv))
core = torch.cat(outs, dim=1).transpose(0, 1)[None]
core = net.norm(core, z)
return net.out_proj(core.reshape(1, T, H * Dv))
def load_rows(root: str, state_len: int, rows: int, row_len: int) -> EosRows:
cfg = TextConfig(json.load(open(f"{root}/backbone/config.json")))
state = {k: v.float() for k, v in load_file(f"{root}/backbone/model.safetensors").items()}
model = EosRows(cfg, state_len, rows, row_len)
model.decoder.load_merged(state)
model.norm.load_state_dict({"weight": state["norm.weight"]})
model.embed_tokens = nn.Embedding.from_pretrained(state["embed_tokens.weight"], freeze=True)
head = {k: v.float() for k, v in load_file(f"{root}/decision_head.safetensors").items()}
model.load_state_dict(head, strict=False)
return model.eval()
class EosChunked(EosFused):
"""Packed questions with a chunked delta rule (chunk C) instead of one P x P segment inverse.
Per chunk, a token starts either from the carried recurrent state (its question began in an earlier chunk,
`cont` = 1) or from the prefix's final state S0 (its question began in this chunk). `seg_chunks` [M, C, C] is the
segment mask restricted to each chunk (j <= t, same question); `last_seg` [M, C] marks the tokens in the question
of the chunk's last token, the only one whose state is carried to the next chunk."""
def __init__(self, cfg: TextConfig, state_len: int, packed_len: int, chunk: int = 64):
assert state_len % chunk == 0 and packed_len % chunk == 0
super().__init__(cfg, state_len, packed_len, lane=packed_len, chunk=chunk)
self.C, self.M = chunk, packed_len // chunk
self.register_buffer("eye_c", torch.eye(chunk), persistent=False)
def forward(self, input_ids, cos, sin, valid, tail_onehot, segment, lag_keep, lag_tail, seg_chunks, cont,
last_seg, cand_idx, query_idx):
self._extra = (seg_chunks, cont, last_seg)
return super().forward(input_ids, cos, sin, valid, tail_onehot, segment, lag_keep, lag_tail, cand_idx, query_idx)
def _delta(self, net, x, valid, tail_onehot, segment, lag_keep, lag_tail):
seg_chunks, cont, last_seg = self._extra
cfg, S, P, C, M = self.cfg, self.S, self.P, self.C, self.M
T = S + P
H, Dk, Dv, kh = cfg.lin_v_heads, cfg.lin_k_dim, cfg.lin_v_dim, cfg.lin_k_heads
rep = H // kh
pre = net.in_proj_qkv(x)[0]
pre_s, pre_p = pre[:S], pre[S:]
tail = torch.matmul(tail_onehot, pre_s)
weight = net.conv1d.weight[:, 0]
lags = weight.shape[1] - 1
conv_s, conv_p = pre_s * weight[:, lags], pre_p * weight[:, lags]
for s in range(1, lags + 1):
conv_s = conv_s + F.pad(pre_s[: S - s], (0, 0, s, 0)) * weight[:, lags - s]
shifted = F.pad(pre_p[: P - s], (0, 0, s, 0)) * lag_keep[s - 1][:, None] + torch.matmul(lag_tail[s - 1], tail)
conv_p = conv_p + shifted * weight[:, lags - s]
qkv = F.silu(torch.cat([conv_s, conv_p], dim=0))
q = qkv[:, : net.key_dim].reshape(T, kh, Dk)
k = qkv[:, net.key_dim : 2 * net.key_dim].reshape(T, kh, Dk)
v = qkv[:, 2 * net.key_dim :].reshape(T, H, Dv).transpose(0, 1)
z = net.in_proj_z(x).reshape(1, T, H, Dv)
beta = torch.sigmoid(net.in_proj_b(x))[0].transpose(0, 1)
g = (-torch.exp(net.A_log) * F.softplus(net.in_proj_a(x) + net.dt_bias))[0].transpose(0, 1)
q = q * torch.rsqrt((q * q).sum(-1, keepdim=True) + 1e-6)
k = k * torch.rsqrt((k * k).sum(-1, keepdim=True) + 1e-6)
q = (q[:, :, None].expand(T, kh, rep, Dk).reshape(T, H, Dk) * (Dk**-0.5)).transpose(0, 1)
k = k[:, :, None].expand(T, kh, rep, Dk).reshape(T, H, Dk).transpose(0, 1)
N = net.num_chunks
new_v, k_cumdecay, q_dec, k_dec, attn_intra, chunk_decay = kev_stages._delta_chunk_terms(
net, q[:, :S].reshape(H, N, net.chunk, Dk), k[:, :S].reshape(H, N, net.chunk, Dk),
v[:, :S].reshape(H, N, net.chunk, Dv), (beta[:, :S] * valid).reshape(H, N, net.chunk, 1),
(g[:, :S] * valid).reshape(H, N, net.chunk))
state = torch.zeros(H, Dk, Dv, dtype=x.dtype)
outs = []
for i in range(N):
v_new = new_v[:, i] - torch.matmul(k_cumdecay[:, i], state)
outs.append(torch.matmul(q_dec[:, i], state) + torch.matmul(attn_intra[:, i], v_new))
state = state * chunk_decay[:, i] + torch.matmul(k_dec[:, i].transpose(-1, -2), v_new)
s0 = state
# questions: packed, chunked, segment-aware
qp = q[:, S:].reshape(H, M, C, Dk)
kp = k[:, S:].reshape(H, M, C, Dk)
vp = v[:, S:].reshape(H, M, C, Dv)
bp = beta[:, S:].reshape(H, M, C, 1)
gp = g[:, S:].reshape(H, M, C, 1)
cum = torch.matmul(seg_chunks, gp).squeeze(-1) # [H, M, C]: decay since max(chunk start, question start)
pair_decay = torch.exp((cum[..., :, None] - cum[..., None, :]) * seg_chunks + (seg_chunks - 1.0) * 1e4)
k_beta = kp * bp
kkt = torch.matmul(k_beta, kp.transpose(-1, -2)) * pair_decay
attn = torch.matmul(qp, kp.transpose(-1, -2)) * pair_decay
inv = kev_stages._unit_lower_inverse((self.eye_c + kkt * (seg_chunks - self.eye_c)).reshape(H * M, C, C), C)
inv = inv.reshape(H, M, C, C)
decay = torch.exp(cum)[..., None] # [H, M, C, 1]
new_v = torch.matmul(inv, vp * bp)
k_cd = torch.matmul(inv, k_beta * decay)
q_dec = qp * decay
ls = last_seg[None, :, :]
to_last = (torch.exp((cum[..., -1:] - cum) * ls) * ls)[..., None] # [H, M, C, 1]
k_last = kp * to_last
last_decay = torch.exp(cum[..., -1:])[..., None] # [H, M, 1, 1]
cont_c = cont.reshape(M, C, 1)
carried = s0
for i in range(M):
c = cont_c[i]
v_new = new_v[:, i] - c * torch.matmul(k_cd[:, i], carried) - (1.0 - c) * torch.matmul(k_cd[:, i], s0)
outs.append(c * torch.matmul(q_dec[:, i], carried) + (1.0 - c) * torch.matmul(q_dec[:, i], s0)
+ torch.matmul(attn[:, i], v_new))
if i + 1 < M:
cl = cont_c[i, -1:]
init = cl * carried + (1.0 - cl) * s0
carried = init * last_decay[:, i] + torch.matmul(k_last[:, i].transpose(-1, -2), v_new)
core = torch.cat(outs, dim=1).transpose(0, 1)[None]
core = net.norm(core, z)
return net.out_proj(core.reshape(1, T, H * Dv))
def load_chunked(root: str, state_len: int, packed_len: int) -> EosChunked:
cfg = TextConfig(json.load(open(f"{root}/backbone/config.json")))
state = {k: v.float() for k, v in load_file(f"{root}/backbone/model.safetensors").items()}
model = EosChunked(cfg, state_len, packed_len)
model.decoder.load_merged(state)
model.norm.load_state_dict({"weight": state["norm.weight"]})
model.embed_tokens = nn.Embedding.from_pretrained(state["embed_tokens.weight"], freeze=True)
model.load_state_dict({k: v.float() for k, v in load_file(f"{root}/decision_head.safetensors").items()},
strict=False)
return model.eval()
def chunk_inputs(segment: np.ndarray, C: int = 64):
"""seg_chunks [M, C, C], cont [P], last_seg [M, C] from the packed segment mask [P, P]."""
P = segment.shape[0]
M = P // C
seg_chunks = np.stack([segment[i * C : (i + 1) * C, i * C : (i + 1) * C] for i in range(M)])
starts = np.argmax(segment > 0, axis=1) # first position of each token's question (segment row's first 1)
cont = (starts < (np.arange(P) // C) * C).astype(np.float32)
last_seg = np.stack([segment[(i + 1) * C - 1, i * C : (i + 1) * C] for i in range(M)])
return seg_chunks.astype(np.float32), cont, last_seg.astype(np.float32)