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