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