pc-sho-dlm-code / src /model.py
zotowata's picture
Ensure at least one masked token per sample
8ea617a verified
Raw History Blame Contribute Delete
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
# =============================================================================
@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()}")