Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,70 @@
|
||||
"""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])
|
||||
Reference in New Issue
Block a user