Files
DS4Server/tools/mtplx-session-codec-reference.py
T

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