"""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)