Download src/model.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 74.7 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/model.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/src/model.py
-
curl -L -o model.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/model.py
74.7 kB
| """ | |
| PC-SHO-DLM: Predictive-Coding Diffusion Language Models | |
| with Precision-Conditioned Second-Order Settling | |
| Reference implementation of the core architecture. | |
| """ | |
| import math | |
| from dataclasses import dataclass | |
| from typing import Optional, Tuple | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # ============================================================================= | |
| # Configuration | |
| # ============================================================================= | |
| class PCSHOConfig: | |
| """Configuration for PC-SHO-DLM.""" | |
| vocab_size: int = 32000 | |
| max_seq_len: int = 1024 | |
| d_model: int = 512 | |
| n_heads: int = 8 | |
| n_layers: int = 12 | |
| d_ff: int = 2048 | |
| dropout: float = 0.1 | |
| # Diffusion | |
| n_diffusion_steps: int = 1000 | |
| mask_token_id: int = 0 # [MASK] token id | |
| # Predictive coding settling | |
| n_settling_steps: int = 8 | |
| anchor_rho: float = 0.01 | |
| denoising_lambda: float = 5.0 # For honest entropy-based settling (~30-50% of energy) | |
| # Second-order dynamics | |
| eta_base: float = 0.1 | |
| c_eta: float = 1.0 | |
| gamma_min: float = 0.1 | |
| gamma_max: float = 0.99 | |
| mass: float = 1.0 | |
| # State continuation | |
| velocity_decay: float = 0.9 | |
| # Adaptive settling (soft-gate) | |
| settling_threshold: float = 0.5 | |
| settling_temperature: float = 0.1 # Sharpness of soft gate sigmoid | |
| entropy_coeff: float = 1.0 | |
| error_coeff: float = 0.1 | |
| # Feedback predictor | |
| feedback_rank: int = 64 # Low-rank feedback, 0 = tied weights | |
| # Online learning during inference | |
| online_learn_lr: float = 1e-5 | |
| online_learn_grad_clip: float = 0.1 | |
| online_learn_min_energy_ratio: float = 0.95 # only update if energy decreased by 5%+ | |
| # Unified post-settle learning | |
| unified_param_lr: float = 1e-3 | |
| unified_precision_lr_scale: float = 0.1 | |
| unified_task_lr_scale: float = 1.0 | |
| # Temporal hierarchy (per-layer scaling of SHO dynamics) | |
| mass_scale: float = 0.5 # mass increases with depth: m_l = mass * (1 + scale * l/L) | |
| gamma_scale: float = 0.3 # damping decreases with depth | |
| eta_scale: float = 0.3 # step size decreases with depth | |
| # Homeostatic plasticity (adaptive anchoring per layer) | |
| homeostatic_target: float = 0.0 # 0 = auto-detect from first batch | |
| homeostatic_rate: float = 0.0001 # very slow adaptation | |
| homeostatic_rho_min: float = 0.005 | |
| homeostatic_rho_max: float = 0.02 # never exceed 2x the base anchor_rho | |
| # Lateral inhibition — DISABLED (breaks settling convergence via random sampling) | |
| lateral_inhibition: float = 0.0 | |
| settling_budget_fraction: float = 0.5 | |
| # ============================================================================= | |
| # Precision Head | |
| # ============================================================================= | |
| class PrecisionHead(nn.Module): | |
| """Learned tokenwise/channelwise precision for a single layer. | |
| Predicts diagonal precision from hidden states and timestep. | |
| Uses factored (token x channel) parameterization. | |
| """ | |
| def __init__(self, d_model: int, n_diffusion_steps: int, eps: float = 1e-6): | |
| super().__init__() | |
| self.eps = eps | |
| # Token-level precision | |
| self.tok_proj = nn.Sequential( | |
| nn.Linear(d_model, d_model // 4), | |
| nn.GELU(), | |
| nn.Linear(d_model // 4, 1), | |
| ) | |
| # Channel-level precision (shared across tokens) | |
| self.ch_proj = nn.Sequential( | |
| nn.Linear(d_model, d_model // 4), | |
| nn.GELU(), | |
| nn.Linear(d_model // 4, d_model), | |
| ) | |
| # Timestep modulation | |
| self.time_embed = nn.Embedding(n_diffusion_steps + 1, d_model) | |
| def forward( | |
| self, h: torch.Tensor, t: torch.Tensor | |
| ) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """ | |
| Args: | |
| h: (batch, seq_len, d_model) | |
| t: (batch,) integer timesteps | |
| Returns: | |
| tok_precision: (batch, seq_len, 1) | |
| ch_precision: (batch, 1, d_model) | |
| """ | |
| t_emb = self.time_embed(t).unsqueeze(1) # (B, 1, D) | |
| h_mod = h + t_emb | |
| # Aggregate over sequence for channel precision | |
| h_mean = h_mod.mean(dim=1, keepdim=True) # (B, 1, D) | |
| tok_precision = F.softplus(self.tok_proj(h_mod)).clamp(max=10.0) + self.eps # (B, S, 1) | |
| ch_precision = F.softplus(self.ch_proj(h_mean)).clamp(max=10.0) + self.eps # (B, 1, D) | |
| return tok_precision, ch_precision | |
| # ============================================================================= | |
| # Bidirectional Transformer Block (Forward / Bottom-Up) | |
| # ============================================================================= | |
| class BidirectionalTransformerBlock(nn.Module): | |
| """Standard bidirectional transformer block (bottom-up pathway).""" | |
| def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1): | |
| super().__init__() | |
| self.self_attn = nn.MultiheadAttention( | |
| d_model, n_heads, dropout=dropout, batch_first=True | |
| ) | |
| self.ff = nn.Sequential( | |
| nn.Linear(d_model, d_ff), | |
| nn.GELU(), | |
| nn.Linear(d_ff, d_model), | |
| nn.Dropout(dropout), | |
| ) | |
| self.norm1 = nn.LayerNorm(d_model) | |
| self.norm2 = nn.LayerNorm(d_model) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| # Self-attention (bidirectional, no causal mask) | |
| residual = x | |
| x = self.norm1(x) | |
| attn_out, _ = self.self_attn(x, x, x) | |
| x = residual + self.dropout(attn_out) | |
| # Feedforward | |
| residual = x | |
| x = self.norm2(x) | |
| x = residual + self.ff(x) | |
| return x | |
| # ============================================================================= | |
| # Feedback Predictor (Top-Down Pathway) | |
| # ============================================================================= | |
| class FeedbackPredictor(nn.Module): | |
| """Lightweight top-down predictor g_l: predicts h_l from h_{l+1}. | |
| Uses low-rank projection by default for parameter efficiency. | |
| """ | |
| def __init__(self, d_model: int, rank: int = 64): | |
| super().__init__() | |
| if rank > 0 and rank < d_model: | |
| self.proj = nn.Sequential( | |
| nn.LayerNorm(d_model), | |
| nn.Linear(d_model, rank, bias=False), | |
| nn.GELU(), | |
| nn.Linear(rank, d_model, bias=False), | |
| ) | |
| else: | |
| # Full-rank feedback | |
| self.proj = nn.Sequential( | |
| nn.LayerNorm(d_model), | |
| nn.Linear(d_model, d_model), | |
| nn.GELU(), | |
| nn.Linear(d_model, d_model), | |
| ) | |
| def forward(self, h_above: torch.Tensor) -> torch.Tensor: | |
| return self.proj(h_above) | |
| # ============================================================================= | |
| # Lightweight Settling Block (Linear Attention for Inner Loop) | |
| # ============================================================================= | |
| class LinearAttention(nn.Module): | |
| """O(n*d^2) linear attention with recurrent SSM dual form. | |
| Supports two modes: | |
| - Batched: O(n*d^2) over full sequence (training) | |
| - Recurrent: O(d^2) per token with fixed-size state (inference) | |
| The recurrent state S = K^T @ V is the Mamba-2 duality: | |
| same math as linear attention but O(1) memory per new token. | |
| This eliminates KV-cache entirely — 393KB total for unlimited context. | |
| Per-head precision implements attention as active inference. | |
| """ | |
| def __init__(self, d_model: int, n_heads: int, dropout: float = 0.0): | |
| super().__init__() | |
| self.n_heads = n_heads | |
| self.d_head = d_model // n_heads | |
| self.q_proj = nn.Linear(d_model, d_model) | |
| self.k_proj = nn.Linear(d_model, d_model) | |
| self.v_proj = nn.Linear(d_model, d_model) | |
| self.o_proj = nn.Linear(d_model, d_model) | |
| # Per-head precision (attention as active inference) | |
| self.head_precision = nn.Linear(d_model, n_heads) | |
| # Recurrent state: S[h] = (d_head, d_head), z[h] = (d_head,) per head | |
| self._recurrent_state = None # (B, H, d, d) | |
| self._recurrent_norm = None # (B, H, d) | |
| def _elu_feature(self, x): | |
| return F.elu(x, alpha=1.0) + 1.0 | |
| def reset_recurrent_state(self): | |
| """Reset for new sequence.""" | |
| self._recurrent_state = None | |
| self._recurrent_norm = None | |
| def forward(self, x: torch.Tensor, recurrent: bool = False) -> torch.Tensor: | |
| """Forward pass. | |
| Args: | |
| x: (B, S, D) input | |
| recurrent: if True, use O(1) recurrent mode (incremental state update) | |
| """ | |
| B, S, D = x.shape | |
| H, d = self.n_heads, self.d_head | |
| Q = self._elu_feature(self.q_proj(x).view(B, S, H, d).transpose(1, 2)) | |
| K = self._elu_feature(self.k_proj(x).view(B, S, H, d).transpose(1, 2)) | |
| V = self.v_proj(x).view(B, S, H, d).transpose(1, 2) | |
| # Per-head precision: active inference gating | |
| pi = F.softplus(self.head_precision(x.mean(dim=1))) # (B, H) | |
| Q = Q * pi.unsqueeze(-1).unsqueeze(-1) | |
| if recurrent: | |
| # === RECURRENT MODE: O(d^2) per token, replaces KV-cache === | |
| # State: S = sum_t k_t @ v_t^T, shape (B, H, d, d) | |
| # Output: o_t = q_t @ S_t / (q_t @ z_t) | |
| device = x.device | |
| if self._recurrent_state is None: | |
| self._recurrent_state = torch.zeros(B, H, d, d, device=device) | |
| self._recurrent_norm = torch.zeros(B, H, d, device=device) | |
| outputs = [] | |
| for t in range(S): | |
| k_t = K[:, :, t, :] # (B, H, d) | |
| v_t = V[:, :, t, :] | |
| q_t = Q[:, :, t, :] | |
| # Update recurrent state: S += k @ v^T | |
| self._recurrent_state = self._recurrent_state + torch.einsum("bhd,bhe->bhde", k_t, v_t) | |
| self._recurrent_norm = self._recurrent_norm + k_t | |
| # Query: o = q @ S / (q @ z) | |
| o_t = torch.einsum("bhd,bhde->bhe", q_t, self._recurrent_state) | |
| z_t = torch.einsum("bhd,bhd->bh", q_t, self._recurrent_norm).unsqueeze(-1) + 1e-6 | |
| outputs.append(o_t / z_t) | |
| out = torch.stack(outputs, dim=2).transpose(1, 2).reshape(B, S, D) | |
| else: | |
| # === BATCHED MODE: O(n * d^2) over full sequence === | |
| KV = torch.einsum("bhsd,bhse->bhde", K, V) | |
| QKV = torch.einsum("bhsd,bhde->bhse", Q, KV) | |
| normalizer = torch.einsum("bhsd,bhd->bhs", Q, K.sum(dim=2)).unsqueeze(-1) + 1e-6 | |
| out = (QKV / normalizer).transpose(1, 2).reshape(B, S, D) | |
| # Cache the final recurrent state for potential continuation | |
| self._recurrent_state = KV.detach() | |
| self._recurrent_norm = K.sum(dim=2).detach() | |
| return self.o_proj(out) | |
| class LightweightSettlingBlock(nn.Module): | |
| """Settling block with linear attention — used in the K inner iterations. | |
| Same structure as BidirectionalTransformerBlock but O(n) instead of O(n^2). | |
| """ | |
| def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1): | |
| super().__init__() | |
| self.attn = LinearAttention(d_model, n_heads, dropout) | |
| self.ff = nn.Sequential( | |
| nn.Linear(d_model, d_ff), | |
| nn.GELU(), | |
| nn.Linear(d_ff, d_model), | |
| nn.Dropout(dropout), | |
| ) | |
| self.norm1 = nn.LayerNorm(d_model) | |
| self.norm2 = nn.LayerNorm(d_model) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| residual = x | |
| x = self.norm1(x) | |
| x = residual + self.dropout(self.attn(x)) | |
| residual = x | |
| x = self.norm2(x) | |
| x = residual + self.ff(x) | |
| return x | |
| # ============================================================================= | |
| # Masking Diffusion Schedule | |
| # ============================================================================= | |
| class MaskDiffusionSchedule(nn.Module): | |
| """Absorbing-mask corruption schedule for discrete diffusion.""" | |
| def __init__(self, n_steps: int): | |
| super().__init__() | |
| self.n_steps = n_steps | |
| # Cosine masking schedule | |
| alphas = torch.cos( | |
| (torch.arange(n_steps + 1) / n_steps) * (math.pi / 2) | |
| ) ** 2 | |
| # alpha_bar: probability of NOT being masked | |
| # At t=0: alpha_bar ~ 1 (nothing masked) | |
| # At t=T: alpha_bar ~ 0 (everything masked) | |
| self.register_buffer("alpha_bar", alphas) | |
| def mask_rate(self, t: torch.Tensor) -> torch.Tensor: | |
| """Probability of a token being masked at timestep t.""" | |
| return 1.0 - self.alpha_bar[t] | |
| def corrupt( | |
| self, | |
| x_0: torch.Tensor, | |
| t: torch.Tensor, | |
| mask_token_id: int, | |
| ) -> Tuple[torch.Tensor, torch.Tensor]: | |
| """Apply forward corruption: mask tokens with probability alpha_bar_t. | |
| Args: | |
| x_0: (batch, seq_len) clean token ids | |
| t: (batch,) timesteps | |
| mask_token_id: id of [MASK] token | |
| Returns: | |
| x_t: corrupted sequence | |
| mask: boolean mask of corrupted positions | |
| """ | |
| mask_prob = self.mask_rate(t).unsqueeze(-1) # (B, 1) | |
| mask = torch.rand_like(x_0.float()) < mask_prob # (B, S) | |
| empty_rows = ~mask.any(dim=1) | |
| if empty_rows.any(): | |
| forced = torch.randint(0, x_0.shape[1], (int(empty_rows.sum().item()),), device=x_0.device) | |
| mask[empty_rows] = False | |
| mask[empty_rows, forced] = True | |
| x_t = x_0.clone() | |
| x_t[mask] = mask_token_id | |
| return x_t, mask | |
| # ============================================================================= | |
| # PC-SHO-DLM Core Model | |
| # ============================================================================= | |
| class PCSHODLM(nn.Module): | |
| """PC-SHO-DLM: Predictive-Coding Diffusion Language Model | |
| with Precision-Conditioned Second-Order Settling. | |
| """ | |
| def __init__(self, config: PCSHOConfig): | |
| super().__init__() | |
| self.config = config | |
| self._inference_updater = None # lazy init for online learning | |
| self._unified_batch_updater = None # lazy init for post-settle unified training | |
| self._honest_settling = True # no peeking at x_0 during settling | |
| # Embeddings | |
| self.token_embed = nn.Embedding(config.vocab_size, config.d_model) | |
| self.pos_embed = nn.Embedding(config.max_seq_len, config.d_model) | |
| self.time_embed = nn.Embedding(config.n_diffusion_steps + 1, config.d_model) | |
| # Bottom-up transformer blocks | |
| self.forward_blocks = nn.ModuleList([ | |
| BidirectionalTransformerBlock( | |
| config.d_model, config.n_heads, config.d_ff, config.dropout | |
| ) | |
| for _ in range(config.n_layers) | |
| ]) | |
| # Top-down feedback predictors | |
| self.feedback_blocks = nn.ModuleList([ | |
| FeedbackPredictor(config.d_model, config.feedback_rank) | |
| for _ in range(config.n_layers - 1) # No feedback from above top layer | |
| ]) | |
| # Precision heads (one per layer for both up and down) | |
| self.precision_up = nn.ModuleList([ | |
| PrecisionHead(config.d_model, config.n_diffusion_steps) | |
| for _ in range(config.n_layers) | |
| ]) | |
| self.precision_down = nn.ModuleList([ | |
| PrecisionHead(config.d_model, config.n_diffusion_steps) | |
| for _ in range(config.n_layers - 1) | |
| ]) | |
| # Readout head | |
| self.readout_norm = nn.LayerNorm(config.d_model) | |
| self.readout = nn.Linear(config.d_model, config.vocab_size, bias=False) | |
| # Diffusion schedule | |
| self.schedule = MaskDiffusionSchedule(config.n_diffusion_steps) | |
| # Homeostatic plasticity: adaptive per-layer anchoring | |
| self.register_buffer('layer_rho', torch.full((config.n_layers,), config.anchor_rho)) | |
| self.register_buffer('activity_ema', torch.ones(config.n_layers)) | |
| # Lightweight settling blocks (linear attention for inner loop) | |
| self.settling_blocks = nn.ModuleList([ | |
| LightweightSettlingBlock(config.d_model, config.n_heads, config.d_ff, config.dropout) | |
| for _ in range(config.n_layers) | |
| ]) | |
| def n_active_layers(self) -> int: | |
| """Number of active layers (supports dynamic insertion/removal).""" | |
| return len(self.forward_blocks) | |
| def embed_input( | |
| self, x_t: torch.Tensor, t: torch.Tensor | |
| ) -> torch.Tensor: | |
| """Create input embeddings from corrupted tokens and timestep.""" | |
| B, S = x_t.shape | |
| positions = torch.arange(S, device=x_t.device).unsqueeze(0) | |
| h = ( | |
| self.token_embed(x_t) | |
| + self.pos_embed(positions) | |
| + self.time_embed(t).unsqueeze(1) | |
| ) | |
| return h | |
| def amortized_forward_pass( | |
| self, h_0: torch.Tensor | |
| ) -> list[torch.Tensor]: | |
| """Single feedforward pass for initialization.""" | |
| h_init = [h_0] | |
| h = h_0 | |
| for block in self.forward_blocks: | |
| h = block(h) | |
| h_init.append(h) | |
| return h_init # len = n_layers + 1 | |
| def compute_predictions( | |
| self, | |
| h: list[torch.Tensor], | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor]]: | |
| """Compute bottom-up and top-down predictions. | |
| Args: | |
| h: list of hidden states [h_0, h_1, ..., h_L] | |
| Returns: | |
| mu_up: bottom-up proposals [mu_1^up, ..., mu_L^up] | |
| mu_down: top-down predictions [mu_0^down, ..., mu_{L-1}^down] | |
| """ | |
| # Bottom-up: mu_l^up = f_l(h_{l-1}) | |
| # Always use forward_blocks — they carry trained weights. | |
| # settling_blocks (linear attention) are available for explicit long-context mode | |
| # but should NOT replace forward_blocks during standard settling. | |
| mu_up = [] | |
| for l, block in enumerate(self.forward_blocks): | |
| mu_up.append(block(h[l])) | |
| # Top-down: mu_l^down = g_l(h_{l+1}) | |
| mu_down = [] | |
| for l, fb_block in enumerate(self.feedback_blocks): | |
| mu_down.append(fb_block(h[l + 2])) # g_l predicts h_{l+1} from h_{l+2} | |
| return mu_up, mu_down | |
| def compute_energy_gradient( | |
| self, | |
| h: list[torch.Tensor], | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], float]: | |
| """Compute gradient of E_t w.r.t. each h_l. | |
| Returns: | |
| grad_h: gradient for each layer [grad_h_1, ..., grad_h_L] | |
| eps_up: bottom-up errors | |
| eps_down: top-down errors | |
| energy: scalar energy value | |
| """ | |
| config = self.config | |
| L = config.n_layers | |
| # Compute predictions | |
| mu_up, mu_down = self.compute_predictions(h) | |
| # Compute errors | |
| eps_up = [h[l + 1] - mu_up[l] for l in range(L)] | |
| eps_down = [h[l + 1] - mu_down[l] for l in range(L - 1)] | |
| # Compute precisions | |
| gamma_precs = [] # bottom-up precisions | |
| lambda_precs = [] # top-down precisions | |
| for l in range(L): | |
| tok_p, ch_p = self.precision_up[l](h[l + 1], t) | |
| gamma_precs.append((tok_p, ch_p)) | |
| for l in range(L - 1): | |
| tok_p, ch_p = self.precision_down[l](h[l + 1], t) | |
| lambda_precs.append((tok_p, ch_p)) | |
| # Energy computation — per-token normalization. | |
| # All terms use mean over D (channels), sum over B×S (batch×sequence). | |
| # This makes prediction errors and denoising on the SAME per-token scale, | |
| # so denoising_lambda directly controls their relative importance. | |
| # Gradients remain large enough (sum over B×S) for effective settling. | |
| energy = 0.0 | |
| # Bottom-up consistency: Σ_{B,S} mean_D(0.5 * ||eps||_P^2) | |
| for l in range(L): | |
| tok_p, ch_p = gamma_precs[l] | |
| weighted_err = eps_up[l] * tok_p * ch_p # (B, S, D) | |
| energy += 0.5 * (weighted_err * eps_up[l]).mean(dim=-1).sum() | |
| # Top-down consistency | |
| for l in range(L - 1): | |
| tok_p, ch_p = lambda_precs[l] | |
| weighted_err = eps_down[l] * tok_p * ch_p | |
| energy += 0.5 * (weighted_err * eps_down[l]).mean(dim=-1).sum() | |
| # Denoising loss at top layer — sum over masked tokens (already per-token). | |
| # HONEST SETTLING: entropy minimization pushes toward confident predictions | |
| # without peeking at x_0. Cross-entropy with x_0 is used in param updates only. | |
| logits = self.readout(self.readout_norm(h[L])) # (B, S, V) | |
| mask_logits = logits[mask] | |
| if mask_logits.numel() > 0: | |
| if self._honest_settling: | |
| probs = F.softmax(mask_logits, dim=-1) | |
| entropy = -(probs * (probs + 1e-10).log()).sum(dim=-1) # per-token entropy | |
| denoising_loss = entropy.sum() # sum over masked tokens | |
| else: | |
| mask_targets = x_0[mask] | |
| denoising_loss = F.cross_entropy(mask_logits, mask_targets, reduction="sum") | |
| t_frac = t.float().mean() / config.n_diffusion_steps | |
| lambda_t = config.denoising_lambda * (1.0 + (1.0 - t_frac)) | |
| energy += lambda_t * denoising_loss | |
| # Anchoring: per-token normalization | |
| for l in range(1, L + 1): | |
| rho_l = self.layer_rho[min(l - 1, len(self.layer_rho) - 1)] | |
| energy += rho_l * ((h[l] - h_init[l]) ** 2).mean(dim=-1).sum() | |
| # Lateral inhibition | |
| if config.lateral_inhibition > 0: | |
| for l in range(1, L + 1): | |
| h_norm = F.normalize(h[l], dim=-1) | |
| S_dim = h_norm.shape[1] | |
| if S_dim > 16: | |
| idx = torch.randint(0, S_dim, (16,), device=h_norm.device) | |
| sim = torch.einsum("bsd,btd->bst", h_norm[:, idx], h_norm[:, idx]) | |
| mask_diag = 1.0 - torch.eye(16, device=sim.device).unsqueeze(0) | |
| energy += config.lateral_inhibition * (F.relu(sim - 0.5) * mask_diag).sum() | |
| # Compute gradients via autograd (local operation) | |
| h_params = [h[l + 1] for l in range(L)] | |
| grads = torch.autograd.grad( | |
| energy, h_params, create_graph=False, retain_graph=False | |
| ) | |
| grad_h = list(grads) | |
| return grad_h, eps_up, eps_down, energy.item() | |
| def compute_aggregate_precision( | |
| self, | |
| h: list[torch.Tensor], | |
| t: torch.Tensor, | |
| active_mask: Optional[torch.Tensor] = None, | |
| ) -> list[float]: | |
| """Compute aggregate precision per layer for controller.""" | |
| agg_precs = [] | |
| for l in range(self.config.n_layers): | |
| tok_p, ch_p = self.precision_up[l](h[l + 1], t) | |
| # Mean precision over active tokens | |
| if active_mask is not None: | |
| mean_p = (tok_p.squeeze(-1) * active_mask.float()).sum() / ( | |
| active_mask.float().sum() + 1e-8 | |
| ) | |
| else: | |
| mean_p = tok_p.mean() | |
| agg_precs.append(mean_p.item()) | |
| return agg_precs | |
| def _resolve_param_lr( | |
| self, | |
| param_lr: Optional[float] = None, | |
| param_lr_scale: Optional[float] = None, | |
| ) -> float: | |
| """Resolve backward-compatible unified learning-rate arguments.""" | |
| if param_lr is None and param_lr_scale is None: | |
| return self.config.unified_param_lr | |
| if param_lr is None: | |
| return float(param_lr_scale) | |
| if param_lr_scale is None: | |
| return float(param_lr) | |
| if not math.isclose(float(param_lr), float(param_lr_scale), rel_tol=1e-6, abs_tol=1e-12): | |
| raise ValueError( | |
| "Received conflicting values for param_lr and param_lr_scale." | |
| ) | |
| return float(param_lr) | |
| def _compute_active_tokens( | |
| self, | |
| h: list[torch.Tensor], | |
| settling_step: int, | |
| ) -> Optional[torch.Tensor]: | |
| """Compute soft activity gates for adaptive settling.""" | |
| if settling_step <= 0: | |
| return None | |
| with torch.no_grad(): | |
| uncertainty = self.compute_token_uncertainty(h) | |
| return torch.sigmoid( | |
| (uncertainty - self.config.settling_threshold) | |
| / self.config.settling_temperature | |
| ) | |
| def _apply_state_dynamics( | |
| self, | |
| h: list[torch.Tensor], | |
| v: list[torch.Tensor], | |
| grad_h: list[torch.Tensor], | |
| agg_precs: list[float], | |
| active_tokens: Optional[torch.Tensor] = None, | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor]]: | |
| """Apply the shared SHO latent-state update used by all settling paths.""" | |
| config = self.config | |
| L = self.n_active_layers | |
| h_new = [h[0]] | |
| v_new = [] | |
| for l in range(L): | |
| prec = agg_precs[l] | |
| layer_frac = l / max(L - 1, 1) | |
| m_l = config.mass * (1.0 + config.mass_scale * layer_frac) | |
| eta_scale = 1.0 / (1.0 + config.eta_scale * layer_frac) | |
| eta = config.eta_base * eta_scale / (1.0 + config.c_eta * prec) | |
| gamma_raw = 2.0 * math.sqrt(m_l * prec) | |
| gamma = max(config.gamma_min, min(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 | |
| if active_tokens is not None: | |
| gate = active_tokens.unsqueeze(-1) | |
| h_l_new = gate * h_l_new + (1.0 - gate) * h[l + 1].detach() | |
| v_l_new = gate * v_l_new + (1.0 - gate) * v[l] | |
| v_new.append(v_l_new) | |
| h_new.append(h_l_new) | |
| return h_new, v_new | |
| def _apply_homeostatic_update(self, h: list[torch.Tensor]) -> None: | |
| """Adapt per-layer anchoring after settling.""" | |
| config = self.config | |
| if not self.training or config.homeostatic_rate <= 0: | |
| return | |
| L = self.n_active_layers | |
| with torch.no_grad(): | |
| for l in range(min(L, len(self.layer_rho))): | |
| activity = h[l + 1].detach().norm(dim=-1).mean().item() | |
| self.activity_ema[l] = 0.99 * self.activity_ema[l] + 0.01 * activity | |
| target = ( | |
| config.homeostatic_target | |
| if config.homeostatic_target > 0 | |
| else self.activity_ema[l] | |
| ) | |
| error = self.activity_ema[l] - target | |
| self.layer_rho[l] = ( | |
| self.layer_rho[l] + config.homeostatic_rate * error | |
| ).clamp(config.homeostatic_rho_min, config.homeostatic_rho_max) | |
| def _masked_cross_entropy( | |
| self, | |
| h_top: torch.Tensor, | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """Compute masked-token cross-entropy from top-layer hidden states.""" | |
| logits = self.readout(self.readout_norm(h_top)) | |
| mask_logits = logits[mask] | |
| mask_targets = x_0[mask] | |
| if mask_logits.numel() > 0: | |
| return F.cross_entropy(mask_logits, mask_targets) | |
| return torch.tensor(0.0, device=h_top.device) | |
| def measure_energy( | |
| self, | |
| h: list[torch.Tensor], | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| ) -> float: | |
| """Measure predictive-coding energy for a fixed latent state.""" | |
| with torch.enable_grad(): | |
| h_eval = [h[0].detach()] | |
| for l in range(1, self.n_active_layers + 1): | |
| h_eval.append(h[l].detach().requires_grad_(True)) | |
| _, _, _, energy = self.compute_energy_gradient(h_eval, h_init, x_0, mask, t) | |
| return energy | |
| def settling_step( | |
| self, | |
| h: list[torch.Tensor], | |
| v: list[torch.Tensor], | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| active_tokens: Optional[torch.Tensor] = None, | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor], float]: | |
| """One step of precision-conditioned second-order settling. | |
| Args: | |
| h: current hidden states [h_0, h_1, ..., h_L] | |
| v: current velocities [v_1, ..., v_L] | |
| h_init: initialization states | |
| x_0: clean tokens (for denoising loss) | |
| mask: masked positions | |
| t: timesteps | |
| active_tokens: (B, S) boolean mask of tokens still settling | |
| Returns: | |
| h_new: updated hidden states | |
| v_new: updated velocities | |
| energy: current energy value | |
| """ | |
| config = self.config | |
| L = config.n_layers | |
| # Enable grad for h (layers 1..L) - use enable_grad to work inside no_grad contexts | |
| with torch.enable_grad(): | |
| for l in range(1, L + 1): | |
| h[l] = h[l].detach().requires_grad_(True) | |
| # Compute energy gradient | |
| grad_h, eps_up, eps_down, energy = self.compute_energy_gradient( | |
| h, h_init, x_0, mask, t | |
| ) | |
| # Get aggregate precisions for controller | |
| with torch.no_grad(): | |
| agg_precs = self.compute_aggregate_precision(h, t, active_tokens) | |
| h_new, v_new = self._apply_state_dynamics( | |
| h, v, grad_h, agg_precs, active_tokens | |
| ) | |
| return h_new, v_new, energy | |
| def compute_token_uncertainty( | |
| self, h: list[torch.Tensor], eps_up: Optional[list] = None | |
| ) -> torch.Tensor: | |
| """Compute per-token uncertainty for adaptive settling. | |
| Args: | |
| h: hidden states [h_0, ..., h_L] | |
| Returns: | |
| uncertainty: (batch, seq_len) uncertainty scores | |
| """ | |
| config = self.config | |
| logits = self.readout(self.readout_norm(h[-1])) | |
| probs = F.softmax(logits, dim=-1) | |
| # Entropy of predicted distribution | |
| entropy = -(probs * (probs + 1e-10).log()).sum(dim=-1) # (B, S) | |
| # Prediction error magnitude (if available) | |
| if eps_up is not None: | |
| err_magnitude = sum( | |
| e.detach().norm(dim=-1) for e in eps_up | |
| ) * config.error_coeff | |
| else: | |
| err_magnitude = 0.0 | |
| uncertainty = config.entropy_coeff * entropy + err_magnitude | |
| return uncertainty | |
| def settle( | |
| self, | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| prev_h: Optional[list[torch.Tensor]] = None, | |
| prev_v: Optional[list[torch.Tensor]] = None, | |
| return_errors: bool = False, | |
| ): | |
| """Run K steps of second-order predictive-coding settling. | |
| Args: | |
| h_init: feedforward initialization | |
| x_0: clean tokens | |
| mask: masked positions | |
| t: timesteps | |
| prev_h: warm-start hidden states from previous diffusion step | |
| prev_v: warm-start velocities from previous diffusion step | |
| Returns: | |
| h: settled hidden states | |
| v: final velocities | |
| energies: energy trace over settling steps | |
| """ | |
| config = self.config | |
| L = config.n_layers | |
| # Initialize: warm start or from scratch | |
| if prev_h is not None and prev_v is not None: | |
| h = [h_init[0]] + [ph.detach() for ph in prev_h[1:]] | |
| v = [config.velocity_decay * pv.detach() for pv in prev_v] | |
| else: | |
| h = [hi.detach() for hi in h_init] | |
| v = [torch.zeros_like(h_init[l + 1]) for l in range(L)] | |
| energies = [] | |
| self._in_settling = True # flag for lightweight blocks | |
| for k in range(config.n_settling_steps): | |
| active_tokens = self._compute_active_tokens(h, k) | |
| h, v, energy = self.settling_step( | |
| h, v, h_init, x_0, mask, t, active_tokens | |
| ) | |
| energies.append(energy) | |
| self._in_settling = False # done settling | |
| self._apply_homeostatic_update(h) | |
| # Stationarity residual — expensive, only compute periodically | |
| stationarity_residual = None | |
| # Optionally compute final prediction errors (for online learning) | |
| final_errors = None | |
| if return_errors: | |
| with torch.no_grad(): | |
| mu_up, mu_down = self.compute_predictions(h) | |
| eps_up = [h[l + 1] - mu_up[l] for l in range(L)] | |
| eps_down = [h[l + 1] - mu_down[l] for l in range(L - 1)] | |
| final_errors = (eps_up, eps_down) | |
| return h, v, energies, stationarity_residual, final_errors | |
| def forward( | |
| self, | |
| x_0: torch.Tensor, | |
| t: Optional[torch.Tensor] = None, | |
| ) -> dict: | |
| """Training forward pass. | |
| Uses the amortized forward pass (with full gradient flow) for the | |
| training loss. Settling runs under no_grad for energy monitoring only. | |
| This ensures all forward blocks receive gradient signal from the | |
| denoising loss, not just inter-layer consistency. | |
| """ | |
| B, S = x_0.shape | |
| device = x_0.device | |
| # Sample timestep | |
| if t is None: | |
| t = torch.randint(1, self.config.n_diffusion_steps + 1, (B,), device=device) | |
| # Corrupt | |
| x_t, mask = self.schedule.corrupt(x_0, t, self.config.mask_token_id) | |
| # Embed | |
| h_0 = self.embed_input(x_t, t) | |
| # Amortized forward pass (gradients flow through all forward blocks) | |
| h_init = self.amortized_forward_pass(h_0) | |
| # Decode from amortized pass — all layers get gradient signal | |
| logits = self.readout(self.readout_norm(h_init[-1])) | |
| # Compute denoising loss | |
| mask_logits = logits[mask] | |
| mask_targets = x_0[mask] | |
| if mask_logits.numel() > 0: | |
| loss = F.cross_entropy(mask_logits, mask_targets) | |
| else: | |
| loss = torch.tensor(0.0, device=device) | |
| # Skip settling during training — it's K× extra compute for monitoring only. | |
| # Settling is used at inference time (generate method). | |
| return { | |
| "loss": loss, | |
| "logits": logits, | |
| "energies": [], | |
| "mask": mask, | |
| "h_settled": None, | |
| "v_final": None, | |
| "stationarity_residual": None, | |
| } | |
| def settled_forward( | |
| self, | |
| x_0: torch.Tensor, | |
| t: Optional[torch.Tensor] = None, | |
| prev_h: Optional[list[torch.Tensor]] = None, | |
| prev_v: Optional[list[torch.Tensor]] = None, | |
| ) -> dict: | |
| """Evaluate a batch using the settled hidden states.""" | |
| B = x_0.shape[0] | |
| device = x_0.device | |
| if t is None: | |
| t = torch.randint(1, self.config.n_diffusion_steps + 1, (B,), device=device) | |
| x_t, mask = self.schedule.corrupt(x_0, t, self.config.mask_token_id) | |
| h_0 = self.embed_input(x_t, t) | |
| h_init = self.amortized_forward_pass(h_0) | |
| h_settled, v_final, energies, stationarity, _ = self.settle( | |
| h_init, x_0, mask, t, prev_h=prev_h, prev_v=prev_v | |
| ) | |
| logits = self.readout(self.readout_norm(h_settled[-1])) | |
| mask_logits = logits[mask] | |
| mask_targets = x_0[mask] | |
| if mask_logits.numel() > 0: | |
| loss = F.cross_entropy(mask_logits, mask_targets) | |
| else: | |
| loss = torch.tensor(0.0, device=device) | |
| return { | |
| "loss": loss, | |
| "logits": logits, | |
| "energies": energies, | |
| "mask": mask, | |
| "h_settled": h_settled, | |
| "v_final": v_final, | |
| "stationarity_residual": stationarity, | |
| } | |
| def generate( | |
| self, | |
| seq_len: int, | |
| batch_size: int = 1, | |
| device: str = "cpu", | |
| online_learn: bool = False, | |
| ) -> torch.Tensor: | |
| """Generate sequences via reverse diffusion with state continuation. | |
| Args: | |
| seq_len: length of sequences to generate | |
| batch_size: number of sequences | |
| online_learn: if True, apply local weight updates during generation | |
| (the model learns from its own denoising in real time) | |
| Returns: | |
| x_0: (batch, seq_len) generated token ids | |
| """ | |
| config = self.config | |
| T = config.n_diffusion_steps | |
| mask_id = config.mask_token_id | |
| # Initialize online learning if requested | |
| if online_learn: | |
| if self._inference_updater is None: | |
| self._inference_updater = InferenceUpdater( | |
| self, lr=config.online_learn_lr, grad_clip=config.online_learn_grad_clip | |
| ) | |
| self._inference_updater.snapshot() | |
| # Use no_grad unless online learning needs gradients | |
| ctx = torch.enable_grad() if online_learn else torch.no_grad() | |
| with ctx: | |
| # Start fully masked | |
| x = torch.full( | |
| (batch_size, seq_len), mask_id, dtype=torch.long, device=device | |
| ) | |
| prev_h = None | |
| prev_v = None | |
| for t_val in range(T, 0, -1): | |
| t = torch.full((batch_size,), t_val, dtype=torch.long, device=device) | |
| # Only process still-masked positions | |
| mask = (x == mask_id) | |
| if not mask.any(): | |
| break | |
| # Embed | |
| h_0 = self.embed_input(x, t) | |
| # Amortized pass | |
| h_init = self.amortized_forward_pass(h_0) | |
| # Settle with warm start (return errors if online learning) | |
| h_settled, v_final, energies, _, errors = self.settle( | |
| h_init, x, mask, t, | |
| prev_h=prev_h, prev_v=prev_v, | |
| return_errors=online_learn, | |
| ) | |
| # Online learning: update weights from settling errors | |
| if online_learn and errors is not None and len(energies) >= 2: | |
| energy_ratio = energies[-1] / max(energies[0], 1e-8) | |
| if energy_ratio < config.online_learn_min_energy_ratio: | |
| eps_up, eps_down = errors | |
| self._inference_updater.update_from_errors( | |
| h_settled, eps_up, eps_down, x, mask, t | |
| ) | |
| # State continuation | |
| prev_h = h_settled | |
| prev_v = v_final | |
| # Decode | |
| logits = self.readout(self.readout_norm(h_settled[-1])) | |
| # Determine how many tokens to unmask this step | |
| n_masked = mask.sum(dim=-1) | |
| t_prev = t_val - 1 | |
| target_mask_rate = self.schedule.mask_rate( | |
| torch.tensor([t_prev], device=device) | |
| ).squeeze() | |
| n_to_keep_masked = (target_mask_rate * seq_len).long() | |
| # For each sequence, unmask highest-confidence masked tokens | |
| probs = F.softmax(logits, dim=-1) | |
| confidence = probs.max(dim=-1).values # (B, S) | |
| confidence[~mask] = float("inf") # Don't re-unmask | |
| for b in range(batch_size): | |
| b_mask_idx = mask[b].nonzero(as_tuple=True)[0] | |
| if len(b_mask_idx) == 0: | |
| continue | |
| b_conf = confidence[b, b_mask_idx] | |
| n_to_unmask = max(0, len(b_mask_idx) - n_to_keep_masked.item()) | |
| if n_to_unmask > 0: | |
| _, top_idx = b_conf.topk(min(n_to_unmask, len(b_conf))) | |
| unmask_positions = b_mask_idx[top_idx] | |
| predicted_tokens = logits[b, unmask_positions].argmax(dim=-1) | |
| x[b, unmask_positions] = predicted_tokens | |
| # Final: unmask any remaining | |
| still_masked = (x == mask_id) | |
| if still_masked.any(): | |
| t_final = torch.ones(batch_size, dtype=torch.long, device=device) | |
| h_0 = self.embed_input(x, t_final) | |
| h_init = self.amortized_forward_pass(h_0) | |
| h_settled, _, _, _, _ = self.settle( | |
| h_init, x, still_masked, t_final, | |
| prev_h=prev_h, prev_v=prev_v, | |
| ) | |
| logits = self.readout(self.readout_norm(h_settled[-1])) | |
| final_tokens = logits[still_masked].argmax(dim=-1) | |
| x[still_masked] = final_tokens | |
| return x | |
| def unified_step( | |
| self, | |
| h: list[torch.Tensor], | |
| v: list[torch.Tensor], | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| param_lr: Optional[float] = None, | |
| param_lr_scale: Optional[float] = None, | |
| active_tokens: Optional[torch.Tensor] = None, | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor], float]: | |
| """Experimental microstep-coupled inference-and-learning primitive. | |
| The repaired training stack uses ``unified_train_batch`` for | |
| post-settle updates. This method is retained for modules that | |
| explicitly want a per-microstep in-place parameter update. | |
| """ | |
| config = self.config | |
| scaled_lr = self._resolve_param_lr(param_lr, param_lr_scale) | |
| L = self.n_active_layers | |
| # Track energies for adaptive two-timescale ratio | |
| if not hasattr(self, '_unified_energies'): | |
| self._unified_energies = [] | |
| # Enable grad for h | |
| with torch.enable_grad(): | |
| for l in range(1, L + 1): | |
| h[l] = h[l].detach().requires_grad_(True) | |
| # Compute energy gradient w.r.t. hidden states | |
| grad_h, eps_up, eps_down, energy = self.compute_energy_gradient( | |
| h, h_init, x_0, mask, t | |
| ) | |
| self._unified_energies.append(energy) | |
| # Get precisions for controller | |
| with torch.no_grad(): | |
| agg_precs = self.compute_aggregate_precision(h, t, active_tokens) | |
| # Fast latent dynamics use the same controller as canonical settling. | |
| h_new, v_new = self._apply_state_dynamics( | |
| h, v, grad_h, agg_precs, active_tokens | |
| ) | |
| # Slow parameter updates remain experimental in this primitive. | |
| for l in range(L): | |
| h_below = h[l].detach() | |
| h_current = h[l + 1].detach() | |
| # Bottom-up: local gradient for forward block | |
| mu_up = self.forward_blocks[l](h_below) | |
| with torch.no_grad(): | |
| tok_p, ch_p = self.precision_up[l](h_current, t) | |
| eps = h_current - mu_up | |
| local_loss = 0.5 * (eps * tok_p * ch_p * eps).mean(dim=-1).sum() | |
| local_loss.backward() | |
| with torch.no_grad(): | |
| nn.utils.clip_grad_norm_(self.forward_blocks[l].parameters(), 1.0) | |
| for p in self.forward_blocks[l].parameters(): | |
| if p.grad is not None: | |
| p.data -= scaled_lr * p.grad | |
| p.grad.zero_() | |
| # Top-down feedback | |
| for l in range(L - 1): | |
| h_above = h[l + 2].detach() | |
| h_current = h[l + 1].detach() | |
| mu_down = self.feedback_blocks[l](h_above) | |
| with torch.no_grad(): | |
| tok_p, ch_p = self.precision_down[l](h_current, t) | |
| eps = h_current - mu_down | |
| local_loss = 0.5 * (eps * tok_p * ch_p * eps).mean(dim=-1).sum() | |
| local_loss.backward() | |
| with torch.no_grad(): | |
| nn.utils.clip_grad_norm_(self.feedback_blocks[l].parameters(), 1.0) | |
| for p in self.feedback_blocks[l].parameters(): | |
| if p.grad is not None: | |
| p.data -= scaled_lr * p.grad | |
| p.grad.zero_() | |
| # Task-directed update: readout loss propagates through all blocks. | |
| # This gives every component (forward blocks, readout, embeddings) | |
| # direct task signal — the cross-entropy gradient tells each layer | |
| # what features to produce for correct token prediction. | |
| # Combined with the local consistency updates above, each forward block | |
| # receives TWO gradients: (1) local prediction error and (2) task loss. | |
| x_t_input = x_0.clone() | |
| x_t_input[mask] = config.mask_token_id | |
| h_0_fresh = self.embed_input(x_t_input, t) | |
| h_through = h_0_fresh | |
| for block in self.forward_blocks: | |
| h_through = block(h_through) | |
| logits = self.readout(self.readout_norm(h_through)) | |
| mask_logits = logits[mask] | |
| mask_targets = x_0[mask] | |
| if mask_logits.numel() > 0: | |
| readout_loss = F.cross_entropy(mask_logits, mask_targets) | |
| readout_loss.backward() | |
| # Update ALL trainable components with task signal | |
| all_params = list(self.parameters()) | |
| with torch.no_grad(): | |
| nn.utils.clip_grad_norm_(all_params, 1.0) | |
| for p in all_params: | |
| if p.grad is not None: | |
| p.data -= scaled_lr * p.grad | |
| p.grad.zero_() | |
| # Precision heads: calibrate using settled error statistics | |
| prec_params = list(self.precision_up.parameters()) + list(self.precision_down.parameters()) | |
| prec_loss = 0.0 | |
| for l in range(L): | |
| tok_p, ch_p = self.precision_up[l](h[l + 1].detach(), t) | |
| eps = (h[l + 1].detach() - self.forward_blocks[l](h[l].detach()).detach()) | |
| # Gaussian NLL: 0.5 * P * eps^2 - 0.5 * log(P) | |
| prec_loss = prec_loss + 0.5 * (eps * tok_p * ch_p * eps).mean(dim=-1).sum() | |
| prec_loss = prec_loss - 0.5 * (tok_p.log().mean(dim=-1).sum() + ch_p.log().mean(dim=-1).sum()) | |
| if isinstance(prec_loss, torch.Tensor) and prec_loss.requires_grad: | |
| prec_loss.backward() | |
| with torch.no_grad(): | |
| nn.utils.clip_grad_norm_(prec_params, 1.0) | |
| for p in prec_params: | |
| if p.grad is not None: | |
| p.data -= scaled_lr * 0.1 * p.grad # slower for precision | |
| p.grad.zero_() | |
| return h_new, v_new, energy | |
| def unified_settle( | |
| self, | |
| h_init: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| param_lr: Optional[float] = None, | |
| param_lr_scale: Optional[float] = None, | |
| prev_h: Optional[list[torch.Tensor]] = None, | |
| prev_v: Optional[list[torch.Tensor]] = None, | |
| ) -> Tuple[list[torch.Tensor], list[torch.Tensor], list[float]]: | |
| """Run K experimental microstep-coupled unified updates.""" | |
| config = self.config | |
| L = self.n_active_layers | |
| scaled_lr = self._resolve_param_lr(param_lr, param_lr_scale) | |
| self._unified_energies = [] # reset | |
| if prev_h is not None and prev_v is not None: | |
| h = [h_init[0]] + [ph.detach() for ph in prev_h[1:]] | |
| v = [config.velocity_decay * pv.detach() for pv in prev_v] | |
| else: | |
| h = [hi.detach() for hi in h_init] | |
| v = [torch.zeros_like(h_init[l + 1]) for l in range(L)] | |
| energies = [] | |
| for k in range(config.n_settling_steps): | |
| active_tokens = self._compute_active_tokens(h, k) | |
| h, v, energy = self.unified_step( | |
| h, v, h_init, x_0, mask, t, | |
| param_lr=scaled_lr, | |
| active_tokens=active_tokens, | |
| ) | |
| energies.append(energy) | |
| self._apply_homeostatic_update(h) | |
| return h, v, energies | |
| def unified_train_batch( | |
| self, | |
| x_0: torch.Tensor, | |
| t: Optional[torch.Tensor] = None, | |
| param_lr: Optional[float] = None, | |
| param_lr_scale: Optional[float] = None, | |
| ) -> dict: | |
| """Canonical unified training step: settle first, then update.""" | |
| lr = self._resolve_param_lr(param_lr, param_lr_scale) | |
| if self._unified_batch_updater is None: | |
| self._unified_batch_updater = UnifiedParameterUpdater( | |
| self, | |
| lr_forward=lr, | |
| lr_feedback=lr, | |
| lr_readout=lr, | |
| lr_precision=lr * self.config.unified_precision_lr_scale, | |
| lr_embed=lr, | |
| task_lr_scale=self.config.unified_task_lr_scale, | |
| ) | |
| self._unified_batch_updater.set_learning_rates( | |
| lr_forward=lr, | |
| lr_feedback=lr, | |
| lr_readout=lr, | |
| lr_precision=lr * self.config.unified_precision_lr_scale, | |
| lr_embed=lr, | |
| task_lr_scale=self.config.unified_task_lr_scale, | |
| ) | |
| return self._unified_batch_updater.step({"input_ids": x_0}, t=t) | |
| def post_settle_update( | |
| self, | |
| h_settled: list[torch.Tensor], | |
| x_input: torch.Tensor, | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| param_lr: Optional[float] = None, | |
| param_lr_scale: Optional[float] = None, | |
| stationarity: float = 0.0, | |
| energies: Optional[list[float]] = None, | |
| ) -> dict: | |
| """Apply canonical unified updates from an externally settled state.""" | |
| lr = self._resolve_param_lr(param_lr, param_lr_scale) | |
| if self._unified_batch_updater is None: | |
| self._unified_batch_updater = UnifiedParameterUpdater( | |
| self, | |
| lr_forward=lr, | |
| lr_feedback=lr, | |
| lr_readout=lr, | |
| lr_precision=lr * self.config.unified_precision_lr_scale, | |
| lr_embed=lr, | |
| task_lr_scale=self.config.unified_task_lr_scale, | |
| ) | |
| self._unified_batch_updater.set_learning_rates( | |
| lr_forward=lr, | |
| lr_feedback=lr, | |
| lr_readout=lr, | |
| lr_precision=lr * self.config.unified_precision_lr_scale, | |
| lr_embed=lr, | |
| task_lr_scale=self.config.unified_task_lr_scale, | |
| ) | |
| return self._unified_batch_updater.step_from_settled( | |
| h_settled, | |
| x_input=x_input, | |
| x_0=x_0, | |
| mask=mask, | |
| t=t, | |
| stationarity=stationarity, | |
| energies=energies, | |
| ) | |
| # ============================================================================= | |
| # Local Parameter Update (for training without global backprop) | |
| # ============================================================================= | |
| class LocalParameterUpdater: | |
| """Implements local parameter updates for PC-SHO-DLM. | |
| After settling, updates each layer's parameters using only | |
| local prediction errors and local Jacobians. | |
| This is the 'globally backprop-free' training procedure. | |
| """ | |
| def __init__( | |
| self, | |
| model: PCSHODLM, | |
| lr_forward: float = 1e-4, | |
| lr_feedback: float = 1e-4, | |
| lr_readout: float = 1e-4, | |
| lr_precision: float = 1e-4, | |
| lr_embed: Optional[float] = None, | |
| ): | |
| self.model = model | |
| self.lr_forward = lr_forward | |
| self.lr_feedback = lr_feedback | |
| self.lr_readout = lr_readout | |
| self.lr_precision = lr_precision | |
| self.lr_embed = lr_forward if lr_embed is None else lr_embed | |
| # Separate optimizers for each component | |
| self.forward_optimizers = [ | |
| torch.optim.AdamW(block.parameters(), lr=lr_forward) | |
| for block in model.forward_blocks | |
| ] | |
| self.feedback_optimizers = [ | |
| torch.optim.AdamW(block.parameters(), lr=lr_feedback) | |
| for block in model.feedback_blocks | |
| ] | |
| self.readout_optimizer = torch.optim.AdamW( | |
| list(model.readout.parameters()) + list(model.readout_norm.parameters()), | |
| lr=lr_readout, | |
| ) | |
| self.precision_optimizer = torch.optim.AdamW( | |
| list(model.precision_up.parameters()) | |
| + list(model.precision_down.parameters()), | |
| lr=lr_precision, | |
| ) | |
| self.embed_optimizer = torch.optim.AdamW( | |
| list(model.token_embed.parameters()) | |
| + list(model.pos_embed.parameters()) | |
| + list(model.time_embed.parameters()), | |
| lr=self.lr_embed, | |
| ) | |
| def set_learning_rates( | |
| self, | |
| lr_forward: Optional[float] = None, | |
| lr_feedback: Optional[float] = None, | |
| lr_readout: Optional[float] = None, | |
| lr_precision: Optional[float] = None, | |
| lr_embed: Optional[float] = None, | |
| task_lr_scale: Optional[float] = None, | |
| ) -> None: | |
| """Update optimizer learning rates without recreating state.""" | |
| if lr_forward is not None: | |
| self.lr_forward = lr_forward | |
| for opt in self.forward_optimizers: | |
| for group in opt.param_groups: | |
| group["lr"] = lr_forward | |
| if lr_feedback is not None: | |
| self.lr_feedback = lr_feedback | |
| for opt in self.feedback_optimizers: | |
| for group in opt.param_groups: | |
| group["lr"] = lr_feedback | |
| if lr_readout is not None: | |
| self.lr_readout = lr_readout | |
| for group in self.readout_optimizer.param_groups: | |
| group["lr"] = lr_readout | |
| if lr_precision is not None: | |
| self.lr_precision = lr_precision | |
| for group in self.precision_optimizer.param_groups: | |
| group["lr"] = lr_precision | |
| if lr_embed is not None: | |
| self.lr_embed = lr_embed | |
| for group in self.embed_optimizer.param_groups: | |
| group["lr"] = lr_embed | |
| def local_update( | |
| self, | |
| h_settled: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| update_readout: bool = True, | |
| ) -> float: | |
| """Perform local parameter updates using settled prediction errors. | |
| This is the key: we only need local Jacobians, not global backprop. | |
| """ | |
| model = self.model | |
| config = model.config | |
| L = config.n_layers | |
| readout_loss_value = 0.0 | |
| # Re-enable gradients for parameters | |
| # Compute fresh predictions from settled states | |
| for l in range(L): | |
| h_below = h_settled[l].detach() | |
| h_current = h_settled[l + 1].detach() | |
| # Bottom-up: local loss for forward block | |
| self.forward_optimizers[l].zero_grad() | |
| mu_up = model.forward_blocks[l](h_below) | |
| # Get precision | |
| with torch.no_grad(): | |
| tok_p, ch_p = model.precision_up[l](h_current, t) | |
| # Precision-weighted error | |
| eps_up = h_current - mu_up | |
| local_loss_up = 0.5 * (eps_up * tok_p * ch_p * eps_up).mean(dim=-1).sum() | |
| local_loss_up.backward() | |
| torch.nn.utils.clip_grad_norm_(model.forward_blocks[l].parameters(), 1.0) | |
| self.forward_optimizers[l].step() | |
| # Top-down: local loss for feedback blocks | |
| for l in range(L - 1): | |
| self.feedback_optimizers[l].zero_grad() | |
| h_above = h_settled[l + 2].detach() | |
| h_current = h_settled[l + 1].detach() | |
| mu_down = model.feedback_blocks[l](h_above) | |
| with torch.no_grad(): | |
| tok_p, ch_p = model.precision_down[l](h_current, t) | |
| eps_down = h_current - mu_down | |
| local_loss_down = 0.5 * (eps_down * tok_p * ch_p * eps_down).mean(dim=-1).sum() | |
| local_loss_down.backward() | |
| torch.nn.utils.clip_grad_norm_(model.feedback_blocks[l].parameters(), 1.0) | |
| self.feedback_optimizers[l].step() | |
| if update_readout: | |
| self.readout_optimizer.zero_grad() | |
| readout_loss = model._masked_cross_entropy(h_settled[L].detach(), x_0, mask) | |
| if readout_loss.requires_grad: | |
| readout_loss.backward() | |
| torch.nn.utils.clip_grad_norm_( | |
| list(model.readout.parameters()) + list(model.readout_norm.parameters()), | |
| 1.0, | |
| ) | |
| self.readout_optimizer.step() | |
| readout_loss_value = float(readout_loss.item()) | |
| return readout_loss_value | |
| def task_update( | |
| self, | |
| x_t: torch.Tensor, | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| *, | |
| train_forward: bool, | |
| train_embed: bool, | |
| train_readout: bool, | |
| ) -> float: | |
| """Apply task loss to the amortized path using explicit optimizer groups.""" | |
| model = self.model | |
| for opt in self.forward_optimizers: | |
| opt.zero_grad() | |
| self.embed_optimizer.zero_grad() | |
| self.readout_optimizer.zero_grad() | |
| with torch.enable_grad(): | |
| h_0_fresh = model.embed_input(x_t, t) | |
| h_init_fresh = model.amortized_forward_pass(h_0_fresh) | |
| task_loss = model._masked_cross_entropy(h_init_fresh[model.config.n_layers], x_0, mask) | |
| if task_loss.requires_grad: | |
| task_loss.backward() | |
| if train_forward: | |
| torch.nn.utils.clip_grad_norm_(model.forward_blocks.parameters(), 1.0) | |
| for opt in self.forward_optimizers: | |
| opt.step() | |
| if train_embed: | |
| torch.nn.utils.clip_grad_norm_( | |
| list(model.token_embed.parameters()) | |
| + list(model.pos_embed.parameters()) | |
| + list(model.time_embed.parameters()), | |
| 1.0, | |
| ) | |
| self.embed_optimizer.step() | |
| if train_readout: | |
| torch.nn.utils.clip_grad_norm_( | |
| list(model.readout.parameters()) + list(model.readout_norm.parameters()), | |
| 1.0, | |
| ) | |
| self.readout_optimizer.step() | |
| if not train_forward: | |
| for opt in self.forward_optimizers: | |
| opt.zero_grad() | |
| if not train_embed: | |
| self.embed_optimizer.zero_grad() | |
| if not train_readout: | |
| self.readout_optimizer.zero_grad() | |
| return float(task_loss.item()) | |
| return 0.0 | |
| def precision_update( | |
| self, | |
| h_settled: list[torch.Tensor], | |
| t: torch.Tensor, | |
| ) -> float: | |
| """Calibrate precision heads from settled prediction errors.""" | |
| model = self.model | |
| self.precision_optimizer.zero_grad() | |
| for l in range(1, model.config.n_layers + 1): | |
| h_settled[l] = h_settled[l].detach() | |
| mu_up_list, _ = model.compute_predictions(h_settled) | |
| prec_loss = 0.0 | |
| for l in range(model.config.n_layers): | |
| tok_p, ch_p = model.precision_up[l](h_settled[l + 1].detach(), t) | |
| eps = h_settled[l + 1].detach() - mu_up_list[l].detach() | |
| prec_loss = prec_loss + 0.5 * (eps * tok_p * ch_p * eps).mean() | |
| prec_loss = prec_loss - 0.5 * (tok_p.log().mean() + ch_p.log().mean()) | |
| if isinstance(prec_loss, torch.Tensor) and prec_loss.requires_grad: | |
| prec_loss.backward() | |
| torch.nn.utils.clip_grad_norm_( | |
| list(model.precision_up.parameters()) + list(model.precision_down.parameters()), | |
| 1.0, | |
| ) | |
| self.precision_optimizer.step() | |
| return float(prec_loss.item()) | |
| return 0.0 | |
| def step(self, batch: dict, t: Optional[torch.Tensor] = None): | |
| """Complete training step: corrupt, settle, locally update.""" | |
| model = self.model | |
| x_0 = batch["input_ids"] | |
| B, S = x_0.shape | |
| device = x_0.device | |
| # Sample timestep | |
| if t is None: | |
| t = torch.randint(1, model.config.n_diffusion_steps + 1, (B,), device=device) | |
| # Corrupt | |
| 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 (no global gradients needed) | |
| with torch.no_grad(): | |
| h_settled, v_final, energies, stationarity, _ = model.settle(h_init, x_0, mask, t) | |
| # Local updates | |
| self.local_update(h_settled, x_0, mask, t, update_readout=True) | |
| self.task_update( | |
| x_t, x_0, mask, t, | |
| train_forward=False, | |
| train_embed=True, | |
| train_readout=False, | |
| ) | |
| precision_loss = self.precision_update(h_settled, t) | |
| with torch.no_grad(): | |
| settled_eval = float(model._masked_cross_entropy(h_settled[model.config.n_layers], x_0, mask).item()) | |
| h_0_eval = model.embed_input(x_t, t) | |
| h_init_eval = model.amortized_forward_pass(h_0_eval) | |
| amortized_eval = float( | |
| model._masked_cross_entropy(h_init_eval[model.config.n_layers], x_0, mask).item() | |
| ) | |
| return { | |
| "loss": settled_eval, | |
| "settled_loss": settled_eval, | |
| "amortized_loss": amortized_eval, | |
| "precision_loss": precision_loss, | |
| "energies": energies, | |
| "stationarity": stationarity if stationarity else 0.0, | |
| "n_masked": mask.sum().item(), | |
| } | |
| class UnifiedParameterUpdater(LocalParameterUpdater): | |
| """Canonical unified trainer: settle first, then update all trainable components.""" | |
| def __init__( | |
| self, | |
| model: PCSHODLM, | |
| lr_forward: float = 1e-3, | |
| lr_feedback: float = 1e-3, | |
| lr_readout: float = 1e-3, | |
| lr_precision: float = 1e-4, | |
| lr_embed: float = 1e-3, | |
| task_lr_scale: float = 1.0, | |
| ): | |
| super().__init__( | |
| model, | |
| lr_forward=lr_forward, | |
| lr_feedback=lr_feedback, | |
| lr_readout=lr_readout, | |
| lr_precision=lr_precision, | |
| lr_embed=lr_embed, | |
| ) | |
| self.task_lr_scale = task_lr_scale | |
| def set_learning_rates( | |
| self, | |
| lr_forward: Optional[float] = None, | |
| lr_feedback: Optional[float] = None, | |
| lr_readout: Optional[float] = None, | |
| lr_precision: Optional[float] = None, | |
| lr_embed: Optional[float] = None, | |
| task_lr_scale: Optional[float] = None, | |
| ) -> None: | |
| super().set_learning_rates( | |
| lr_forward=lr_forward, | |
| lr_feedback=lr_feedback, | |
| lr_readout=lr_readout, | |
| lr_precision=lr_precision, | |
| lr_embed=lr_embed, | |
| ) | |
| if task_lr_scale is not None: | |
| self.task_lr_scale = task_lr_scale | |
| def step(self, batch: dict, t: Optional[torch.Tensor] = None): | |
| """Unified training step with post-settle local updates and amortized task learning.""" | |
| model = self.model | |
| x_0 = batch["input_ids"] | |
| B, S = x_0.shape | |
| device = x_0.device | |
| if t is None: | |
| 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) | |
| with torch.no_grad(): | |
| h_settled, v_final, energies, stationarity, _ = model.settle(h_init, x_0, mask, t) | |
| return self.step_from_settled( | |
| h_settled, | |
| x_input=x_t, | |
| x_0=x_0, | |
| mask=mask, | |
| t=t, | |
| stationarity=stationarity if stationarity else 0.0, | |
| energies=energies, | |
| v_final=v_final, | |
| ) | |
| def step_from_settled( | |
| self, | |
| h_settled: list[torch.Tensor], | |
| *, | |
| x_input: torch.Tensor, | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| stationarity: float = 0.0, | |
| energies: Optional[list[float]] = None, | |
| v_final: Optional[list[torch.Tensor]] = None, | |
| ) -> dict: | |
| """Apply unified updates from a caller-provided settled state.""" | |
| model = self.model | |
| self.local_update(h_settled, x_0, mask, t, update_readout=True) | |
| task_lr_scale = self.task_lr_scale | |
| base_forward_lr = self.lr_forward | |
| base_embed_lr = self.lr_embed | |
| if task_lr_scale != 1.0: | |
| self.set_learning_rates( | |
| lr_forward=base_forward_lr * task_lr_scale, | |
| lr_embed=base_embed_lr * task_lr_scale, | |
| ) | |
| task_loss = self.task_update( | |
| x_input, x_0, mask, t, | |
| train_forward=True, | |
| train_embed=True, | |
| train_readout=False, | |
| ) | |
| if task_lr_scale != 1.0: | |
| self.set_learning_rates( | |
| lr_forward=base_forward_lr, | |
| lr_embed=base_embed_lr, | |
| ) | |
| precision_loss = self.precision_update(h_settled, t) | |
| with torch.no_grad(): | |
| settled_eval = float(model._masked_cross_entropy(h_settled[model.config.n_layers], x_0, mask).item()) | |
| h_0_eval = model.embed_input(x_input, t) | |
| h_init_eval = model.amortized_forward_pass(h_0_eval) | |
| amortized_eval = float( | |
| model._masked_cross_entropy(h_init_eval[model.config.n_layers], x_0, mask).item() | |
| ) | |
| return { | |
| "loss": settled_eval, | |
| "settled_loss": settled_eval, | |
| "amortized_loss": amortized_eval, | |
| "task_loss": task_loss, | |
| "precision_loss": precision_loss, | |
| "energies": energies or [], | |
| "stationarity": stationarity if stationarity else 0.0, | |
| "n_masked": mask.sum().item(), | |
| "h_settled": h_settled, | |
| "v_final": v_final, | |
| "mask": mask, | |
| "t": t, | |
| "energy": energies[-1] if energies else model.measure_energy(h_settled, h_settled, x_0, mask, t), | |
| } | |
| # ============================================================================= | |
| # Online Learning During Inference | |
| # ============================================================================= | |
| class InferenceUpdater: | |
| """Lightweight local updater for online learning during generation. | |
| Uses prediction errors already computed during settling — no additional | |
| forward/backward passes needed. This is the key PC advantage: the settling | |
| process computes everything required for weight updates. | |
| Uses SGD with momentum (not AdamW) to minimize memory overhead at inference. | |
| """ | |
| def __init__(self, model: PCSHODLM, lr: float = 1e-5, grad_clip: float = 0.1): | |
| self.model = model | |
| self.lr = lr | |
| self.grad_clip = grad_clip | |
| self._snapshot = None | |
| # Per-layer SGD optimizers (lightweight) | |
| self.forward_opts = [ | |
| torch.optim.SGD(block.parameters(), lr=lr, momentum=0.9) | |
| for block in model.forward_blocks | |
| ] | |
| self.feedback_opts = [ | |
| torch.optim.SGD(block.parameters(), lr=lr, momentum=0.9) | |
| for block in model.feedback_blocks | |
| ] | |
| self.readout_opt = torch.optim.SGD( | |
| list(model.readout.parameters()) + list(model.readout_norm.parameters()), | |
| lr=lr, momentum=0.9, | |
| ) | |
| def snapshot(self): | |
| """Save parameter snapshot for potential rollback.""" | |
| self._snapshot = { | |
| name: p.data.clone() for name, p in self.model.named_parameters() | |
| } | |
| def rollback(self): | |
| """Restore parameters from snapshot.""" | |
| if self._snapshot is None: | |
| return | |
| for name, p in self.model.named_parameters(): | |
| if name in self._snapshot: | |
| p.data.copy_(self._snapshot[name]) | |
| def update_from_errors( | |
| self, | |
| h_settled: list[torch.Tensor], | |
| eps_up: list[torch.Tensor], | |
| eps_down: list[torch.Tensor], | |
| x_0: torch.Tensor, | |
| mask: torch.Tensor, | |
| t: torch.Tensor, | |
| ) -> float: | |
| """Apply local weight updates using pre-computed prediction errors. | |
| Returns total update norm for monitoring. | |
| """ | |
| model = self.model | |
| L = model.n_active_layers | |
| total_norm = 0.0 | |
| # Bottom-up blocks: minimize precision-weighted eps_up | |
| for l in range(L): | |
| self.forward_opts[l].zero_grad() | |
| h_below = h_settled[l].detach() | |
| mu_up = model.forward_blocks[l](h_below) | |
| with torch.no_grad(): | |
| tok_p, ch_p = model.precision_up[l](h_settled[l + 1].detach(), t) | |
| eps = h_settled[l + 1].detach() - mu_up | |
| loss = 0.5 * (eps * tok_p * ch_p * eps).sum() / x_0.shape[0] | |
| loss.backward() | |
| nn.utils.clip_grad_norm_(model.forward_blocks[l].parameters(), self.grad_clip) | |
| self.forward_opts[l].step() | |
| total_norm += loss.item() | |
| # Top-down blocks | |
| for l in range(L - 1): | |
| self.feedback_opts[l].zero_grad() | |
| h_above = h_settled[l + 2].detach() | |
| mu_down = model.feedback_blocks[l](h_above) | |
| with torch.no_grad(): | |
| tok_p, ch_p = model.precision_down[l](h_settled[l + 1].detach(), t) | |
| eps = h_settled[l + 1].detach() - mu_down | |
| loss = 0.5 * (eps * tok_p * ch_p * eps).sum() / x_0.shape[0] | |
| loss.backward() | |
| nn.utils.clip_grad_norm_(model.feedback_blocks[l].parameters(), self.grad_clip) | |
| self.feedback_opts[l].step() | |
| # Readout | |
| self.readout_opt.zero_grad() | |
| logits = model.readout(model.readout_norm(h_settled[L].detach())) | |
| mask_logits = logits[mask] | |
| mask_targets = x_0[mask] | |
| if mask_logits.numel() > 0: | |
| rl = F.cross_entropy(mask_logits, mask_targets) | |
| rl.backward() | |
| nn.utils.clip_grad_norm_( | |
| list(model.readout.parameters()) + list(model.readout_norm.parameters()), | |
| self.grad_clip, | |
| ) | |
| self.readout_opt.step() | |
| return total_norm | |
| # ============================================================================= | |
| # Dynamic Layer Manager | |
| # ============================================================================= | |
| class DynamicLayerManager: | |
| """Manages runtime insertion and removal of layers. | |
| Exploits PC locality: a new layer only needs to become consistent | |
| with its immediate neighbors. Distant layers are unaffected. | |
| """ | |
| def __init__(self, model: PCSHODLM): | |
| self.model = model | |
| self.history: list[dict] = [] | |
| def insert_layer(self, position: int, init_strategy: str = "identity") -> None: | |
| """Insert a new layer at the given position. | |
| Args: | |
| position: index in the layer stack (0-indexed) | |
| init_strategy: "identity" (near-identity init) or "random" | |
| """ | |
| model = self.model | |
| config = model.config | |
| device = next(model.parameters()).device | |
| # Create new modules | |
| new_forward = BidirectionalTransformerBlock( | |
| config.d_model, config.n_heads, config.d_ff, config.dropout | |
| ).to(device) | |
| new_feedback = FeedbackPredictor(config.d_model, config.feedback_rank).to(device) | |
| new_prec_up = PrecisionHead(config.d_model, config.n_diffusion_steps).to(device) | |
| new_prec_down = PrecisionHead(config.d_model, config.n_diffusion_steps).to(device) | |
| # Identity initialization: zero out the residual path so output ≈ input | |
| if init_strategy == "identity": | |
| with torch.no_grad(): | |
| # Zero the final linear in the feedforward and attention output | |
| for name, param in new_forward.named_parameters(): | |
| if "ff.3" in name or "ff.2" in name: # last linear in FF | |
| param.zero_() | |
| # Insert into module lists | |
| model.forward_blocks.insert(position, new_forward) | |
| model.precision_up.insert(position, new_prec_up) | |
| # Feedback and precision_down have n_layers-1 entries | |
| fb_pos = min(position, len(model.feedback_blocks)) | |
| model.feedback_blocks.insert(fb_pos, new_feedback) | |
| model.precision_down.insert(min(position, len(model.precision_down)), new_prec_down) | |
| # Update config | |
| config.n_layers += 1 | |
| self.history.append({ | |
| "action": "insert", | |
| "position": position, | |
| "strategy": init_strategy, | |
| "new_n_layers": config.n_layers, | |
| }) | |
| def remove_layer(self, position: int) -> None: | |
| """Remove a layer at the given position.""" | |
| model = self.model | |
| config = model.config | |
| if config.n_layers <= 2: | |
| raise ValueError("Cannot remove layer: minimum 2 layers required") | |
| del model.forward_blocks[position] | |
| del model.precision_up[position] | |
| fb_pos = min(position, len(model.feedback_blocks) - 1) | |
| del model.feedback_blocks[fb_pos] | |
| del model.precision_down[min(position, len(model.precision_down) - 1)] | |
| config.n_layers -= 1 | |
| self.history.append({ | |
| "action": "remove", | |
| "position": position, | |
| "new_n_layers": config.n_layers, | |
| }) | |
| # ============================================================================= | |
| # Utility: Model Size | |
| # ============================================================================= | |
| def count_parameters(model: nn.Module) -> int: | |
| return sum(p.numel() for p in model.parameters() if p.requires_grad) | |
| if __name__ == "__main__": | |
| # Quick sanity check | |
| config = PCSHOConfig( | |
| vocab_size=1000, | |
| max_seq_len=128, | |
| d_model=256, | |
| n_heads=4, | |
| n_layers=4, | |
| d_ff=512, | |
| n_diffusion_steps=100, | |
| n_settling_steps=4, | |
| ) | |
| model = PCSHODLM(config) | |
| print(f"PC-SHO-DLM parameters: {count_parameters(model):,}") | |
| # Test forward pass | |
| x_0 = torch.randint(1, 1000, (2, 64)) | |
| output = model(x_0) | |
| print(f"Loss: {output['loss'].item():.4f}") | |
| print(f"Energy trace: {output['energies']}") | |
| print(f"Logits shape: {output['logits'].shape}") | |
| # Test generation | |
| generated = model.generate(seq_len=32, batch_size=2) | |
| print(f"Generated shape: {generated.shape}") | |
| print(f"Generated sample: {generated[0, :10].tolist()}") | |