""" PC-SHO-DLM Assumption Diagnostics Empirical diagnostics to validate the theoretical assumptions: 1. Stationarity residual (envelope theorem support) 2. Local Hessian spectrum via HVP + Lanczos (curvature/kappa estimation) 3. Energy monotonicity statistics (discrete-time Lyapunov check) 4. Precision calibration (predicted precision vs empirical squared error) 5. Warm-start drift measurement (tracking theorem support) 6. Token freeze accuracy (soft-gate reliability) 7. Contraction estimation (solver contractivity) """ import json import math import os import sys from pathlib import Path from typing import Optional import torch import torch.nn.functional as F sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) from model import PCSHODLM, PCSHOConfig # ============================================================================= # 1. Stationarity Residual # ============================================================================= def measure_stationarity_residual( model: PCSHODLM, x_0: torch.Tensor, settling_steps_list: list[int] = [2, 4, 8, 16, 32], device: str = "cpu", ) -> dict: """Measure ||grad_h E_t(h^(K))|| as a function of K. Validates envelope theorem: low residual means post-settling local updates are approximately exact gradients of the reduced objective. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) results = {} for K in settling_steps_list: model.config.n_settling_steps = K model.config.settling_threshold = float("inf") h = [hi.detach() for hi in h_init] v = [torch.zeros_like(h_init[l + 1]) for l in range(model.config.n_layers)] for k in range(K): h, v, energy = model.settling_step(h, v, h_init, x_0, mask, t, None) # Compute final gradient norm L = model.config.n_layers for l in range(1, L + 1): h[l] = h[l].detach().requires_grad_(True) with torch.enable_grad(): _, _, _, energy = model.compute_energy_gradient(h, h_init, x_0, mask, t) grad_norms = [] for l in range(1, L + 1): if h[l].grad is not None: grad_norms.append(h[l].grad.detach().norm().item()) # Re-compute via autograd h_params = [h[l + 1] for l in range(L)] for l in range(1, L + 1): h[l] = h[l].detach().requires_grad_(True) with torch.enable_grad(): grad_h, _, _, energy_val = model.compute_energy_gradient(h, h_init, x_0, mask, t) total_residual = sum(g.detach().norm().item() for g in grad_h) results[K] = { "stationarity_residual": total_residual, "energy": energy_val, } print(f" K={K:3d}: residual={total_residual:.4f}, energy={energy_val:.2f}") return results # ============================================================================= # 2. Local Hessian Spectrum via HVP + Lanczos # ============================================================================= def hessian_vector_product( model: PCSHODLM, h: list[torch.Tensor], h_init: list[torch.Tensor], x_0: torch.Tensor, mask: torch.Tensor, t: torch.Tensor, v_dir: list[torch.Tensor], ) -> list[torch.Tensor]: """Compute Hessian-vector product d^2 E / dh^2 @ v using double autograd.""" L = model.config.n_layers with torch.enable_grad(): # First: compute gradient with create_graph=True for second derivatives for l in range(1, L + 1): h[l] = h[l].detach().requires_grad_(True) # Re-compute energy directly (need graph for second derivative) mu_up, mu_down = model.compute_predictions(h) eps_up = [h[l + 1] - mu_up[l] for l in range(L)] energy = 0.0 for l in range(L): energy += 0.5 * (eps_up[l] ** 2).sum() energy += model.config.anchor_rho * sum( ((h[l + 1] - h_init[l + 1]) ** 2).sum() for l in range(L) ) h_params = [h[l + 1] for l in range(L)] grads = torch.autograd.grad(energy, h_params, create_graph=True) # Dot product with direction dot = sum((g * vd).sum() for g, vd in zip(grads, v_dir)) # Second derivative hvp = torch.autograd.grad(dot, h_params, create_graph=False) return [hv.detach() for hv in hvp] def estimate_hessian_spectrum_lanczos( model: PCSHODLM, x_0: torch.Tensor, n_lanczos_steps: int = 20, device: str = "cpu", ) -> dict: """Estimate extremal eigenvalues of the energy Hessian via Lanczos iteration. Returns approximate lambda_min, lambda_max, and effective kappa. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape L = model.config.n_layers t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) # Settle to approximate fixed point model.config.n_settling_steps = 16 model.config.settling_threshold = float("inf") h_settled, _, _, _, _ = model.settle(h_init, x_0, mask, t) # Lanczos iteration # Flatten h into a single vector for eigenvalue computation def flatten_h(h_list): return torch.cat([h_list[l + 1].reshape(-1) for l in range(L)]) def unflatten_v(v_flat): parts = [] offset = 0 for l in range(L): size = h_settled[l + 1].numel() parts.append(v_flat[offset:offset + size].reshape_as(h_settled[l + 1])) offset += size return parts dim = sum(h_settled[l + 1].numel() for l in range(L)) # Initialize random vector q = torch.randn(dim, device=device) q = q / q.norm() alphas = [] betas = [0.0] Q = [q] for j in range(min(n_lanczos_steps, dim)): v_dir = unflatten_v(q) hvp = hessian_vector_product(model, list(h_settled), h_init, x_0, mask, t, v_dir) w = torch.cat([hv.reshape(-1) for hv in hvp]) alpha = w.dot(q).item() alphas.append(alpha) if j > 0: w = w - betas[-1] * Q[-2] w = w - alpha * q beta = w.norm().item() if beta < 1e-10: break betas.append(beta) q = w / beta Q.append(q) # Build tridiagonal matrix and compute eigenvalues n = len(alphas) T_mat = torch.zeros(n, n, device=device) for i in range(n): T_mat[i, i] = alphas[i] for i in range(n - 1): T_mat[i, i + 1] = betas[i + 2] if i + 2 < len(betas) else 0 T_mat[i + 1, i] = betas[i + 2] if i + 2 < len(betas) else 0 eigenvalues = torch.linalg.eigvalsh(T_mat).cpu().numpy() lambda_min = float(max(eigenvalues.min(), 1e-8)) lambda_max = float(eigenvalues.max()) kappa = lambda_max / lambda_min if lambda_min > 0 else float("inf") print(f" Hessian spectrum: lambda_min={lambda_min:.4f}, lambda_max={lambda_max:.4f}, kappa={kappa:.2f}") return { "eigenvalues": eigenvalues.tolist(), "lambda_min": lambda_min, "lambda_max": lambda_max, "effective_kappa": kappa, } # ============================================================================= # 3. Energy Monotonicity Statistics # ============================================================================= def measure_energy_monotonicity( model: PCSHODLM, x_0: torch.Tensor, K: int = 32, device: str = "cpu", ) -> dict: """Track energy at each microstep and report monotonicity violations. Tests whether discrete-time settling maintains the Lyapunov guarantee. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) model.config.n_settling_steps = K model.config.settling_threshold = float("inf") h = [hi.detach() for hi in h_init] v = [torch.zeros_like(h_init[l + 1]) for l in range(model.config.n_layers)] energies = [] kinetic_energies = [] for k in range(K): h, v, energy = model.settling_step(h, v, h_init, x_0, mask, t, None) energies.append(energy) # Kinetic energy: 0.5 * sum ||v_l||^2 ke = 0.5 * sum(vl.detach().norm().item() ** 2 for vl in v) kinetic_energies.append(ke) # Compute total energy (potential + kinetic) total_energies = [e + k for e, k in zip(energies, kinetic_energies)] # Count monotonicity violations violations = 0 for i in range(1, len(total_energies)): if total_energies[i] > total_energies[i - 1] + 1e-6: violations += 1 print(f" Energy monotonicity: {violations}/{K-1} violations") print(f" Energy range: {energies[0]:.2f} -> {energies[-1]:.2f}") return { "potential_energies": energies, "kinetic_energies": kinetic_energies, "total_energies": total_energies, "n_violations": violations, "violation_rate": violations / max(1, K - 1), } # ============================================================================= # 4. Precision Calibration # ============================================================================= def measure_precision_calibration( model: PCSHODLM, x_0: torch.Tensor, n_samples: int = 10, device: str = "cpu", ) -> dict: """Compare predicted precision to empirical squared prediction errors. If precision heads are well-calibrated, predicted precision should correlate inversely with actual error magnitude. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape predicted_precs = [] empirical_sq_errors = [] for _ in range(n_samples): t = torch.randint(1, model.config.n_diffusion_steps + 1, (B,), device=device) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) model.config.n_settling_steps = 8 h_settled, _, _, _, _ = model.settle(h_init, x_0, mask, t) # Get predictions and precisions at each layer for l in range(model.config.n_layers): with torch.no_grad(): mu_up = model.forward_blocks[l](h_settled[l]) tok_p, ch_p = model.precision_up[l](h_settled[l + 1], t) # Empirical squared error eps = (h_settled[l + 1] - mu_up) sq_err = (eps ** 2).mean(dim=-1) # (B, S) # Predicted precision (token-level) pred_p = tok_p.squeeze(-1) # (B, S) predicted_precs.append(pred_p.cpu()) empirical_sq_errors.append(sq_err.cpu()) # Compute correlation all_prec = torch.cat([p.flatten() for p in predicted_precs]) all_err = torch.cat([e.flatten() for e in empirical_sq_errors]) # Pearson correlation between precision and 1/error inv_err = 1.0 / (all_err + 1e-8) correlation = torch.corrcoef(torch.stack([all_prec, inv_err]))[0, 1].item() # Spearman rank correlation (approximate via sorting) n = min(len(all_prec), 10000) idx = torch.randperm(len(all_prec))[:n] prec_ranks = all_prec[idx].argsort().argsort().float() err_ranks = all_err[idx].argsort().argsort().float() # Higher precision should correlate with lower error (higher error rank) rank_corr = torch.corrcoef(torch.stack([prec_ranks, err_ranks]))[0, 1].item() print(f" Precision calibration:") print(f" Pearson(precision, 1/error): {correlation:.4f}") print(f" Rank correlation(precision, error_rank): {rank_corr:.4f}") return { "pearson_prec_inv_error": correlation, "rank_correlation": rank_corr, "mean_precision": all_prec.mean().item(), "mean_sq_error": all_err.mean().item(), } # ============================================================================= # 5. Contraction Estimation # ============================================================================= def estimate_contraction_factor( model: PCSHODLM, x_0: torch.Tensor, perturbation_scale: float = 0.01, K: int = 8, device: str = "cpu", ) -> dict: """Estimate solver contraction factor by perturbation experiments. Perturb h^k slightly and measure how fast the perturbation decays. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape L = model.config.n_layers t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) model.config.n_settling_steps = K model.config.settling_threshold = float("inf") # Run unperturbed h_clean = [hi.detach() for hi in h_init] v_clean = [torch.zeros_like(h_init[l + 1]) for l in range(L)] # Run perturbed h_pert = [hi.detach().clone() for hi in h_init] for l in range(1, L + 1): h_pert[l] = h_pert[l] + perturbation_scale * torch.randn_like(h_pert[l]) v_pert = [torch.zeros_like(h_init[l + 1]) for l in range(L)] distances = [] initial_dist = sum( (h_pert[l + 1] - h_clean[l + 1]).norm().item() for l in range(L) ) distances.append(initial_dist) for k in range(K): h_clean, v_clean, _ = model.settling_step( h_clean, v_clean, h_init, x_0, mask, t, None ) h_pert, v_pert, _ = model.settling_step( h_pert, v_pert, h_init, x_0, mask, t, None ) dist = sum( (h_pert[l + 1] - h_clean[l + 1]).detach().norm().item() for l in range(L) ) distances.append(dist) # Estimate contraction per step contraction_factors = [] for i in range(1, len(distances)): if distances[i - 1] > 1e-10: contraction_factors.append(distances[i] / distances[i - 1]) avg_contraction = sum(contraction_factors) / len(contraction_factors) if contraction_factors else 1.0 print(f" Contraction estimation:") print(f" Initial perturbation distance: {initial_dist:.6f}") print(f" Final distance: {distances[-1]:.6f}") print(f" Average contraction factor: {avg_contraction:.4f}") print(f" Contractive: {avg_contraction < 1.0}") return { "distances": distances, "contraction_factors": contraction_factors, "avg_contraction": avg_contraction, "is_contractive": avg_contraction < 1.0, } # ============================================================================= # 6. Token Freeze Accuracy # ============================================================================= def measure_freeze_accuracy( model: PCSHODLM, x_0: torch.Tensor, device: str = "cpu", ) -> dict: """Measure fraction of soft-frozen tokens that later change. Runs settling with and without soft-freezing, compares final tokens. Tests whether early confident predictions are reliable. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) # Run without adaptive settling (full compute) model.config.settling_threshold = float("inf") h_full, _, _, _, _ = model.settle(h_init, x_0, mask, t) logits_full = model.readout(model.readout_norm(h_full[-1])) tokens_full = logits_full.argmax(dim=-1) # Run with adaptive settling (soft-freeze) model.config.settling_threshold = 0.5 h_adaptive, _, _, _, _ = model.settle(h_init, x_0, mask, t) logits_adaptive = model.readout(model.readout_norm(h_adaptive[-1])) tokens_adaptive = logits_adaptive.argmax(dim=-1) # Check: at masked positions, how often do frozen-early tokens match full-compute tokens masked_match = (tokens_full[mask] == tokens_adaptive[mask]).float() match_rate = masked_match.mean().item() if masked_match.numel() > 0 else 1.0 # Compute mid-settling predictions to identify "early frozen" tokens model.config.n_settling_steps = model.config.n_settling_steps // 2 model.config.settling_threshold = float("inf") h_mid, _, _, _, _ = model.settle(h_init, x_0, mask, t) logits_mid = model.readout(model.readout_norm(h_mid[-1])) tokens_mid = logits_mid.argmax(dim=-1) # "Early confident" = tokens that don't change from mid to full early_stable = (tokens_mid[mask] == tokens_full[mask]).float() early_stable_rate = early_stable.mean().item() if early_stable.numel() > 0 else 1.0 print(f" Token freeze accuracy:") print(f" Adaptive vs full match rate: {match_rate:.4f}") print(f" Early-stable token rate: {early_stable_rate:.4f}") return { "adaptive_vs_full_match": match_rate, "early_stable_rate": early_stable_rate, } # ============================================================================= # Run All Diagnostics # ============================================================================= def run_all_diagnostics( output_dir: str = "diagnostics_output", device: str = "cpu", ): """Run complete assumption diagnostic suite.""" output_path = Path(output_dir) output_path.mkdir(parents=True, exist_ok=True) config = PCSHOConfig( vocab_size=257, max_seq_len=128, d_model=256, n_heads=4, n_layers=4, d_ff=512, n_diffusion_steps=100, n_settling_steps=8, mask_token_id=0, ) model = PCSHODLM(config) x_0 = torch.randint(1, 257, (4, 128)) all_results = {} print("=" * 60) print("PC-SHO-DLM Assumption Diagnostics") print("=" * 60) print("\n1. Stationarity Residual (Envelope Theorem)") all_results["stationarity"] = measure_stationarity_residual( model, x_0, [2, 4, 8, 16], device ) print("\n2. Local Hessian Spectrum (HVP + Lanczos)") all_results["hessian"] = estimate_hessian_spectrum_lanczos( model, x_0, n_lanczos_steps=10, device=device ) print("\n3. Energy Monotonicity") all_results["monotonicity"] = measure_energy_monotonicity( model, x_0, K=16, device=device ) print("\n4. Precision Calibration") all_results["precision"] = measure_precision_calibration( model, x_0, n_samples=5, device=device ) print("\n5. Contraction Estimation") all_results["contraction"] = estimate_contraction_factor( model, x_0, device=device ) print("\n6. Token Freeze Accuracy") all_results["freeze_accuracy"] = measure_freeze_accuracy( model, x_0, device=device ) # Save # Convert numpy arrays to lists for JSON def convert(obj): import numpy as np if isinstance(obj, np.ndarray): return obj.tolist() if isinstance(obj, dict): return {k: convert(v) for k, v in obj.items()} if isinstance(obj, list): return [convert(i) for i in obj] return obj with open(output_path / "diagnostics.json", "w") as f: json.dump(convert(all_results), f, indent=2) print(f"\nResults saved to {output_path}/diagnostics.json") return all_results if __name__ == "__main__": import argparse parser = argparse.ArgumentParser() parser.add_argument("--output", type=str, default="diagnostics_output") parser.add_argument("--device", type=str, default="cpu") args = parser.parse_args() run_all_diagnostics(args.output, args.device)