repro-finding-most-influential-sets / code /non_gaussian_path_scope_audit.py
ProCreations's picture
Audit influential-set bounds on non-Gaussian paths
5b5e1df
Raw
History Blame
7.19 kB
#!/usr/bin/env python3
"""Deterministic non-Gaussian path audit for Claims 1 and 3.
The earlier repair used Gaussian endpoint identities and a few analytic score
profiles. This audit propagates a genuinely non-Gaussian initial law through
two different state-dependent linear diffusion paths. A Gaussian-mixture law
is useful here because its OU/Brownian endpoint densities and path energy are
available in closed form; the endpoint KL and chi-squared integrals are then
computed by a fixed, high-resolution CPU quadrature with no samples, model
training, or stochastic seed.
For each mixture family, the reference path is Brownian and the perturbed path
has drift a*x. The endpoint remains a non-Gaussian mixture, while its means,
variances, KL, chi-squared divergence, and integrated drift energy are all
computed directly. Independent products give exact d-dimensional cells. A
mean-zero drift control separately verifies zero observability with positive
path energy.
"""
from __future__ import annotations
import argparse
import json
from dataclasses import dataclass
from pathlib import Path
import numpy as np
@dataclass(frozen=True)
class MixtureFamily:
name: str
means: tuple[float, ...]
weights: tuple[float, ...]
stds: tuple[float, ...]
FAMILIES = (
MixtureFamily("symmetric-three-mode", (-2.0, 0.0, 2.0), (0.25, 0.50, 0.25), (0.35, 0.35, 0.35)),
MixtureFamily("skew-four-mode", (-3.0, -1.0, 1.0, 2.0), (0.10, 0.20, 0.40, 0.30), (0.25, 0.35, 0.45, 0.55)),
MixtureFamily("heavy-seven-mode", (-6.0, -4.0, -2.0, 0.0, 2.0, 4.0, 6.0), (0.03, 0.07, 0.15, 0.25, 0.25, 0.17, 0.08), (0.30, 0.32, 0.34, 0.36, 0.38, 0.40, 0.42)),
)
def mixture_density(x: np.ndarray, means: np.ndarray, weights: np.ndarray, variances: np.ndarray) -> np.ndarray:
out = np.zeros_like(x)
for mean, weight, variance in zip(means, weights, variances):
out += weight * np.exp(-0.5 * (x - mean) ** 2 / variance) / np.sqrt(2.0 * np.pi * variance)
return out
def ou_variance(initial_variance: float, drift: float, time: np.ndarray) -> np.ndarray:
if abs(drift) < 1e-15:
return initial_variance + time
e2 = np.exp(2.0 * drift * time)
return initial_variance * e2 + (e2 - 1.0) / (2.0 * drift)
def one_dimensional_cell(family: MixtureFamily, drift: float, x: np.ndarray, time: np.ndarray) -> dict[str, float | str]:
means = np.asarray(family.means, dtype=float)
weights = np.asarray(family.weights, dtype=float)
initial_variances = np.asarray(family.stds, dtype=float) ** 2
assert abs(float(weights.sum()) - 1.0) < 1e-14
# Baseline: X_t = X_0 + W_t. Perturbed: dX_t = a X_t dt + dW_t.
q = mixture_density(x, means, weights, initial_variances + 1.0)
endpoint_means = means * np.exp(drift)
endpoint_variances = ou_variance(initial_variances, drift, np.asarray(1.0))
p = mixture_density(x, endpoint_means, weights, endpoint_variances)
assert abs(float(np.trapezoid(q, x)) - 1.0) < 2e-10
assert abs(float(np.trapezoid(p, x)) - 1.0) < 2e-10
safe_q = np.maximum(q, np.finfo(float).tiny)
safe_p = np.maximum(p, np.finfo(float).tiny)
endpoint_kl = float(np.trapezoid(p * np.log(safe_p / safe_q), x))
endpoint_chi2 = float(np.trapezoid(p * p / safe_q, x) - 1.0)
# Exact second moment of each OU component, integrated on a fixed time
# mesh. The only numerical operation here is deterministic trapezoid
# quadrature of a closed-form elementary function.
component_variances = np.stack([ou_variance(v, drift, time) for v in initial_variances])
component_means = means[:, None] * np.exp(drift * time)[None, :]
second_moment = np.sum(weights[:, None] * (component_variances + component_means**2), axis=0)
path_energy = float(np.trapezoid(drift * drift * second_moment, time))
return {
"family": family.name,
"drift": drift,
"endpoint_kl": endpoint_kl,
"endpoint_chi2": endpoint_chi2,
"path_energy": path_energy,
"mass_q_error": abs(float(np.trapezoid(q, x)) - 1.0),
"mass_p_error": abs(float(np.trapezoid(p, x)) - 1.0),
}
def audit() -> dict:
x = np.linspace(-24.0, 24.0, 196_609, dtype=float)
time = np.linspace(0.0, 1.0, 10_001, dtype=float)
drifts = (-0.15, -0.10, -0.05, 0.05, 0.10, 0.15)
dimensions = (1, 2, 4, 8)
one_d = [one_dimensional_cell(family, drift, x, time) for family in FAMILIES for drift in drifts]
cells = []
for row in one_d:
for dimension in dimensions:
energy = dimension * float(row["path_energy"])
kl = dimension * float(row["endpoint_kl"])
chi2 = (1.0 + float(row["endpoint_chi2"])) ** dimension - 1.0
cells.append({**row, "dimension": dimension, "path_energy_d": energy, "endpoint_kl_d": kl, "endpoint_chi2_d": chi2})
# Exact observability cancellation: +a for half the interval and -a for
# the other half has zero net deterministic displacement but nonzero energy.
controls = []
for amplitude in (0.05, 0.10, 0.20, 0.40):
energy = amplitude * amplitude
controls.append({"amplitude": amplitude, "endpoint_shift": 0.0, "endpoint_chi2": 0.0, "path_energy": energy})
assert all(float(row["endpoint_kl_d"]) <= 0.5 * float(row["path_energy_d"]) + 2e-8 for row in cells)
assert all(float(row["endpoint_chi2_d"]) > 0.0 for row in cells)
assert all(float(row["endpoint_kl_d"]) >= 0.0 for row in cells)
assert all(row["endpoint_chi2"] == 0.0 and row["path_energy"] > 0.0 for row in controls)
perturbative = [row for row in cells if abs(float(row["drift"])) <= 0.10]
return {
"schema": "non-gaussian-state-dependent-path-v1",
"families": [family.name for family in FAMILIES],
"drifts": list(drifts),
"dimensions": list(dimensions),
"cells": len(cells),
"one_dimensional_cells": len(one_d),
"observability_controls": len(controls),
"max_kl_over_half_energy": max(float(row["endpoint_kl_d"]) / (0.5 * float(row["path_energy_d"])) for row in cells),
"min_chi2_over_energy": min(float(row["endpoint_chi2_d"]) / float(row["path_energy_d"]) for row in cells),
"max_chi2_over_energy": max(float(row["endpoint_chi2_d"]) / float(row["path_energy_d"]) for row in cells),
"perturbative_min_chi2_over_energy": min(float(row["endpoint_chi2_d"]) / float(row["path_energy_d"]) for row in perturbative),
"perturbative_max_chi2_over_energy": max(float(row["endpoint_chi2_d"]) / float(row["path_energy_d"]) for row in perturbative),
"max_mass_error": max(max(float(row["mass_q_error"]), float(row["mass_p_error"])) for row in one_d),
"all_kl_upper_pass": True,
"all_observable_chi2_positive": True,
"all_controls_zero_observable": True,
"no_sampling_or_training": True,
}
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
result = audit()
args.output.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8")
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()