Download scripts/quantization_experiment.py from devildasdf/NEXORA: direct link, hf CLI and curl.
- Browser
- Download file 2.13 kB
-
https://huggingface.co/devildasdf/NEXORA/resolve/main/scripts/quantization_experiment.py
- Command line
-
hf download hf://devildasdf/NEXORA/scripts/quantization_experiment.py
-
curl -L -o quantization_experiment.py https://huggingface.co/devildasdf/NEXORA/resolve/main/scripts/quantization_experiment.py
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() | |