Download testing/forward_oracle.py from PYTHAI/bankml: direct link, hf CLI and curl.
- Browser
- Download file 20.2 kB
-
https://huggingface.co/spaces/PYTHAI/bankml/resolve/main/testing/forward_oracle.py
- Command line
-
hf download hf://spaces/PYTHAI/bankml/testing/forward_oracle.py
-
curl -L -o forward_oracle.py https://huggingface.co/spaces/PYTHAI/bankml/resolve/main/testing/forward_oracle.py
20.2 kB
| #!/usr/bin/env python3 | |
| # SPDX-License-Identifier: MIT OR Apache-2.0 | |
| """P3's oracle for the first operations of the forward pass, from the SHIPPED ggml (llama.cpp b11192): the same ops | |
| llama.cpp's Qwen3 graph uses — get_rows on the Q1_0 token embedding (inp_embd), then rms_norm(eps) and mul by | |
| blk.0.attn_norm.weight (attn_norm-0) — built as a ggml graph through ctypes and computed by the release's own CPU | |
| backend (libggml-cpu-haswell.so), on real token ids. Writes one line per token: id, sha256 of the embedding row's f32 | |
| bytes, sha256 of the normed row's f32 bytes, and the first four normed values (for eyes). | |
| usage: python3 testing/forward_oracle.py GGUF LIBDIR [OUT] (OUT default .models/oracle-forward)""" | |
| import ctypes as C, hashlib, mmap, os, struct, sys | |
| from pathlib import Path | |
| gguf, libdir = sys.argv[1], sys.argv[2] | |
| out = Path(sys.argv[3] if len(sys.argv) > 3 else Path(__file__).resolve().parents[1] / ".models" / "oracle-forward") | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| from gguf_guard import parse # noqa: E402 | |
| base = C.CDLL(os.path.join(libdir, "libggml-base.so"), mode=C.RTLD_GLOBAL) | |
| cpu = C.CDLL(os.path.join(libdir, "libggml-cpu-haswell.so"), mode=C.RTLD_GLOBAL) | |
| cpu.ggml_cpu_init() | |
| class InitParams(C.Structure): | |
| _fields_ = [("mem_size", C.c_size_t), ("mem_buffer", C.c_void_p), ("no_alloc", C.c_bool)] | |
| P = C.c_void_p | |
| base.ggml_init.restype = P | |
| base.ggml_init.argtypes = [InitParams] | |
| for f, args in (("ggml_mul_mat", [P, P, P]), ("ggml_new_tensor_1d", [P, C.c_int, C.c_int64]), ("ggml_new_tensor_2d", [P, C.c_int, C.c_int64, C.c_int64]), | |
| ("ggml_get_rows", [P, P, P]), ("ggml_rms_norm", [P, P, C.c_float]), ("ggml_mul", [P, P, P]), ("ggml_new_graph", [P]), | |
| ("ggml_get_data", [P])): | |
| getattr(base, f).restype = P | |
| getattr(base, f).argtypes = args | |
| base.ggml_build_forward_expand.argtypes = [P, P] | |
| base.ggml_nbytes.restype = C.c_size_t | |
| base.ggml_nbytes.argtypes = [P] | |
| base.ggml_free.argtypes = [P] | |
| cpu.ggml_graph_compute_with_ctx.argtypes = [P, P, C.c_int] | |
| F32, I32, Q1_0 = 0, 26, 41 | |
| fh = open(gguf, "rb") | |
| mm = mmap.mmap(fh.fileno(), 0, access=mmap.ACCESS_READ) | |
| ver, kv, tensors, align, data_start = parse(mm[:64 << 20]) | |
| T = {t["name"]: t for t in tensors} | |
| eps = float(kv["qwen3.attention.layer_norm_rms_epsilon"]) | |
| emb, norm = T["token_embd.weight"], T["blk.0.attn_norm.weight"] | |
| n_embd, n_vocab = emb["dims"] if "dims" in emb else emb["ne"] | |
| assert emb["type"] == Q1_0 and norm["type"] == F32 | |
| emb_bytes = n_vocab * (n_embd // 128) * 18 | |
| # the ids: a real prompt through the tokenizer's oracle cases, the chat markers, and the table's edges | |
| ids = [0, 1, 2, 13, 198, 220, 271, 872, 13048, 151643, 151644, 151645, 151667, 151668, n_vocab - 1] | |
| cases = Path(__file__).resolve().parents[1] / ".models" / "oracle-tokenizer" / "cases.jsonl" | |
| if cases.exists(): | |
| import json | |
| for line in cases.read_text().splitlines()[:40]: | |
| ids += json.loads(line)["ids"][:24] | |
| ids = list(dict.fromkeys(ids)) # unique, in order | |
| ctx = base.ggml_init(InitParams(emb_bytes + n_embd * 4 * (3 * len(ids) + 8) + (64 << 20), None, False)) | |
| w = base.ggml_new_tensor_2d(ctx, Q1_0, n_embd, n_vocab) | |
| C.memmove(base.ggml_get_data(w), C.c_char_p(mm[data_start + emb["offset"]: data_start + emb["offset"] + emb_bytes]), emb_bytes) | |
| nw = base.ggml_new_tensor_1d(ctx, F32, n_embd) | |
| C.memmove(base.ggml_get_data(nw), C.c_char_p(mm[data_start + norm["offset"]: data_start + norm["offset"] + n_embd * 4]), n_embd * 4) | |
| tok = base.ggml_new_tensor_1d(ctx, I32, len(ids)) | |
| C.memmove(base.ggml_get_data(tok), (C.c_int32 * len(ids))(*ids), 4 * len(ids)) | |
| inp = base.ggml_get_rows(ctx, w, tok) # inp_embd | |
| cur = base.ggml_mul(ctx, base.ggml_rms_norm(ctx, inp, C.c_float(eps)), nw) # attn_norm-0 | |
| gf = base.ggml_new_graph(ctx) | |
| base.ggml_build_forward_expand(gf, cur) | |
| cpu.ggml_graph_compute_with_ctx(ctx, gf, 3) | |
| row = n_embd * 4 | |
| a = C.string_at(base.ggml_get_data(inp), row * len(ids)) | |
| b = C.string_at(base.ggml_get_data(cur), row * len(ids)) | |
| out.mkdir(parents=True, exist_ok=True) | |
| with open(out / "embed_norm.tsv", "w") as f: | |
| for k, i in enumerate(ids): | |
| ra, rb = a[k * row:(k + 1) * row], b[k * row:(k + 1) * row] | |
| first = " ".join(f"{x:.6g}" for x in struct.unpack("<4f", rb[:16])) | |
| f.write(f"{i}\t{hashlib.sha256(ra).hexdigest()}\t{hashlib.sha256(rb).hexdigest()}\t{first}\n") | |
| base.ggml_free(ctx) | |
| print(f"{len(ids)} tokens: inp_embd and attn_norm-0 rows from the shipped ggml (eps {eps:g}, 3 threads) → {out / 'embed_norm.tsv'}") | |
| # ---- step four: layer 0's Q, K, V — projections, per-head RMS norms, YaRN RoPE (NEOX), as llama.cpp's qwen3 graph ---- | |
| # llama.cpp's own rope parameters for this model (llama-context.cpp, b11192), computed in float as it computes them | |
| libm = C.CDLL("libm.so.6") | |
| libm.logf.restype = C.c_float | |
| libm.logf.argtypes = [C.c_float] | |
| f32 = lambda x: struct.unpack("<f", struct.pack("<f", x))[0] # noqa: E731 | |
| factor = f32(float(kv.get("qwen3.rope.scaling.factor", 1.0))) | |
| freq_scale = f32(1.0 / factor) | |
| logfac = libm.logf(f32(1.0 / freq_scale)) | |
| mscale = f32(f32(f32(0.1 * 1.0) * logfac) + 1.0) if factor > 1.0 else 1.0 # get_mscale(factor, 1) | |
| attn_factor = f32(mscale * f32(1.0 / f32(1.0 + f32(0.1 * logfac)))) # *= 1/(1 + 0.1·logf(factor)) | |
| n_ctx_orig = int(kv.get("qwen3.rope.scaling.original_context_length", kv["qwen3.context_length"])) | |
| freq_base = float(kv["qwen3.rope.freq_base"]) | |
| n_head, n_head_kv, head = int(kv["qwen3.attention.head_count"]), int(kv["qwen3.attention.head_count_kv"]), int(kv["qwen3.attention.key_length"]) | |
| prompt = [151644, 8948, 198, 2610, 525, 20739, 5409, 13, 151645, 198, 151644, 872, 198, 3838, 1558, 264, 22725, 12118, 30, 151645, 198, | |
| 151644, 77091, 198, 151667, 271, 151668, 271] # "<|im_start|>system\nYou are Savante.<|im_end|>\n<|im_start|>user\n…" | |
| n = len(prompt) | |
| def tensor_2d(name, typ, ne0, ne1, bpr): | |
| t = T[name] | |
| b = mm[data_start + t["offset"]: data_start + t["offset"] + bpr * ne1] | |
| x = base.ggml_new_tensor_2d(ctx2, typ, ne0, ne1) | |
| C.memmove(base.ggml_get_data(x), C.c_char_p(b), len(b)) | |
| return x | |
| def tensor_1d(name, ne0): | |
| t = T[name] | |
| x = base.ggml_new_tensor_1d(ctx2, F32, ne0) | |
| C.memmove(base.ggml_get_data(x), C.c_char_p(mm[data_start + t["offset"]: data_start + t["offset"] + 4 * ne0]), 4 * ne0) | |
| return x | |
| for f, args in (("ggml_reshape_3d", [P, P, C.c_int64, C.c_int64, C.c_int64]), | |
| ("ggml_rope_ext", [P, P, P, P, C.c_int, C.c_int, C.c_int, C.c_float, C.c_float, C.c_float, C.c_float, C.c_float, C.c_float])): | |
| getattr(base, f).restype = P | |
| getattr(base, f).argtypes = args | |
| bpr = n_embd // 128 * 18 | |
| ctx2 = base.ggml_init(InitParams(emb_bytes + 2 * (n_embd * bpr) + 2 * (n_head_kv * head * bpr) + (256 << 20), None, False)) | |
| w2 = base.ggml_new_tensor_2d(ctx2, Q1_0, n_embd, n_vocab) | |
| C.memmove(base.ggml_get_data(w2), C.c_char_p(mm[data_start + emb["offset"]: data_start + emb["offset"] + emb_bytes]), emb_bytes) | |
| tok2 = base.ggml_new_tensor_1d(ctx2, I32, n) | |
| C.memmove(base.ggml_get_data(tok2), (C.c_int32 * n)(*prompt), 4 * n) | |
| pos = base.ggml_new_tensor_1d(ctx2, I32, n) | |
| C.memmove(base.ggml_get_data(pos), (C.c_int32 * n)(*range(n)), 4 * n) | |
| far = [7 + 2341 * i for i in range(n)] # out to 63 214: long-context positions, where theta is a long product | |
| pos_far = base.ggml_new_tensor_1d(ctx2, I32, n) | |
| C.memmove(base.ggml_get_data(pos_far), (C.c_int32 * n)(*far), 4 * n) | |
| xn = base.ggml_mul(ctx2, base.ggml_rms_norm(ctx2, base.ggml_get_rows(ctx2, w2, tok2), C.c_float(eps)), tensor_1d("blk.0.attn_norm.weight", n_embd)) | |
| outs = {} | |
| for nm, wname, nh, qn in (("Q", "blk.0.attn_q.weight", n_head, "blk.0.attn_q_norm.weight"), ("K", "blk.0.attn_k.weight", n_head_kv, "blk.0.attn_k_norm.weight"), | |
| ("V", "blk.0.attn_v.weight", n_head_kv, None)): | |
| proj = base.ggml_mul_mat(ctx2, tensor_2d(wname, Q1_0, n_embd, nh * head, bpr), xn) if hasattr(base, "ggml_mul_mat") else None | |
| outs[nm] = proj | |
| if qn: | |
| r = base.ggml_reshape_3d(ctx2, proj, head, nh, n) | |
| normed = base.ggml_mul(ctx2, base.ggml_rms_norm(ctx2, r, C.c_float(eps)), tensor_1d(qn, head)) | |
| outs[nm + "_normed"] = normed | |
| outs[nm + "_rope"] = base.ggml_rope_ext(ctx2, normed, pos, None, head, 2, n_ctx_orig, C.c_float(freq_base), C.c_float(freq_scale), | |
| C.c_float(1.0), C.c_float(attn_factor), C.c_float(32.0), C.c_float(1.0)) | |
| outs[nm + "_rope_far"] = base.ggml_rope_ext(ctx2, normed, pos_far, None, head, 2, n_ctx_orig, C.c_float(freq_base), C.c_float(freq_scale), | |
| C.c_float(1.0), C.c_float(attn_factor), C.c_float(32.0), C.c_float(1.0)) | |
| # ---- step five: layer 0's attention — the K/V cache in f16, causal mask, flash attention (the CPU reference path: | |
| # fewer than 64 query rows and fewer than 512 KV cells), then wo and the residual (llama.cpp's kqv_out, attn_out, ffn_inp) | |
| F16 = 1 | |
| for f, args in (("ggml_permute", [P, P, C.c_int, C.c_int, C.c_int, C.c_int]), ("ggml_cpy", [P, P, P]), | |
| ("ggml_new_tensor_3d", [P, C.c_int, C.c_int64, C.c_int64, C.c_int64]), ("ggml_reshape_2d", [P, P, C.c_int64, C.c_int64]), | |
| ("ggml_add", [P, P, P]), | |
| ("ggml_flash_attn_ext", [P, P, P, P, P, C.c_float, C.c_float, C.c_float])): | |
| getattr(base, f).restype = P | |
| getattr(base, f).argtypes = args | |
| base.ggml_prec_set_acc.argtypes = [P, C.c_int] | |
| k16 = base.ggml_cpy(ctx2, outs["K_rope"], base.ggml_new_tensor_3d(ctx2, F16, head, n_head_kv, n)) | |
| v16 = base.ggml_cpy(ctx2, base.ggml_reshape_3d(ctx2, outs["V"], head, n_head_kv, n), base.ggml_new_tensor_3d(ctx2, F16, head, n_head_kv, n)) | |
| mask = base.ggml_new_tensor_2d(ctx2, F16, n, n) | |
| C.memmove(base.ggml_get_data(mask), (C.c_uint16 * (n * n))(*[0 if j <= i else 0xFC00 for i in range(n) for j in range(n)]), 2 * n * n) | |
| kq_scale = f32(1.0 / f32(head ** 0.5)) | |
| fa = base.ggml_flash_attn_ext(ctx2, base.ggml_permute(ctx2, outs["Q_rope"], 0, 2, 1, 3), base.ggml_permute(ctx2, k16, 0, 2, 1, 3), | |
| base.ggml_permute(ctx2, v16, 0, 2, 1, 3), mask, C.c_float(kq_scale), C.c_float(0.0), C.c_float(0.0)) | |
| base.ggml_prec_set_acc(fa, 10) # GGML_PREC_F32, as llama-graph sets it | |
| kqv = base.ggml_reshape_2d(ctx2, fa, n_head * head, n) | |
| attn_out = base.ggml_mul_mat(ctx2, tensor_2d("blk.0.attn_output.weight", Q1_0, n_head * head, n_embd, (n_head * head) // 128 * 18), kqv) | |
| emb0 = base.ggml_get_rows(ctx2, w2, tok2) | |
| outs["kqv_out"] = kqv | |
| outs["attn_out"] = attn_out | |
| outs["ffn_inp"] = base.ggml_add(ctx2, attn_out, emb0) | |
| # ---- step six: layer 0's feed-forward block — ffn_norm, gate and up, SwiGLU (ggml's vectorized expf), down, the | |
| # residual: llama.cpp's ffn_norm, ffn_swiglu, ffn_out and l_out, the layer's output | |
| base.ggml_swiglu_split.restype = P | |
| base.ggml_swiglu_split.argtypes = [P, P, P] | |
| n_ff = int(kv["qwen3.feed_forward_length"]) | |
| fn = base.ggml_mul(ctx2, base.ggml_rms_norm(ctx2, outs["ffn_inp"], C.c_float(eps)), tensor_1d("blk.0.ffn_norm.weight", n_embd)) | |
| gate = base.ggml_mul_mat(ctx2, tensor_2d("blk.0.ffn_gate.weight", Q1_0, n_embd, n_ff, bpr), fn) | |
| up = base.ggml_mul_mat(ctx2, tensor_2d("blk.0.ffn_up.weight", Q1_0, n_embd, n_ff, bpr), fn) | |
| sw = base.ggml_swiglu_split(ctx2, gate, up) | |
| down = base.ggml_mul_mat(ctx2, tensor_2d("blk.0.ffn_down.weight", Q1_0, n_ff, n_embd, n_ff // 128 * 18), sw) | |
| outs["ffn_norm"] = fn | |
| outs["ffn_swiglu"] = sw | |
| outs["ffn_out"] = down | |
| outs["l_out"] = base.ggml_add(ctx2, down, outs["ffn_inp"]) | |
| gf2 = base.ggml_new_graph(ctx2) | |
| for k in outs: | |
| base.ggml_build_forward_expand(gf2, outs[k]) | |
| cpu.ggml_graph_compute_with_ctx(ctx2, gf2, 3) | |
| with open(out / "qkv.tsv", "w") as f: | |
| f.write(f"# attn_factor {attn_factor!r} bits {struct.unpack('<I', struct.pack('<f', attn_factor))[0]:08x} · freq_scale {freq_scale} · n_ctx_orig {n_ctx_orig} · freq_base {freq_base} · kq_scale bits {struct.unpack('<I', struct.pack('<f', kq_scale))[0]:08x} · far {' '.join(map(str, far))} · tokens {' '.join(map(str, prompt))}\n") | |
| for k in outs: | |
| width = n_ff if k == "ffn_swiglu" else n_embd if k in ("kqv_out", "attn_out", "ffn_inp", "ffn_norm", "ffn_out", "l_out") else (n_head if k.startswith("Q") else n_head_kv) * head | |
| raw = C.string_at(base.ggml_get_data(outs[k]), width * 4 * n) | |
| for p_ in range(n): | |
| row = raw[p_ * width * 4:(p_ + 1) * width * 4] | |
| f.write(f"{k}\t{p_}\t{hashlib.sha256(row).hexdigest()}\t{' '.join(f'{x:.6g}' for x in struct.unpack('<3f', row[:12]))}\n") | |
| print(f"layer 0 Q/K/V for {n} tokens (projections, normed, roped; attn_factor {attn_factor!r}) → {out / 'qkv.tsv'}") | |
| # ---- SwiGLU across the whole range: ggml's vectorized expf has a separate path for |x| > ~87 that a prompt never | |
| # reaches; sweep silu(x)·1 over ±120 (and the edges) through the shipped ggml_swiglu_split, every output's bits | |
| import random | |
| rnd = random.Random(11192) | |
| xs = [f32(-120 + 240 * i / 16383) for i in range(16384)] + [f32(rnd.uniform(-30, 30)) for _ in range(8192)] + \ | |
| [0.0, -0.0, 1e-30, -1e-30, 87.0, 88.5, 89.0, 103.9, 104.0, -87.0, -88.5, -89.0, -103.9, -104.0, -126.0, 126.0, 150.0, -150.0] | |
| xs += [0.0] * (-len(xs) % 8) | |
| m = len(xs) | |
| ctx3 = base.ggml_init(InitParams(64 << 20, None, False)) | |
| tx, tone = base.ggml_new_tensor_1d(ctx3, F32, m), base.ggml_new_tensor_1d(ctx3, F32, m) | |
| C.memmove(base.ggml_get_data(tx), (C.c_float * m)(*xs), 4 * m) | |
| C.memmove(base.ggml_get_data(tone), (C.c_float * m)(*([1.0] * m)), 4 * m) | |
| ty = base.ggml_swiglu_split(ctx3, tx, tone) | |
| gf3 = base.ggml_new_graph(ctx3) | |
| base.ggml_build_forward_expand(gf3, ty) | |
| cpu.ggml_graph_compute_with_ctx(ctx3, gf3, 3) | |
| ys = struct.unpack(f"<{m}I", C.string_at(base.ggml_get_data(ty), 4 * m)) | |
| with open(out / "swiglu.tsv", "w") as f: | |
| for x, y in zip(xs, ys): | |
| f.write(f"{struct.unpack('<I', struct.pack('<f', x))[0]:08x}\t{y:08x}\n") | |
| print(f"swiglu sweep: {m} values → {out / 'swiglu.tsv'}") | |
| # ---- step nine: a micro-batch of 64 rows or more takes ggml's TILED flash attention. Layer 0 on a 150-token prompt | |
| # (the chat prompt, then ids from the tokenizer's cases): projections, norms, RoPE, f16 K/V, causal mask, attention | |
| long_ids = list(prompt) | |
| for i in ids: | |
| if len(long_ids) >= 150: | |
| break | |
| if i < n_vocab and i not in (151643, 151645): | |
| long_ids.append(i) | |
| nl = len(long_ids) | |
| ctx2 = base.ggml_init(InitParams(emb_bytes + 4 * (n_embd * bpr) + (512 << 20), None, False)) | |
| w2 = base.ggml_new_tensor_2d(ctx2, Q1_0, n_embd, n_vocab) | |
| C.memmove(base.ggml_get_data(w2), C.c_char_p(mm[data_start + emb["offset"]: data_start + emb["offset"] + emb_bytes]), emb_bytes) | |
| tl = base.ggml_new_tensor_1d(ctx2, I32, nl) | |
| C.memmove(base.ggml_get_data(tl), (C.c_int32 * nl)(*long_ids), 4 * nl) | |
| pl = base.ggml_new_tensor_1d(ctx2, I32, nl) | |
| C.memmove(base.ggml_get_data(pl), (C.c_int32 * nl)(*range(nl)), 4 * nl) | |
| xl = base.ggml_mul(ctx2, base.ggml_rms_norm(ctx2, base.ggml_get_rows(ctx2, w2, tl), C.c_float(eps)), tensor_1d("blk.0.attn_norm.weight", n_embd)) | |
| def qk_long(wname, nname, nh): | |
| r = base.ggml_reshape_3d(ctx2, base.ggml_mul_mat(ctx2, tensor_2d(wname, Q1_0, n_embd, nh * head, bpr), xl), head, nh, nl) | |
| r = base.ggml_mul(ctx2, base.ggml_rms_norm(ctx2, r, C.c_float(eps)), tensor_1d(nname, head)) | |
| return base.ggml_rope_ext(ctx2, r, pl, None, head, 2, n_ctx_orig, C.c_float(freq_base), C.c_float(freq_scale), | |
| C.c_float(1.0), C.c_float(attn_factor), C.c_float(32.0), C.c_float(1.0)) | |
| ql = qk_long("blk.0.attn_q.weight", "blk.0.attn_q_norm.weight", n_head) | |
| kl = qk_long("blk.0.attn_k.weight", "blk.0.attn_k_norm.weight", n_head_kv) | |
| vl = base.ggml_reshape_3d(ctx2, base.ggml_mul_mat(ctx2, tensor_2d("blk.0.attn_v.weight", Q1_0, n_embd, n_head_kv * head, bpr), xl), head, n_head_kv, nl) | |
| k16 = base.ggml_cpy(ctx2, kl, base.ggml_new_tensor_3d(ctx2, F16, head, n_head_kv, nl)) | |
| v16 = base.ggml_cpy(ctx2, vl, base.ggml_new_tensor_3d(ctx2, F16, head, n_head_kv, nl)) | |
| maskl = base.ggml_new_tensor_2d(ctx2, F16, nl, nl) | |
| C.memmove(base.ggml_get_data(maskl), (C.c_uint16 * (nl * nl))(*[0 if j <= i else 0xFC00 for i in range(nl) for j in range(nl)]), 2 * nl * nl) | |
| fal = base.ggml_flash_attn_ext(ctx2, base.ggml_permute(ctx2, ql, 0, 2, 1, 3), base.ggml_permute(ctx2, k16, 0, 2, 1, 3), | |
| base.ggml_permute(ctx2, v16, 0, 2, 1, 3), maskl, C.c_float(kq_scale), C.c_float(0.0), C.c_float(0.0)) | |
| base.ggml_prec_set_acc(fal, 10) | |
| kqvl = base.ggml_reshape_2d(ctx2, fal, n_head * head, nl) | |
| gfl = base.ggml_new_graph(ctx2) | |
| base.ggml_build_forward_expand(gfl, kqvl) | |
| cpu.ggml_graph_compute_with_ctx(ctx2, gfl, 3) | |
| raw = C.string_at(base.ggml_get_data(kqvl), n_embd * 4 * nl) | |
| with open(out / "tiled.tsv", "w") as f: | |
| f.write(f"# layer 0 kqv_out, one micro-batch of {nl} rows (ggml's tiled flash attention) · tokens {' '.join(map(str, long_ids))}\n") | |
| for p_ in range(nl): | |
| row = raw[p_ * n_embd * 4:(p_ + 1) * n_embd * 4] | |
| f.write(f"kqv_out\t{p_}\t{hashlib.sha256(row).hexdigest()}\t{' '.join(f'{x:.6g}' for x in struct.unpack('<3f', row[:12]))}\n") | |
| print(f"tiled attention: layer 0 kqv_out for a {nl}-row micro-batch → {out / 'tiled.tsv'}") | |
| # ---- step ten: a single-token decode whose padded KV length reaches 512 takes ggml's SPLIT-KV flash attention: the | |
| # padded cells cut into one chunk per thread, a partial softmax per chunk, then a reduction. Layer 0's attention for | |
| # the last row of the 150-token prompt, repeated to longer contexts (K/V rows re-used cyclically, positions kept), | |
| # at padded lengths 512–1024 and at 3 and 4 threads (the result depends on the thread count) | |
| qrow = C.string_at(base.ggml_get_data(ql), 0) # (keep ql alive) | |
| kraw = C.string_at(base.ggml_get_data(k16), head * n_head_kv * 2 * nl) | |
| vraw = C.string_at(base.ggml_get_data(v16), head * n_head_kv * 2 * nl) | |
| qraw = C.string_at(base.ggml_get_data(ql), head * n_head * 4 * nl) | |
| with open(out / "splitkv.tsv", "w") as f: | |
| f.write(f"# layer 0 kqv_out for one query row (the last of the 150-token prompt) over n cells; cells cycle the prompt's K/V rows · tokens {' '.join(map(str, long_ids))}\n") | |
| for cells in (257, 300, 511, 512, 513, 700, 1000): | |
| padded = max(256, -(-cells // 256) * 256) | |
| for nth in (3, 4): | |
| c3 = base.ggml_init(InitParams(64 << 20, None, False)) | |
| qt = base.ggml_new_tensor_3d(c3, F32, head, 1, n_head) # [DK, n_tokens = 1, heads] | |
| qb = b"".join(qraw[((nl - 1) * n_head + h) * head * 4:((nl - 1) * n_head + h + 1) * head * 4] for h in range(n_head)) | |
| C.memmove(base.ggml_get_data(qt), C.c_char_p(qb), len(qb)) | |
| kt = base.ggml_new_tensor_3d(c3, F16, head, padded, n_head_kv) # [DK, n_kv, heads_kv] | |
| vt = base.ggml_new_tensor_3d(c3, F16, head, padded, n_head_kv) | |
| kb, vb = bytearray(head * padded * n_head_kv * 2), bytearray(head * padded * n_head_kv * 2) | |
| for hk in range(n_head_kv): | |
| for c in range(cells): | |
| src = ((c % nl) * n_head_kv + hk) * head * 2 | |
| dst = (hk * padded + c) * head * 2 | |
| kb[dst:dst + head * 2] = kraw[src:src + head * 2] | |
| vb[dst:dst + head * 2] = vraw[src:src + head * 2] | |
| C.memmove(base.ggml_get_data(kt), C.c_char_p(bytes(kb)), len(kb)) | |
| C.memmove(base.ggml_get_data(vt), C.c_char_p(bytes(vb)), len(vb)) | |
| mt = base.ggml_new_tensor_2d(c3, F16, padded, 1) | |
| C.memmove(base.ggml_get_data(mt), (C.c_uint16 * padded)(*[0 if j < cells else 0xFC00 for j in range(padded)]), 2 * padded) | |
| fa3 = base.ggml_flash_attn_ext(c3, qt, kt, vt, mt, C.c_float(kq_scale), C.c_float(0.0), C.c_float(0.0)) | |
| base.ggml_prec_set_acc(fa3, 10) | |
| g3 = base.ggml_new_graph(c3) | |
| base.ggml_build_forward_expand(g3, fa3) | |
| cpu.ggml_graph_compute_with_ctx(c3, g3, nth) | |
| row = C.string_at(base.ggml_get_data(fa3), n_embd * 4) | |
| f.write(f"kqv_out\t{cells}\t{nth}\t{hashlib.sha256(row).hexdigest()}\t{' '.join(f'{x:.6g}' for x in struct.unpack('<3f', row[:12]))}\n") | |
| base.ggml_free(c3) | |
| print(f"split-KV attention: one row over 257–1000 cells at 3 and 4 threads → {out / 'splitkv.tsv'}") | |