import json import time import numpy as np from pruning_mvp import make_net, prune_report, output_interval, forward def run(seed=123, n_cases=32): rng = np.random.default_rng(seed) rows = [] for case in range(n_cases): W, b = make_net(seed=1000 + case) center = rng.uniform(-0.5, 0.5, 2) eps = 0.30 xlo, xhi = center - eps, center + eps hiddenW, hiddenb = W[:-1], b[:-1] base_vars = 2 * sum(len(v) for v in hiddenb) t0 = time.perf_counter() base_lo, base_hi = output_interval(W, b, xlo, xhi) base_time = time.perf_counter() - t0 row = {'case': case, 'baseline_variables': base_vars, 'baseline_interval_seconds': base_time} for tau, name in [(0.0, 'exact_sign'), (1e-6, 'tau_1e-6'), (0.1, 'tau_0.1'), (0.3, 'tau_0.3'), (1.0, 'tau_1.0')]: t0 = time.perf_counter() _, reports = prune_report(hiddenW, hiddenb, xlo, xhi, tau) idea_time = time.perf_counter() - t0 retained = sum(r.retained for r in reports) fixed = sum(r.active + r.inactive for r in reports) row[name] = {'variables': 2 * retained, 'fixed_sign': fixed, 'contribution_pruned': sum(r.pruned_contribution for r in reports), 'seconds': idea_time, # Exact sign mode must preserve the interval exactly. 'certificate_equal': bool(tau == 0.0 and np.array_equal( base_lo, output_interval(W, b, xlo, xhi)[0]) and np.array_equal(base_hi, output_interval(W, b, xlo, xhi)[1]))} rows.append(row) names = ['exact_sign', 'tau_1e-6', 'tau_0.1', 'tau_0.3', 'tau_1.0'] summary = {'cases': n_cases, 'baseline_variables_mean': float(np.mean([r['baseline_variables'] for r in rows])), 'baseline_seconds_mean': float(np.mean([r['baseline_interval_seconds'] for r in rows]))} for name in names: vals = [r[name] for r in rows] summary[name] = { 'variables_mean': float(np.mean([v['variables'] for v in vals])), 'reduction_fraction_mean': float(np.mean([(r['baseline_variables'] - v['variables']) / r['baseline_variables'] for r, v in zip(rows, vals)])), 'fixed_sign_mean': float(np.mean([v['fixed_sign'] for v in vals])), 'contribution_pruned_mean': float(np.mean([v['contribution_pruned'] for v in vals])), 'certificate_equal_all': bool(all(v['certificate_equal'] for v in vals)) if name == 'exact_sign' else None, } # Independent point checks on the first case: output intervals remain sound. W, b = make_net(seed=1000) center = rows[0]['case'] * 0.0 + np.array([0.0, 0.0]) xlo, xhi = center - .3, center + .3 lo, hi = output_interval(W, b, xlo, xhi) points = rng.uniform(xlo, xhi, (5000, 2)) violations = sum(np.any((forward(W, b, p) < lo - 1e-10) | (forward(W, b, p) > hi + 1e-10)) for p in points) result = {'summary': summary, 'independent_soundness_points': 5000, 'independent_soundness_violations': int(violations)} with open('results.json', 'w') as f: json.dump(result, f, indent=2) print(json.dumps(result, indent=2)) if __name__ == '__main__': run()