File size: 2,127 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
"""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()