Download experiments/diagnostics.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 20.3 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/experiments/diagnostics.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/experiments/diagnostics.py
-
curl -L -o diagnostics.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/experiments/diagnostics.py
20.3 kB
| """ | |
| 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) | |