File size: 7,618 Bytes
c2d8a57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
"""
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)