""" PC-SHO-DLM Settling Dynamics Analysis Empirically measures: 1. Convergence rate: first-order vs second-order 2. Energy decrease per microstep 3. Effect of precision conditioning 4. Warm-start vs cold-start efficiency 5. Token-level settling patterns """ import json import math import os import sys from pathlib import Path 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, count_parameters def measure_settling_convergence( model: PCSHODLM, x_0: torch.Tensor, max_k: int = 32, device: str = "cpu", ) -> dict: """Measure how energy decreases over settling steps. Returns energy traces for analysis. """ model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape # Sample a fixed timestep (mid-range) t = torch.full((B,), model.config.n_diffusion_steps // 2, device=device, dtype=torch.long) # Corrupt x_t, mask = model.schedule.corrupt(x_0, t, model.config.mask_token_id) # Embed and initialize h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) # Temporarily override settling steps original_k = model.config.n_settling_steps model.config.n_settling_steps = max_k model.config.settling_threshold = float("inf") # Disable adaptive for clean measurement # Run settling and record energy at each step 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 = [] grad_norms = [] token_uncertainties = [] for k in range(max_k): # Enable grads for energy computation for l in range(1, model.config.n_layers + 1): h[l] = h[l].detach().requires_grad_(True) grad_h, eps_up, eps_down, energy = model.compute_energy_gradient( h, h_init, x_0, mask, t ) energies.append(energy) # Record gradient norm total_grad_norm = sum(g.detach().norm().item() for g in grad_h) grad_norms.append(total_grad_norm) # Record per-token uncertainty uncertainty = model.compute_token_uncertainty(h).detach().cpu() token_uncertainties.append(uncertainty.mean().item()) # Compute aggregate precision for controller agg_precs = model.compute_aggregate_precision(h, t) # Update h_new = [h[0]] v_new = [] for l in range(model.config.n_layers): prec = agg_precs[l] eta = model.config.eta_base / (1.0 + model.config.c_eta * prec) gamma_raw = 2.0 * math.sqrt(model.config.mass * prec) gamma = max(model.config.gamma_min, min(model.config.gamma_max, gamma_raw)) v_l_new = (1.0 - gamma) * v[l] - eta * grad_h[l].detach() h_l_new = h[l + 1].detach() + v_l_new v_new.append(v_l_new) h_new.append(h_l_new) h = h_new v = v_new # Restore model.config.n_settling_steps = original_k return { "energies": energies, "grad_norms": grad_norms, "token_uncertainties": token_uncertainties, } def compare_first_vs_second_order( config_base: PCSHOConfig, x_0: torch.Tensor, max_k: int = 32, device: str = "cpu", ) -> dict: """Compare convergence of first-order vs second-order settling.""" results = {} # Second-order (full model) model_sho = PCSHODLM(config_base).to(device).eval() results["second_order"] = measure_settling_convergence( model_sho, x_0, max_k, device ) # First-order (gamma = 1 forces no momentum) config_fo = PCSHOConfig(**{ k: v for k, v in config_base.__dict__.items() }) config_fo.gamma_min = 1.0 config_fo.gamma_max = 1.0 model_fo = PCSHODLM(config_fo).to(device).eval() # Copy weights for fair comparison model_fo.load_state_dict(model_sho.state_dict(), strict=False) results["first_order"] = measure_settling_convergence( model_fo, x_0, max_k, device ) return results def measure_warm_start_advantage( model: PCSHODLM, x_0: torch.Tensor, n_settling: int = 8, device: str = "cpu", ) -> dict: """Compare warm-start vs cold-start across diffusion steps.""" model = model.to(device).eval() x_0 = x_0.to(device) B, S = x_0.shape warm_energies_per_step = [] cold_energies_per_step = [] # Run a few diffusion steps test_timesteps = [ model.config.n_diffusion_steps, model.config.n_diffusion_steps * 3 // 4, model.config.n_diffusion_steps // 2, model.config.n_diffusion_steps // 4, ] prev_h = None prev_v = None for t_val in test_timesteps: t = torch.full((B,), t_val, 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) # Warm start model.config.n_settling_steps = n_settling h_warm, v_warm, energies_warm, _, _ = model.settle( h_init, x_0, mask, t, prev_h=prev_h, prev_v=prev_v ) warm_energies_per_step.append(energies_warm) # Cold start h_cold, v_cold, energies_cold, _, _ = model.settle( h_init, x_0, mask, t, prev_h=None, prev_v=None ) cold_energies_per_step.append(energies_cold) # Update for next warm start prev_h = h_warm prev_v = v_warm return { "timesteps": test_timesteps, "warm_energies": warm_energies_per_step, "cold_energies": cold_energies_per_step, } def run_full_analysis( output_dir: str = "settling_analysis", device: str = "cpu", ): """Run complete settling dynamics analysis.""" 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, ) # Synthetic data x_0 = torch.randint(1, 257, (4, 128)) # 1. First-order vs second-order comparison print("1. Comparing first-order vs second-order settling...") fo_vs_sho = compare_first_vs_second_order(config, x_0, max_k=32, device=device) with open(output_path / "fo_vs_sho.json", "w") as f: json.dump(fo_vs_sho, f, indent=2) print(f" FO final energy: {fo_vs_sho['first_order']['energies'][-1]:.2f}") print(f" SHO final energy: {fo_vs_sho['second_order']['energies'][-1]:.2f}") # 2. Warm-start analysis print("2. Measuring warm-start advantage...") model = PCSHODLM(config) warm_analysis = measure_warm_start_advantage(model, x_0, device=device) with open(output_path / "warm_start.json", "w") as f: json.dump(warm_analysis, f, indent=2) for i, ts in enumerate(warm_analysis["timesteps"]): warm_final = warm_analysis["warm_energies"][i][-1] cold_final = warm_analysis["cold_energies"][i][-1] print(f" t={ts}: warm={warm_final:.2f}, cold={cold_final:.2f}") print(f"\nResults saved to {output_path}") if __name__ == "__main__": import argparse parser = argparse.ArgumentParser() parser.add_argument("--output", type=str, default="settling_analysis") parser.add_argument("--device", type=str, default="cpu") args = parser.parse_args() run_full_analysis(args.output, args.device)