""" 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 # ============================================================================= @dataclass 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) ]) @property 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, } @torch.no_grad() 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()}")