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