NEXORA / scripts /quantization_experiment.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
2.13 kB
"""CPU dynamic INT8 experiment on the reference model only."""
from pathlib import Path
import copy
import json
import sys
import time
import warnings
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
import numpy as np
import torch
from safetensors.torch import load_file
from nexora.model import NexoraLM, ModelConfig
from nexora.training import batch
from nexora.evaluation import percentiles
def main():
torch.set_num_threads(4)
cfg = ModelConfig(**json.loads(Path("artifacts/tiny/config.json").read_text()))
model = NexoraLM(cfg).eval()
model.load_state_dict(load_file("artifacts/tiny/model.safetensors"))
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
quant = torch.ao.quantization.quantize_dynamic(copy.deepcopy(model), {torch.nn.Linear}, dtype=torch.qint8)
ids = torch.from_numpy(np.load("artifacts/data/validation.npy").astype(np.int64))
x, y = batch(ids, 4, 128, torch.Generator().manual_seed(9), "cpu")
results = {}
logits = {}
for name, m in [("fp32", model), ("dynamic_int8_linear_only", quant)]:
times = []
with torch.no_grad():
m(x, y)
for _ in range(20):
start = time.perf_counter()
output, loss = m(x, y)
times.append(time.perf_counter()-start)
logits[name] = output
results[name] = {"validation_loss": loss.item(), "batch_latency_seconds": percentiles(times), "samples": len(times)}
results["max_logit_difference"] = (logits["fp32"]-logits["dynamic_int8_linear_only"]).abs().max().item()
results["argmax_agreement"] = (logits["fp32"].argmax(-1) == logits["dynamic_int8_linear_only"].argmax(-1)).float().mean().item()
results["limitations"] = "Tiny model, one validation batch, CPU linear-layer INT8 only. Not INT8 embeddings, BF16/FP8/INT4 deployment, coding/reasoning/context/tool quality validation. Do not infer production speedup."
Path("reports/quantization.json").write_text(json.dumps(results, indent=2))
print(json.dumps(results))
if __name__ == "__main__":
main()