from collections import defaultdict from statistics import mean def is_correct(record): return not record.get("error") and record.get("predicted") in record.get("valid_answers", []) def _breakdown(records, key): groups = defaultdict(lambda: {"correct": 0, "total": 0}) for r in records: groups[r[key]]["total"] += 1 groups[r[key]]["correct"] += int(is_correct(r)) return { k: {**v, "accuracy": v["correct"] / v["total"] if v["total"] else 0.0} for k, v in sorted(groups.items()) } def group_consistency(scored): """Share of contrast groups whose every scored variant is correct.""" groups = defaultdict(list) for r in scored: if r.get("group_id"): groups[r["group_id"]].append(is_correct(r)) if not groups: return None, 0 consistent = sum(all(hits) for hits in groups.values()) return consistent / len(groups), len(groups) def evaluate(records, high_conf_threshold=0.90): """Aggregate per-case records into a report. Failed calls (API errors, invalid labels) stay in the denominator and count as incorrect, so an unreliable provider cannot inflate its accuracy. """ scored = [r for r in records if r.get("scored", True)] ambiguous = [r for r in records if not r.get("scored", True)] answered = [r for r in scored if not r.get("error")] errors = [r for r in records if r.get("error")] correct = [r for r in scored if is_correct(r)] wrong_answered = [r for r in answered if not is_correct(r)] high_conf_errors = [ r for r in wrong_answered if r.get("predicted_probability") is not None and r["predicted_probability"] >= high_conf_threshold ] expected_probs = [ r["expected_probability"] for r in answered if r.get("expected_probability") is not None ] ambiguous_answered = [r for r in ambiguous if not r.get("error")] ambiguous_valid = [r for r in ambiguous_answered if is_correct(r)] ambiguous_max_probs = [ max(r["probabilities"].values()) for r in ambiguous_answered if r.get("probabilities") ] consistency, group_total = group_consistency(scored) report = { "scored_total": len(scored), "answered": len(answered), "correct": len(correct), "accuracy": len(correct) / len(scored) if scored else 0.0, "api_errors": len(errors), "high_confidence_threshold": high_conf_threshold, "high_confidence_errors": len(high_conf_errors), "high_confidence_error_rate": len(high_conf_errors) / len(scored) if scored else 0.0, "mean_expected_probability": mean(expected_probs) if expected_probs else None, "ambiguous_total": len(ambiguous), "ambiguous_valid_choice_rate": ( len(ambiguous_valid) / len(ambiguous_answered) if ambiguous_answered else None ), "ambiguous_mean_max_probability": ( mean(ambiguous_max_probs) if ambiguous_max_probs else None ), "by_category": _breakdown(scored, "category"), "by_difficulty": _breakdown(scored, "difficulty"), } if group_total: report["group_total"] = group_total report["group_consistency"] = consistency if any(r.get("domain") for r in scored): report["by_domain"] = _breakdown(scored, "domain") return report