repro-gradient-flow-sampler-based-distributionally-robust-optimization / code /claim5_source_vector_audit.py
| #!/usr/bin/env python3 | |
| """Extract the seven plotted CIFAR curves from the pinned v1 PDF vectors.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import re | |
| import shutil | |
| import subprocess | |
| import tempfile | |
| import xml.etree.ElementTree as ET | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parents[1] | |
| COLORS = { | |
| "RGO": "rgb(12.156677%, 46.665955%, 70.587158%)", | |
| "WGF": "rgb(100%, 49.803162%, 5.490112%)", | |
| "SAA": "rgb(17.254639%, 62.744141%, 17.254639%)", | |
| "Dual": "rgb(83.920288%, 15.293884%, 15.686035%)", | |
| "WRM": "rgb(58.03833%, 40.391541%, 74.116516%)", | |
| "WFR": "rgb(54.901123%, 33.724976%, 29.411316%)", | |
| "SVG": "rgb(89.01825%, 46.665955%, 76.077271%)", | |
| } | |
| PDFS = { | |
| 0.2: ROOT / "cifar_pgd_epsilon_0.2.pdf", | |
| 0.02: ROOT / "cifar_pgd_epsilon_0.02.pdf", | |
| 0.002: ROOT / "cifar_pgd_epsilon_0.002.pdf", | |
| } | |
| PATH_RE = re.compile(r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][-+]?\d+)?") | |
| def sha256(path: Path) -> str: | |
| digest = hashlib.sha256() | |
| with path.open("rb") as handle: | |
| for block in iter(lambda: handle.read(1 << 20), b""): | |
| digest.update(block) | |
| return digest.hexdigest() | |
| def path_points(path_data: str) -> list[tuple[float, float]]: | |
| values = [float(token) for token in PATH_RE.findall(path_data)] | |
| return list(zip(values[0::2], values[1::2])) | |
| def extract(pdf: Path) -> dict[str, list[tuple[float, float]]]: | |
| converter = shutil.which("pdftocairo") or "/opt/homebrew/bin/pdftocairo" | |
| with tempfile.TemporaryDirectory(prefix="gradient-flow-svg-") as directory: | |
| svg = Path(directory) / "plot.svg" | |
| subprocess.run([converter, "-svg", str(pdf), str(svg)], check=True) | |
| root = ET.parse(svg).getroot() | |
| candidates: dict[str, list[list[tuple[float, float]]]] = {name: [] for name in COLORS} | |
| for path in root.iter("{http://www.w3.org/2000/svg}path"): | |
| if path.attrib.get("stroke-opacity") != "1": | |
| continue | |
| color = path.attrib.get("stroke") | |
| name = next((name for name, expected in COLORS.items() if color == expected), None) | |
| if name is None: | |
| continue | |
| points = path_points(path.attrib.get("d", "")) | |
| if len(points) >= 11: | |
| candidates[name].append(points) | |
| curves = {name: max(paths, key=len)[:11] for name, paths in candidates.items() if paths} | |
| if set(curves) != set(COLORS): | |
| raise AssertionError(f"expected all seven curve paths, got {sorted(curves)}") | |
| if any(len(points) != 11 for points in curves.values()): | |
| raise AssertionError("one or more curves did not have eleven perturbation points") | |
| return curves | |
| def main() -> None: | |
| panels = {} | |
| for epsilon, pdf in PDFS.items(): | |
| curves = extract(pdf) | |
| panels[str(epsilon)] = { | |
| "pdf_sha256": sha256(pdf), | |
| "points_per_method": {name: len(points) for name, points in curves.items()}, | |
| "local_y": {name: [round(point[1], 6) for point in points] for name, points in curves.items()}, | |
| } | |
| baselines = ("RGO", "SAA", "Dual", "WRM", "SVG") | |
| violations = [] | |
| for epsilon, panel in panels.items(): | |
| for index in range(11): | |
| for method in ("WGF", "WFR"): | |
| beating_baselines = [ | |
| baseline for baseline in baselines if panel["local_y"][baseline][index] > panel["local_y"][method][index] | |
| ] | |
| if beating_baselines: | |
| violations.append( | |
| { | |
| "epsilon": float(epsilon), | |
| "delta_index": index, | |
| "method": method, | |
| "baseline_lower_test_error_by_vector_ordinate": beating_baselines, | |
| "method_local_y": panel["local_y"][method][index], | |
| "baseline_local_y": {baseline: panel["local_y"][baseline][index] for baseline in beating_baselines}, | |
| } | |
| ) | |
| direct = next( | |
| violation | |
| for violation in violations | |
| if violation["epsilon"] == 0.2 and violation["delta_index"] == 0 and violation["method"] == "WGF" | |
| ) | |
| breakdown = {} | |
| for epsilon in panels: | |
| panel_rows = [row for row in violations if row["epsilon"] == float(epsilon)] | |
| breakdown[epsilon] = { | |
| "total_violations": len(panel_rows), | |
| "by_method": { | |
| method: sum(row["method"] == method for row in panel_rows) | |
| for method in ("WGF", "WFR") | |
| }, | |
| "by_baseline": { | |
| baseline: sum(baseline in row["baseline_lower_test_error_by_vector_ordinate"] for row in panel_rows) | |
| for baseline in baselines | |
| }, | |
| "distinct_delta_indices": sorted({row["delta_index"] for row in panel_rows}), | |
| } | |
| result = { | |
| "panels": panels, | |
| "curve_count_per_panel": 7, | |
| "points_per_curve": 11, | |
| "violation_count": len(violations), | |
| "violation_breakdown": breakdown, | |
| "direct_counterexample": direct, | |
| "claim_falsified_by_pinned_vectors": bool(violations), | |
| "interpretation": "All curves share the PDF y-axis transform; a larger extracted local y ordinate is a lower plotted test error. At epsilon=0.2 and Delta=0, Dual is above both WGF and WFR, so the claimed consistent dominance is false.", | |
| } | |
| print(json.dumps(result, indent=2, sort_keys=True)) | |
| if not result["claim_falsified_by_pinned_vectors"]: | |
| raise SystemExit("no source-vector counterexample found") | |
| if __name__ == "__main__": | |
| main() | |