umm / scripts /diagnose_vstar_eval.py
Orangerl's picture
Update Stage-2 code, evaluations, handoff, and deployment skill
bca45b0 verified
Raw History Blame Contribute Delete
17.7 kB
#!/usr/bin/env python3
"""Audit V* benchmark preparation and analyze errors against target scale."""
from __future__ import annotations
import argparse
import csv
import json
import math
import re
from collections import Counter, defaultdict
from pathlib import Path
from statistics import mean, median
from PIL import Image
def read_jsonl(path: Path) -> list[dict]:
return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line]
def quantile(values: list[float], probability: float) -> float:
ordered = sorted(values)
if not ordered:
raise ValueError("Cannot compute a quantile of an empty list")
position = (len(ordered) - 1) * probability
lower = math.floor(position)
upper = math.ceil(position)
if lower == upper:
return float(ordered[lower])
return float(ordered[lower] * (upper - position) + ordered[upper] * (position - lower))
def summary(values: list[float]) -> dict[str, float]:
return {
"min": min(values),
"p10": quantile(values, 0.10),
"p25": quantile(values, 0.25),
"median": median(values),
"mean": mean(values),
"p75": quantile(values, 0.75),
"p90": quantile(values, 0.90),
"max": max(values),
}
def normalize_text(value: str) -> str:
return re.sub(r"[^a-z0-9]+", " ", value.lower()).strip()
def accuracy_group(rows: list[dict]) -> dict:
samples = len(rows)
correct = sum(bool(row["correct"]) for row in rows)
accuracy = correct / samples if rows else None
lower = None
upper = None
if samples:
z = 1.959963984540054
denominator = 1 + z * z / samples
center = (accuracy + z * z / (2 * samples)) / denominator
margin = (
z
* math.sqrt(accuracy * (1 - accuracy) / samples + z * z / (4 * samples * samples))
/ denominator
)
lower = center - margin
upper = center + margin
return {
"samples": samples,
"correct": correct,
"accuracy": accuracy,
"wilson_95ci": [lower, upper],
}
def grouped(rows: list[dict], key: str) -> dict[str, dict]:
groups: dict[str, list[dict]] = defaultdict(list)
for row in rows:
groups[str(row[key])].append(row)
return {name: accuracy_group(group) for name, group in sorted(groups.items())}
def binned(rows: list[dict], key: str, boundaries: list[float]) -> dict[str, dict]:
groups: dict[str, list[dict]] = defaultdict(list)
for row in rows:
value = float(row[key])
lower = -math.inf
label = ""
for upper in boundaries:
if value <= upper:
label = f"({lower if math.isfinite(lower) else '-inf'}, {upper}]"
break
lower = upper
if not label:
label = f"({boundaries[-1]}, inf)"
groups[label].append(row)
return {name: accuracy_group(group) for name, group in groups.items()}
def mcnemar_exact(discordant_a_only: int, discordant_b_only: int) -> float:
trials = discordant_a_only + discordant_b_only
if trials == 0:
return 1.0
lower = min(discordant_a_only, discordant_b_only)
probability = sum(math.comb(trials, index) for index in range(lower + 1)) / (2**trials)
return min(1.0, 2 * probability)
def paired_comparison(rows: list[dict], first_key: str, second_key: str) -> dict:
first_only = sum(bool(row[first_key]) and not bool(row[second_key]) for row in rows)
second_only = sum(not bool(row[first_key]) and bool(row[second_key]) for row in rows)
both_correct = sum(bool(row[first_key]) and bool(row[second_key]) for row in rows)
both_wrong = sum(not bool(row[first_key]) and not bool(row[second_key]) for row in rows)
first_accuracy = sum(bool(row[first_key]) for row in rows) / len(rows)
second_accuracy = sum(bool(row[second_key]) for row in rows) / len(rows)
return {
"samples": len(rows),
"first_accuracy": first_accuracy,
"second_accuracy": second_accuracy,
"difference_percentage_points": (second_accuracy - first_accuracy) * 100,
"both_correct": both_correct,
"both_wrong": both_wrong,
"first_only_correct": first_only,
"second_only_correct": second_only,
"mcnemar_exact_two_sided_p": mcnemar_exact(first_only, second_only),
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--manifest", type=Path, required=True)
parser.add_argument("--results", type=Path, required=True)
parser.add_argument("--dynamic-results", type=Path)
parser.add_argument("--oracle-results", type=Path)
parser.add_argument("--raw-root", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument("--eval-width", type=int, default=448)
parser.add_argument("--eval-height", type=int, default=448)
args = parser.parse_args()
manifest = read_jsonl(args.manifest)
results = read_jsonl(args.results)
manifest_by_id = {str(row["sample_id"]): row for row in manifest}
results_by_id = {str(row["sample_id"]): row for row in results}
if len(manifest_by_id) != len(manifest) or len(results_by_id) != len(results):
raise ValueError("Duplicate sample IDs detected")
if manifest_by_id.keys() != results_by_id.keys():
raise ValueError("Manifest/result sample IDs differ")
diagnostic_results = {}
for name, path in (("dynamic", args.dynamic_results), ("oracle", args.oracle_results)):
if path is None:
continue
experiment_rows = read_jsonl(path)
experiment_by_id = {str(row["sample_id"]): row for row in experiment_rows}
if len(experiment_by_id) != len(experiment_rows):
raise ValueError(f"Duplicate sample IDs in {name} results")
if experiment_by_id.keys() != manifest_by_id.keys():
raise ValueError(f"Manifest/{name} result sample IDs differ")
diagnostic_results[name] = experiment_by_id
rows: list[dict] = []
mapping_failures: list[dict] = []
dimension_mismatches: list[dict] = []
for sample_id in sorted(manifest_by_id, key=int):
item = manifest_by_id[sample_id]
result = results_by_id[sample_id]
image_path = Path(item["images"][0])
annotation_path = args.raw_root / str(item["config"]) / f"{image_path.stem}.json"
annotation = json.loads(annotation_path.read_text(encoding="utf-8"))
with Image.open(image_path) as image:
width, height = image.size
recorded_width, recorded_height = item["image_sizes"][0]
if (width, height) != (recorded_width, recorded_height):
dimension_mismatches.append(
{
"sample_id": sample_id,
"actual": [width, height],
"manifest": [recorded_width, recorded_height],
}
)
expected_index = ord(str(item["answer"]).upper()) - ord("A")
expected_choice = str(item["choices"][expected_index])
official_correct_sentence = str(annotation["options"][0])
choice_norm = normalize_text(expected_choice)
sentence_norm = normalize_text(official_correct_sentence)
mapping_valid = choice_norm in sentence_norm or sentence_norm in choice_norm
if not mapping_valid:
mapping_failures.append(
{
"sample_id": sample_id,
"choice": expected_choice,
"official_correct_sentence": official_correct_sentence,
}
)
scale_x = args.eval_width / width
scale_y = args.eval_height / height
boxes = annotation["bbox"]
projected_boxes = [
[float(x) * scale_x, float(y) * scale_y, float(w) * scale_x, float(h) * scale_y]
for x, y, w, h in boxes
]
min_projected_side = min(min(box[2], box[3]) for box in projected_boxes)
min_projected_area = min(box[2] * box[3] for box in projected_boxes)
min_area_fraction = min((float(w) * float(h)) / (width * height) for _, _, w, h in boxes)
target_centers = [
((float(x) + float(w) / 2) / width, (float(y) + float(h) / 2) / height)
for x, y, w, h in boxes
]
max_center_distance = max(
math.sqrt(((cx - 0.5) / 0.5) ** 2 + ((cy - 0.5) / 0.5) ** 2)
for cx, cy in target_centers
)
predicted = result.get("predicted")
predicted_choice = None
if predicted:
predicted_index = ord(str(predicted).upper()) - ord("A")
if 0 <= predicted_index < len(item["choices"]):
predicted_choice = item["choices"][predicted_index]
question = str(item["question"]).lower()
if "color" in question:
question_type = "color"
elif "material" in question:
question_type = "material"
elif str(item["config"]) == "relative_position":
question_type = "relative_position"
elif "kind of" in question or "what animal" in question:
question_type = "object_kind"
else:
question_type = "other_attribute"
output_row = {
"sample_id": sample_id,
"config": item["config"],
"question_type": question_type,
"question": item["question"],
"image_path": str(image_path),
"image_width": width,
"image_height": height,
"image_megapixels": width * height / 1_000_000,
"target_count": len(boxes),
"target_objects": " | ".join(annotation["target_object"]),
"min_target_area_fraction": min_area_fraction,
"min_projected_side_448": min_projected_side,
"min_projected_area_448": min_projected_area,
"max_center_distance": max_center_distance,
"choice_count": len(item["choices"]),
"expected": result["expected"],
"expected_choice": expected_choice,
"predicted": predicted,
"predicted_choice": predicted_choice,
"correct": bool(result["correct"]),
"mapping_valid": mapping_valid,
"error": result.get("error"),
}
for experiment_name, experiment_by_id in diagnostic_results.items():
experiment = experiment_by_id[sample_id]
output_row[f"{experiment_name}_predicted"] = experiment.get("predicted")
output_row[f"{experiment_name}_correct"] = bool(experiment.get("correct"))
output_row[f"{experiment_name}_error"] = experiment.get("error")
rows.append(output_row)
labels = Counter(row["expected"] for row in rows)
predictions = Counter(row["predicted"] for row in rows)
confusion: dict[str, Counter] = defaultdict(Counter)
for row in rows:
confusion[str(row["expected"])][str(row["predicted"])] += 1
task_groups: dict[str, list[dict]] = defaultdict(list)
for row in rows:
task_groups[str(row["config"])].append(row)
chance_by_task = {
task: mean(1 / int(row["choice_count"]) for row in task_rows)
for task, task_rows in sorted(task_groups.items())
}
overall_random_chance = mean(1 / int(row["choice_count"]) for row in rows)
tiny_threshold_counts = {}
for threshold in (4, 8, 12, 16, 24, 32):
subset = [row for row in rows if float(row["min_projected_side_448"]) <= threshold]
tiny_threshold_counts[f"le_{threshold}px"] = accuracy_group(subset)
scale_by_config = {}
for config, config_rows in sorted(task_groups.items()):
correct_rows = [row for row in config_rows if row["correct"]]
wrong_rows = [row for row in config_rows if not row["correct"]]
scale_by_config[config] = {
"accuracy_by_projected_min_side_bin": binned(
config_rows, "min_projected_side_448", [4, 8, 12, 16, 24, 32]
),
"min_projected_side_correct": summary(
[float(row["min_projected_side_448"]) for row in correct_rows]
),
"min_projected_side_wrong": summary(
[float(row["min_projected_side_448"]) for row in wrong_rows]
),
}
wrong_smallest = sorted(
(row for row in rows if not row["correct"]),
key=lambda row: float(row["min_projected_side_448"]),
)[:20]
correct_smallest = sorted(
(row for row in rows if row["correct"]),
key=lambda row: float(row["min_projected_side_448"]),
)[:10]
report = {
"scope": {
"samples": len(rows),
"manifest": str(args.manifest.resolve()),
"results": str(args.results.resolve()),
"raw_root": str(args.raw_root.resolve()),
"simulated_eval_size": [args.eval_width, args.eval_height],
},
"integrity": {
"unique_manifest_ids": len(manifest_by_id),
"unique_result_ids": len(results_by_id),
"id_sets_equal": manifest_by_id.keys() == results_by_id.keys(),
"dimension_mismatches": dimension_mismatches,
"label_mapping_failures": mapping_failures,
"inference_errors": sum(bool(row["error"]) for row in rows),
"unparsed": sum(row["predicted"] is None for row in rows),
},
"performance": {
"overall": accuracy_group(rows),
"by_config": grouped(rows, "config"),
"by_question_type": grouped(rows, "question_type"),
"random_choice_chance": overall_random_chance,
"random_choice_chance_by_config": chance_by_task,
"expected_label_distribution": dict(sorted(labels.items())),
"predicted_label_distribution": dict(sorted(predictions.items())),
"confusion": {
expected: dict(sorted(predicted.items()))
for expected, predicted in sorted(confusion.items())
},
},
"image_and_target_scale": {
"image_width": summary([float(row["image_width"]) for row in rows]),
"image_height": summary([float(row["image_height"]) for row in rows]),
"image_megapixels": summary([float(row["image_megapixels"]) for row in rows]),
"min_projected_target_side_448": summary(
[float(row["min_projected_side_448"]) for row in rows]
),
"min_projected_target_area_448": summary(
[float(row["min_projected_area_448"]) for row in rows]
),
"min_target_area_fraction": summary(
[float(row["min_target_area_fraction"]) for row in rows]
),
"accuracy_by_projected_min_side_bin": binned(
rows, "min_projected_side_448", [4, 8, 12, 16, 24, 32]
),
"accuracy_by_image_megapixels": binned(rows, "image_megapixels", [2, 3, 4, 6]),
"accuracy_by_target_center_distance": binned(
rows, "max_center_distance", [0.5, 0.8, 1.0, 1.2]
),
"cumulative_tiny_target_thresholds": tiny_threshold_counts,
"by_config": scale_by_config,
},
"smallest_wrong_samples": wrong_smallest,
"smallest_correct_samples": correct_smallest,
}
if diagnostic_results:
sensitivity = {}
for experiment_name in diagnostic_results:
correct_key = f"{experiment_name}_correct"
sensitivity[experiment_name] = {
"overall": accuracy_group(
[{**row, "correct": row[correct_key]} for row in rows]
),
"by_config": {
config: accuracy_group(
[{**row, "correct": row[correct_key]} for row in config_rows]
)
for config, config_rows in sorted(task_groups.items())
},
"accuracy_by_baseline_projected_min_side_bin": binned(
[{**row, "correct": row[correct_key]} for row in rows],
"min_projected_side_448",
[4, 8, 12, 16, 24, 32],
),
"errors": sum(bool(row[f"{experiment_name}_error"]) for row in rows),
}
if "dynamic" in diagnostic_results:
sensitivity["baseline_vs_dynamic"] = paired_comparison(
rows, "correct", "dynamic_correct"
)
if "oracle" in diagnostic_results:
sensitivity["baseline_vs_oracle"] = paired_comparison(
rows, "correct", "oracle_correct"
)
if "dynamic" in diagnostic_results and "oracle" in diagnostic_results:
sensitivity["dynamic_vs_oracle"] = paired_comparison(
rows, "dynamic_correct", "oracle_correct"
)
report["sensitivity_experiments"] = sensitivity
args.output_dir.mkdir(parents=True, exist_ok=True)
(args.output_dir / "diagnosis.json").write_text(
json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
fieldnames = list(rows[0])
with (args.output_dir / "sample_diagnostics.csv").open(
"w", encoding="utf-8", newline=""
) as handle:
writer = csv.DictWriter(handle, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(rows)
print(json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True))
if __name__ == "__main__":
main()