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
+101
View File
@@ -140,6 +140,80 @@ def routing_tier_drift(
}
def prediction_agreement(
records: Sequence[dict[str, Any]],
actual: Sequence[int],
reference: Sequence[int],
candidate: Sequence[int],
) -> dict[str, Any]:
"""Report label drift from a reference runtime to a candidate artifact."""
if not (
len(records) == len(actual) == len(reference) == len(candidate)
and records
):
raise ValueError("prediction agreement needs equally sized, non-empty vectors")
agreements = [
expected == got for expected, got in zip(reference, candidate)
]
scorable_indexes = [
index
for index, record in enumerate(records)
if record.get("slice") != "vague-eval"
]
transitions = Counter()
disagreements = []
for index, agrees in enumerate(agreements):
if agrees:
continue
reference_correct = reference[index] == actual[index]
candidate_correct = candidate[index] == actual[index]
if reference_correct and not candidate_correct:
transition = "correctToIncorrect"
elif not reference_correct and candidate_correct:
transition = "incorrectToCorrect"
else:
transition = "differentIncorrectLabel"
transitions[transition] += 1
disagreements.append(
{
"promptHash": prompt_hash(records[index]["prompt"]),
"slice": records[index].get("slice", "unknown"),
"expected": LABELS[actual[index]],
"reference": LABELS[reference[index]],
"candidate": LABELS[candidate[index]],
"transition": transition,
}
)
by_slice = {}
for slice_name in sorted(
{record.get("slice", "unknown") for record in records}
):
indexes = [
index
for index, record in enumerate(records)
if record.get("slice", "unknown") == slice_name
]
by_slice[slice_name] = {
"records": len(indexes),
"labelAgreement": (
sum(agreements[index] for index in indexes) / len(indexes)
),
}
return {
"records": len(records),
"labelAgreement": sum(agreements) / len(agreements),
"scoredLabelAgreement": (
sum(agreements[index] for index in scorable_indexes)
/ len(scorable_indexes)
),
"transitionCounts": dict(sorted(transitions.items())),
"bySlice": by_slice,
"disagreements": disagreements,
}
def evaluate(args: argparse.Namespace) -> dict[str, Any]:
try:
import torch
@@ -171,6 +245,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
label_to_id = {label: index for index, label in enumerate(LABELS)}
tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True)
onnx_session = None
reference_model = None
if args.onnx_model is not None:
try:
import onnxruntime as ort
@@ -191,7 +266,15 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
device = torch.device("cpu")
model = None
runtime_name = "onnxruntime-cpu"
if args.compare_pytorch:
reference_model = AutoModelForSequenceClassification.from_pretrained(
args.model_dir,
local_files_only=True,
).to(device)
reference_model.eval()
else:
if args.compare_pytorch:
raise DataError("--compare-pytorch requires --onnx-model")
device = _device(torch, args.device)
model = AutoModelForSequenceClassification.from_pretrained(
args.model_dir, local_files_only=True
@@ -216,6 +299,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
probabilities: list[float] = []
confidences: list[str] = []
margins: list[float] = []
reference_predictions: list[int] = []
inference_batch_size = 1 if onnx_session is not None else args.batch_size
with torch.inference_mode():
for start in range(0, len(records), inference_batch_size):
@@ -226,6 +310,11 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
torch,
)
logits = predict_logits(encoded) / temperature
if reference_model is not None:
reference_logits = reference_model(**encoded).logits
reference_predictions.extend(
reference_logits.argmax(dim=-1).tolist()
)
distribution = torch.softmax(logits, dim=-1)
top = torch.topk(distribution, k=2, dim=-1)
batch_probabilities = top.values[:, 0].tolist()
@@ -419,6 +508,13 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
},
"misclassifications": misclassifications,
}
if reference_model is not None:
report["pytorchParity"] = prediction_agreement(
records,
actual,
reference_predictions,
predicted,
)
write_json(args.report, report)
return report
@@ -431,6 +527,11 @@ def build_parser() -> argparse.ArgumentParser:
type=Path,
help="score a fixed-shape ONNX artifact instead of the PyTorch checkpoint",
)
parser.add_argument(
"--compare-pytorch",
action="store_true",
help="include label-level drift from --model-dir when scoring ONNX",
)
parser.add_argument("--calibration", type=Path, default=DEFAULT_CALIBRATION)
parser.add_argument("--test", type=Path, default=DEFAULT_TEST)
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)