"""Pinned SessionBank codec fixture; no model weights, downloads or inference.""" import hashlib import json import sys from pathlib import Path import mlx.core as mx from mtplx.cache_bank.codec import encode_payload, decode_payload from mtplx.cache_state import CacheSnapshot def digest(value): if value is None: return None mx.eval(value) return {"shape": list(value.shape), "dtype": str(value.dtype).removeprefix("mlx.core."), "sha256": hashlib.sha256(bytes(memoryview(value))).hexdigest()} def state_digests(snapshot): return [None if state is None else [digest(v) for v in state] for state in snapshot.states] def main(directory): directory = Path(directory) directory.mkdir(parents=True, exist_ok=True) base = mx.arange(2 * 513 * 4, dtype=mx.float32).reshape(1, 2, 513, 4) / 16 bf16 = base.astype(mx.bfloat16) # Actual cache shapes can be strided, reversed, broadcast, empty or scalar. small = mx.arange(12, dtype=mx.int32).reshape(1, 3, 4) reversed_view = small[:, :, ::-1] broadcast = mx.broadcast_to(mx.array([7], dtype=mx.uint16), (1, 2, 3)) trunk = CacheSnapshot(states=([bf16, reversed_view, broadcast, mx.array([], dtype=mx.float16)], (base, base[:, :, ::-1, :], small.transpose(0, 2, 1), None), None), meta_states=("", None, None)) head = CacheSnapshot(states=((bf16[:, :, :5, :], base[:, :, :5, :], None, None),), meta_states=(None,)) boundary = CacheSnapshot(states=([bf16[:, :, :2, :], reversed_view], None, None), meta_states=("", None, None)) cases = [] for block in (0, 256, 1024): # Scalar codec quirk: decode preserves the one-element import vector. encoded = encode_payload(cache_snapshot=trunk, logits=mx.array(1.25), hidden=bf16[:, :1, :1, :], mtp_history_snapshot=head, gdn_boundaries=[(2, boundary, None)], has_recurrent=True, block_size=block) folder = directory / str(block) folder.mkdir(exist_ok=True) for name, raw in encoded.tensors.items(): (folder / (name + ".bin")).write_bytes(raw) decoded = decode_payload(encoded.spec, encoded.tensors.__getitem__) legacy = dict(encoded.spec) del legacy["gdn_boundaries"] legacy = decode_payload(legacy, encoded.tensors.__getitem__) assert not legacy.gdn_boundaries no_hidden = {**encoded.spec, "gdn_boundaries": [ {k: v for k, v in b.items() if k != "hidden_last"} for b in encoded.spec["gdn_boundaries"]]} assert decode_payload(no_hidden, encoded.tensors.__getitem__).gdn_boundaries[0][2] is None cases.append(dict(block_size=block, spec=encoded.spec, nbytes=encoded.nbytes, legacy_boundary_count=len(legacy.gdn_boundaries), blobs={k: hashlib.sha256(v).hexdigest() for k, v in encoded.tensors.items()}, decoded=dict(trunk=state_digests(decoded.cache_snapshot), mtp=state_digests(decoded.mtp_history_snapshot), logits=digest(decoded.logits), hidden=digest(decoded.hidden), boundaries=[dict(tokens=n, states=state_digests(s), hidden=digest(h)) for n, s, h in decoded.gdn_boundaries]))) print(json.dumps({"event": "codec_reference_case", "block_size": block, "blobs": len(encoded.tensors), "bytes": encoded.nbytes}), flush=True) (directory / "reference.json").write_text(json.dumps({"cases": cases})) if __name__ == "__main__": main(sys.argv[1])