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