71 lines
3.4 KiB
Python
71 lines
3.4 KiB
Python
"""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])
|