Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 17:28:44 -07:00
parent 93e9c838bb
commit e4403b570d
7 changed files with 408 additions and 19 deletions
+83 -3
View File
@@ -5,15 +5,23 @@ from __future__ import annotations
import argparse
import hashlib
import math
import shutil
import sys
import tempfile
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any, Sequence
import numpy as np
from purpose_data import DataError, load_jsonl, normalize_prompt, write_json
from purpose_data import (
DataError,
load_jsonl,
normalize_prompt,
prompt_hash,
write_json,
)
from train import (
HEAD_TOKENS,
MAX_LENGTH,
@@ -41,6 +49,59 @@ GOLDEN_PROMPTS = (
)
def stratified_calibration_sample(
records: Sequence[dict[str, Any]],
count: int,
*,
seed: int,
) -> list[dict[str, Any]]:
"""Select an exact, deterministic purpose/slice/language calibration sample."""
if not records:
raise DataError("cannot calibrate quantization from an empty validation split")
if count <= 0:
raise DataError("calibration record count must be positive")
target = min(count, len(records))
groups: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list)
for record in records:
language = str(record.get("lang", "unknown")).split("-", 1)[0].casefold()
key = (
str(record.get("purpose", "unknown")),
str(record.get("slice", "unknown")),
language,
)
groups[key].append(record)
allocations = {}
remainders = []
allocated = 0
for key in sorted(groups):
quota = len(groups[key]) * target / len(records)
base = math.floor(quota)
allocations[key] = base
allocated += base
tie_break = hashlib.sha256(f"{seed}\0{key}".encode("utf-8")).hexdigest()
remainders.append((quota - base, tie_break, key))
for _, _, key in sorted(remainders, reverse=True)[: target - allocated]:
allocations[key] += 1
selected = []
for key in sorted(groups):
ranked = sorted(
groups[key],
key=lambda record: hashlib.sha256(
f"{seed}\0{prompt_hash(record['prompt'])}".encode("utf-8")
).hexdigest(),
)
selected.extend(ranked[: allocations[key]])
return sorted(
selected,
key=lambda record: hashlib.sha256(
f"{seed + 1}\0{prompt_hash(record['prompt'])}".encode("utf-8")
).hexdigest(),
)
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
@@ -216,11 +277,16 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
onnx.save(fp16_model, fp16_path)
validation = load_jsonl(args.validation)
calibration_samples = stratified_calibration_sample(
validation,
args.calibration_records,
seed=args.calibration_seed,
)
class Reader(CalibrationDataReader):
def __init__(self) -> None:
self.index = 0
self.samples = validation[: args.calibration_records]
self.samples = calibration_samples
def get_next(self) -> dict[str, np.ndarray] | None:
if self.index >= len(self.samples):
@@ -281,7 +347,20 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
"opset": args.opset,
"fixedInputShape": [1, MAX_LENGTH],
"inputNames": input_names,
"calibrationRecords": min(args.calibration_records, len(validation)),
"calibrationRecords": len(calibration_samples),
"calibrationSeed": args.calibration_seed,
"calibrationSample": {
"strategy": "stratified by purpose, slice, and primary language",
"purposeCounts": dict(
sorted(Counter(item["purpose"] for item in calibration_samples).items())
),
"sliceCounts": dict(
sorted(Counter(item["slice"] for item in calibration_samples).items())
),
"promptHashes": sorted(
prompt_hash(item["prompt"]) for item in calibration_samples
),
},
"shippingArtifact": "int8QDQ",
"shippingBudgetBytes": args.shipping_budget_bytes,
"shippingBudgetPassed": int8_size <= args.shipping_budget_bytes,
@@ -303,6 +382,7 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
parser.add_argument("--opset", type=int, default=17)
parser.add_argument("--calibration-records", type=int, default=256)
parser.add_argument("--calibration-seed", type=int, default=20260730)
parser.add_argument(
"--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES
)