| |
| """Deterministic, version-locked audit of gradient-flow-sampler DRO claims.""" |
| from __future__ import annotations |
|
|
| import argparse |
| import csv |
| import hashlib |
| import json |
| import math |
| from pathlib import Path |
|
|
| import matplotlib |
|
|
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
| import numpy as np |
|
|
|
|
| PAPER_ID = "QRtzkKrbJi" |
| V1_PDF_SHA256 = "796d3cb25e2a9062ab4e81201daa7837c4bf615399063b0ed439c9e17e6db8d6" |
| V1_SOURCE_SHA256 = "35a471bd60c11517db90b8e337064019c5a07faa621718aeca9800d425604488" |
| CURRENT_PDF_SHA256 = "88720453447d8affcefd3e63097c52f512d5b275766ba3e73dff9af2ca1c0f8f" |
| CURRENT_SOURCE_SHA256 = "cc00b5195a2592b24f9b876bf9fd27ada78fa9a46c468e81221983943cb726e5" |
| CLAIMS = [ |
| "The paper introduces a unified PDE gradient flow framework for distributionally robust optimization (DRO) with six concrete algorithms, including Wasserstein Gradient Flow (Algorithm 3) and Wasserstein Fisher-Rao flow (Algorithm 4) variants for entropy-regularized Wasserstein DRO (Section 4, Algorithms 3-4).", |
| "Proposition 1 shows the Wasserstein gradient flow sampler must run for time at least on the order of O((1/λ) log(L/√(λε))) to produce an ε-accurate gradient estimate (Section 4, Proposition 1).", |
| "Theorem 1 proves the outer loop of the gradient-flow-sampler-based DRO algorithm requires O(1/ε²_opt) iterations to reach an ε-stationary point (Section 5, Theorem 1).", |
| "Theorem 2 bounds the total computational complexity of the WGF-based DRO algorithm (Algorithm 3) as Õ(L_Φ L²_U L²_f d² / (λ³_U ε⁴_opt)) (Section 5, Theorem 2).", |
| "On CIFAR-10 adversarial training under PGD attacks, the WFR- and WGF-based DRO methods achieve consistently higher robust accuracy across all perturbation settings compared to baseline DRO methods (Section 6.3).", |
| "Lemma 1 establishes that the entropy-regularized DRO problem is equivalent to a Schrödinger half-bridge problem, enabling sampling from the conditional worst-case distribution (Section 3.1, Lemma 1).", |
| ] |
|
|
|
|
| def sha256(path: Path) -> str: |
| h = hashlib.sha256() |
| with path.open("rb") as handle: |
| for block in iter(lambda: handle.read(1 << 20), b""): |
| h.update(block) |
| return h.hexdigest() |
|
|
|
|
| def write_csv(path: Path, rows: list[dict]) -> None: |
| fieldnames = [] |
| for row in rows: |
| for key in row: |
| if key not in fieldnames: |
| fieldnames.append(key) |
| with path.open("w", newline="", encoding="utf-8") as handle: |
| writer = csv.DictWriter(handle, fieldnames=fieldnames) |
| writer.writeheader() |
| writer.writerows(rows) |
|
|
|
|
| def source_audit(root: Path) -> tuple[list[dict], dict]: |
| v1 = (root / "source_v1" / "main.tex").read_text(encoding="utf-8") |
| current = (root / "source_current" / "main.tex").read_text(encoding="utf-8") |
| anchors = [ |
| (1, "label{alg:SDRO-NGD}", "label{alg:SDRO-WFR}"), |
| (2, "label{prop:gradient_oracle_error_control}", "log{\\frac{L}{\\sqrt{\\lambda}\\epsilon}}"), |
| (3, "label{thm:outer_loop}", "S = O(1/\\epsilon_{\\text{opt}}^2)"), |
| (4, "label{thm:ula}", "\\epsilon_{\\text{opt}}^4"), |
| (5, "label{fig:robustness}", "WFR and WGF, in particular, consistently achieve"), |
| (6, "label{lem:sb-klform}", "Schr\\\"odinger half bridge problem"), |
| ] |
| rows = [] |
| for claim, anchor, companion in anchors: |
| rows.append({ |
| "claim": claim, |
| "v1_anchor": anchor, |
| "anchor_present": int(anchor in v1), |
| "v1_companion": companion, |
| "companion_present": int(companion in v1), |
| }) |
| algorithm_labels = [ |
| "alg:sampler", "alg:GF-DRO", "alg:SDRO-NGD", "alg:SDRO-WFR", |
| "alg:SDRO-SVG", "alg:SDRO_rgo", |
| ] |
| inventory = [{"algorithm": i + 1, "label": label, "present": int(f"label{{{label}}}" in v1)} for i, label in enumerate(algorithm_labels)] |
| return rows + inventory, { |
| "v1_pdf_hash_exact": sha256(root / "paper_v1.pdf") == V1_PDF_SHA256, |
| "v1_source_hash_exact": sha256(root / "source_v1.tar.gz") == V1_SOURCE_SHA256, |
| "current_pdf_hash_exact": sha256(root / "paper_current.pdf") == CURRENT_PDF_SHA256, |
| "current_source_hash_exact": sha256(root / "source_current.tar.gz") == CURRENT_SOURCE_SHA256, |
| "v1_claim_anchors": sum(int(r["anchor_present"] and r["companion_present"]) for r in rows), |
| "six_algorithm_labels": sum(r["present"] for r in inventory), |
| "rate_drift_detected": ( |
| "S = O(1/\\epsilon_{\\text{opt}}^2)" in v1 |
| and "S = O(\\frac1{\\epsilon_{\\text{opt}}^4})" in current |
| and "\\epsilon_{\\text{opt}}^6" in current |
| ), |
| } |
|
|
|
|
| def flow_time_audit() -> tuple[list[dict], dict]: |
| rows = [] |
| for lam in (0.25, 0.5, 1.0, 2.0, 4.0): |
| for lipschitz in (0.75, 1.5, 3.0): |
| for epsilon in (0.2, 0.1, 0.05, 0.02, 0.01): |
| initial_error = lipschitz / math.sqrt(lam) |
| threshold = math.log(initial_error / epsilon) / lam |
| at_threshold = initial_error * math.exp(-lam * threshold) |
| early = initial_error * math.exp(-lam * 0.8 * threshold) |
| rows.append({ |
| "lambda": lam, |
| "L": lipschitz, |
| "epsilon": epsilon, |
| "threshold_time": threshold, |
| "error_at_threshold": at_threshold, |
| "error_at_80pct_time": early, |
| "threshold_pass": int(at_threshold <= epsilon * (1 + 1e-12)), |
| "early_control_fails": int(early > epsilon), |
| }) |
| return rows, { |
| "cells": len(rows), |
| "max_threshold_ratio_error": max(abs(r["error_at_threshold"] / r["epsilon"] - 1.0) for r in rows), |
| "all_thresholds_pass": all(r["threshold_pass"] for r in rows), |
| "all_early_controls_fail": all(r["early_control_fails"] for r in rows), |
| } |
|
|
|
|
| def complexity_audit() -> tuple[list[dict], dict]: |
| rows = [] |
| for eps in (0.4, 0.25, 0.16, 0.1, 0.063, 0.04): |
| outer = eps ** -2 |
| for d in (2, 8, 32, 128): |
| for l_phi, l_u, l_f, lam_u in ((1.0, 1.0, 1.0, 1.0), (2.0, 1.5, 0.75, 0.5)): |
| prefactor = l_phi * l_u**2 * l_f**2 * d**2 / lam_u**3 |
| total = prefactor * eps ** -4 * math.log(1.0 / eps) |
| rows.append({ |
| "epsilon_opt": eps, "dimension": d, "L_phi": l_phi, |
| "L_U": l_u, "L_f": l_f, "lambda_U": lam_u, |
| "outer_iterations_proxy": outer, |
| "outer_normalized": outer * eps**2, |
| "total_complexity_proxy": total, |
| "total_normalized": total * eps**4 / (math.log(1.0 / eps) * prefactor), |
| }) |
| eps_grid = np.asarray(sorted({r["epsilon_opt"] for r in rows})) |
| outer_grid = eps_grid ** -2 |
| plain_total = eps_grid ** -4 |
| slope_outer = float(np.polyfit(np.log(eps_grid), np.log(outer_grid), 1)[0]) |
| slope_total = float(np.polyfit(np.log(eps_grid), np.log(plain_total), 1)[0]) |
| return rows, { |
| "cells": len(rows), |
| "outer_exponent": slope_outer, |
| "total_polynomial_exponent": slope_total, |
| "max_outer_identity_error": max(abs(r["outer_normalized"] - 1.0) for r in rows), |
| "max_total_identity_error": max(abs(r["total_normalized"] - 1.0) for r in rows), |
| } |
|
|
|
|
| def half_bridge_audit() -> tuple[list[dict], dict]: |
| rows = [] |
| xs = np.asarray([-1.0, 0.0, 1.0]) |
| px = np.asarray([0.2, 0.5, 0.3]) |
| ys = np.linspace(-2.0, 2.0, 17) |
| for tau in (0.2, 0.5, 1.0): |
| for epsilon in (0.25, 0.5, 1.0): |
| costs = (xs[:, None] - ys[None, :]) ** 2 |
| potential = 0.35 * ys**2 - 0.4 * ys |
| logits = -(2.0 * tau * potential[None, :] + costs) / epsilon |
| logits -= logits.max(axis=1, keepdims=True) |
| cond = np.exp(logits) |
| cond /= cond.sum(axis=1, keepdims=True) |
| coupling = px[:, None] * cond |
| marginal_y = coupling.sum(axis=0) |
| |
| |
| kkt = 2.0 * tau * potential[None, :] + costs + epsilon * np.log(cond) |
| kkt_residual = float(np.max(np.ptp(kkt, axis=1))) |
| mixture_error = float(np.max(np.abs(marginal_y - np.sum(px[:, None] * cond, axis=0)))) |
| rows.append({ |
| "tau": tau, "epsilon": epsilon, |
| "max_fixed_x_marginal_error": float(np.max(np.abs(coupling.sum(axis=1) - px))), |
| "max_conditional_normalization_error": float(np.max(np.abs(cond.sum(axis=1) - 1.0))), |
| "kkt_residual": kkt_residual, |
| "mixture_identity_error": mixture_error, |
| "marginal_y_mean": float(np.dot(marginal_y, ys)), |
| }) |
| return rows, { |
| "cells": len(rows), |
| "max_fixed_marginal_error": max(r["max_fixed_x_marginal_error"] for r in rows), |
| "max_conditional_normalization_error": max(r["max_conditional_normalization_error"] for r in rows), |
| "max_kkt_residual": max(r["kkt_residual"] for r in rows), |
| "max_mixture_identity_error": max(r["mixture_identity_error"] for r in rows), |
| } |
|
|
|
|
| def source_cifar_audit(root: Path) -> tuple[list[dict], dict]: |
| v1 = (root / "source_v1" / "main.tex").read_text(encoding="utf-8") |
| figures = ["rgo_lam=10_eps=0.2.pdf", "rgo_lam=10_eps=0.02.pdf", "rgo_lam=10_eps=0.002.pdf"] |
| rows = [{"epsilon": eps, "figure": name, "sha256": sha256(root / "source_v1" / name), "nonempty": int((root / "source_v1" / name).stat().st_size > 1000)} for eps, name in zip((0.2, 0.02, 0.002), figures)] |
| return rows, { |
| "three_primary_figures": len(rows) == 3 and all(r["nonempty"] for r in rows), |
| "cifar_setup_present": "features extracted from the real-world image dataset CIFAR-10" in v1 and "vary $\\Delta$ from 0 to 0.08" in v1, |
| "wfr_wgf_source_conclusion_present": "WFR and WGF, in particular, consistently achieve a high degree of robustness across all settings" in v1, |
| } |
|
|
|
|
| def make_figure(path: Path, flow: list[dict], complexity: list[dict], bridge: list[dict]) -> None: |
| fig, axes = plt.subplots(1, 3, figsize=(15, 4.6)) |
| subset = [r for r in flow if r["lambda"] == 1.0 and r["L"] == 1.5] |
| axes[0].loglog([r["epsilon"] for r in subset], [r["threshold_time"] for r in subset], "o-") |
| axes[0].set(title="Proposition 1 flow time", xlabel="epsilon", ylabel="threshold time") |
| eps = sorted({r["epsilon_opt"] for r in complexity}) |
| axes[1].loglog(eps, [e**-2 for e in eps], "o-", label="outer") |
| axes[1].loglog(eps, [e**-4 for e in eps], "s-", label="total polynomial") |
| axes[1].set(title="V1 complexity exponents", xlabel="epsilon_opt", ylabel="normalized work") |
| axes[1].legend(frameon=False) |
| axes[2].plot(range(len(bridge)), [r["marginal_y_mean"] for r in bridge], "o-") |
| axes[2].set(title="Half-bridge Gibbs mixtures", xlabel="parameter cell", ylabel="worst-case mean") |
| for ax in axes: |
| ax.grid(alpha=0.25) |
| fig.suptitle("Gradient-flow DRO: version-locked exact audits", weight="bold") |
| fig.tight_layout() |
| fig.savefig(path, dpi=180) |
| plt.close(fig) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--output-dir", type=Path, default=Path("outputs")) |
| args = parser.parse_args() |
| out = args.output_dir |
| out.mkdir(parents=True, exist_ok=True) |
| root = Path(__file__).resolve().parent |
|
|
| source_rows, source = source_audit(root) |
| flow_rows, flow = flow_time_audit() |
| complexity_rows, complexity = complexity_audit() |
| bridge_rows, bridge = half_bridge_audit() |
| cifar_rows, cifar = source_cifar_audit(root) |
| gates = { |
| "four_primary_hashes_exact": all(source[k] for k in ("v1_pdf_hash_exact", "v1_source_hash_exact", "current_pdf_hash_exact", "current_source_hash_exact")), |
| "six_exact_live_claims": len(CLAIMS) == 6, |
| "six_v1_claim_anchors": source["v1_claim_anchors"] == 6, |
| "six_concrete_algorithms": source["six_algorithm_labels"] == 6, |
| "material_rate_drift_detected": source["rate_drift_detected"], |
| "flow_time_grid_complete": flow["cells"] == 75, |
| "flow_time_threshold_exact": flow["all_thresholds_pass"] and flow["max_threshold_ratio_error"] < 1e-12, |
| "early_flow_destructive_control": flow["all_early_controls_fail"], |
| "complexity_grid_complete": complexity["cells"] == 48, |
| "outer_rate_exponent": abs(complexity["outer_exponent"] + 2.0) < 1e-12, |
| "total_rate_exponent": abs(complexity["total_polynomial_exponent"] + 4.0) < 1e-12, |
| "complexity_identities_exact": complexity["max_outer_identity_error"] < 1e-12 and complexity["max_total_identity_error"] < 1e-12, |
| "half_bridge_grid_complete": bridge["cells"] == 9, |
| "half_bridge_fixed_marginal": bridge["max_fixed_marginal_error"] < 1e-14, |
| "half_bridge_gibbs_kkt": bridge["max_kkt_residual"] < 1e-12, |
| "half_bridge_mixture_identity": bridge["max_mixture_identity_error"] < 1e-14, |
| "three_primary_cifar_figures": cifar["three_primary_figures"], |
| "cifar_setup_and_conclusion_pinned": cifar["cifar_setup_present"] and cifar["wfr_wgf_source_conclusion_present"], |
| } |
| write_csv(out / "source_and_algorithm_audit.csv", source_rows) |
| write_csv(out / "flow_time_thresholds.csv", flow_rows) |
| write_csv(out / "complexity_rates.csv", complexity_rows) |
| write_csv(out / "half_bridge_gibbs.csv", bridge_rows) |
| write_csv(out / "source_cifar_figures.csv", cifar_rows) |
| make_figure(out / "gradient_flow_dro_audit.png", flow_rows, complexity_rows, bridge_rows) |
| result = { |
| "paper_id": PAPER_ID, |
| "claims": CLAIMS, |
| "source": source, |
| "flow_time": flow, |
| "complexity": complexity, |
| "half_bridge": bridge, |
| "cifar_source": cifar, |
| "gates": gates, |
| "all_gates_pass": all(gates.values()), |
| "scope": { |
| "literal_claim_source": "arXiv v1", |
| "current_revision": "material rate-drift control", |
| "cifar_results": "pinned primary figures and source conclusion; not independently rerun", |
| "finite_audits": "exact mechanisms and rate identities; not replacements for universal proofs", |
| }, |
| } |
| (out / "results.json").write_text(json.dumps(result, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") |
| hashes = {p.name: sha256(p) for p in sorted(out.iterdir()) if p.is_file() and p.name != "SHA256SUMS.json"} |
| (out / "SHA256SUMS.json").write_text(json.dumps(hashes, indent=2) + "\n", encoding="utf-8") |
| print(json.dumps({"all_gates_pass": result["all_gates_pass"], "passed": sum(gates.values()), "total": len(gates), "output": str(out)}, indent=2)) |
| if not result["all_gates_pass"]: |
| raise SystemExit("one or more gates failed") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|