Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user