pc-sho-dlm-code / experiments /diagnostics.py
Zae
PC-SHO-DLM: full architecture with MSA integration
c2d8a57
Raw History Blame Contribute Delete
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)