vtava's picture
Upload validated TinyCeNN standalone release
e9f1128 verified
Raw History Blame Contribute Delete
1.14 kB
from pathlib import Path
import json
import torch
from load_model import load_model
root = Path(__file__).resolve().parent
meta = json.loads((root / "standalone_config.json").read_text())
probe = json.loads((root / "validation_probe.json").read_text())
model, tokenizer = load_model(root)
classes = [m.__class__.__name__ for m in model.modules()]
assert meta["expected_class"] in classes, f"Missing custom class: {meta['expected_class']}"
import tinycenn_lm
runtime = Path(tinycenn_lm.__file__).resolve()
assert str(runtime).startswith(str(root)), f"External TinyCeNN runtime used: {runtime}"
ids = tokenizer(probe["prompt"], return_tensors="pt").input_ids.to(next(model.parameters()).device)
with torch.inference_mode():
out = model(input_ids=ids, use_cache=False, return_dict=True)
assert torch.isfinite(out.logits).all()
actual = int(out.logits[0, -1].argmax())
assert actual == int(probe["next_token_id"]), (actual, probe["next_token_id"])
print("STANDALONE_RELOAD_PASS")
print("architecture=", meta["expected_class"])
print("layers=", meta["layers"])
print("runtime=", runtime)
print("next_token=", tokenizer.decode([actual]))