""" Token Information State Module for Q-TensorFormer. Constructs an 8-dimensional normalized information state vector per token: z_t = [S_t, H_t, U_t, A_t, R_t, L_t, M_t, B_t] Signals: - S_t: Attention dispersion / entanglement entropy proxy - H_t: Predictive entropy (-sum p log p from logits) - U_t: Epistemic uncertainty / representation variance - A_t: Cumulative attention importance received by token - R_t: Tensor-Train reconstruction error under current rank - L_t: Latency pressure [0, 1] from hardware cost model - M_t: Memory pressure [0, 1] (DRAM / KV capacity used) - B_t: Bandwidth pressure [0, 1] (memory traffic vs peak BW) Includes per-signal normalization, moving statistics, and individual signal ablation. """ import torch import torch.nn as nn import torch.nn.functional as F import math from typing import Dict, Optional, Tuple, List class TokenInformationState(nn.Module): """ Computes and maintains normalized token information states z_t. """ FEATURE_NAMES = [ "entanglement_dispersion", # S_t "predictive_entropy", # H_t "uncertainty", # U_t "attention_importance", # A_t "reconstruction_error", # R_t "latency_pressure", # L_t "memory_pressure", # M_t "bandwidth_pressure", # B_t ] def __init__(self, d_model: int, n_heads: int = 4, ema_decay: float = 0.95): super().__init__() self.d_model = d_model self.n_heads = n_heads self.ema_decay = ema_decay self.num_features = len(self.FEATURE_NAMES) # Learnable projection for uncertainty and entropy estimation self.uncertainty_probe = nn.Linear(d_model, 1, bias=True) self.entropy_probe = nn.Linear(d_model, 16, bias=True) # Running statistics for normalization (mean and var for each signal) self.register_buffer("running_mean", torch.full((self.num_features,), 0.5)) self.register_buffer("running_var", torch.full((self.num_features,), 0.08)) self.register_buffer("total_updates", torch.tensor(0, dtype=torch.long)) # Ablation mask (1.0 = active, 0.0 = ablated/zeroed) self.register_buffer("ablation_mask", torch.ones(self.num_features)) def set_ablation(self, feature_name: str, enabled: bool): """Enable or ablate an individual signal for ablation experiments.""" if feature_name not in self.FEATURE_NAMES: raise ValueError(f"Unknown feature: {feature_name}. Choose from {self.FEATURE_NAMES}") idx = self.FEATURE_NAMES.index(feature_name) self.ablation_mask[idx] = 1.0 if enabled else 0.0 def reset_ablation(self): """Reset all features to active.""" self.ablation_mask.fill_(1.0) def compute_attention_entropy(self, attn_weights: Optional[torch.Tensor], hidden_states: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]: """ Compute attention dispersion (S_t) and attention importance (A_t). """ if attn_weights is None and hidden_states is not None: # Outgoing attention dispersion proxy from hidden state dot-products scores = torch.matmul(hidden_states, hidden_states.transpose(-2, -1)) / math.sqrt(self.d_model) attn_weights = F.softmax(scores, dim=-1).unsqueeze(1) # (B, 1, T, T) if attn_weights is None: B, T = (hidden_states.shape[0], hidden_states.shape[1]) if hidden_states is not None else (1, 1) dev = hidden_states.device if hidden_states is not None else torch.device("cpu") return torch.full((B, T), 0.5, device=dev), torch.full((B, T), 0.5, device=dev) eps = 1e-9 # Shannon entropy of outgoing attention distribution: S_i = -sum_j A_ij log(A_ij) entropy = -torch.sum(attn_weights * torch.log(attn_weights + eps), dim=-1) # (B, H, T) S_t = entropy.mean(dim=1) # average across heads -> (B, T) # Normalize by maximum possible entropy log(T) T = attn_weights.size(-1) if T > 1: S_t = S_t / math.log(T) incoming = attn_weights.sum(dim=-2) # (B, H, T_keys) A_t = incoming.mean(dim=1) # (B, T) A_t = A_t / max(1.0, T / 4.0) return S_t, A_t def compute_predictive_entropy(self, logits: Optional[torch.Tensor], hidden_states: torch.Tensor, batch_size: int, seq_len: int, device: torch.device) -> torch.Tensor: """ Compute predictive entropy H_t from logits or representation probe. """ if logits is not None: eps = 1e-9 probs = F.softmax(logits, dim=-1) vocab_size = logits.size(-1) H_t = -torch.sum(probs * torch.log(probs + eps), dim=-1) return H_t / math.log(max(2, vocab_size)) else: probe_logits = self.entropy_probe(hidden_states) probs = F.softmax(probe_logits, dim=-1) H_t = -torch.sum(probs * torch.log(probs + 1e-9), dim=-1) return H_t / math.log(16) def compute_uncertainty(self, hidden_states: torch.Tensor) -> torch.Tensor: """ Compute epistemic uncertainty U_t from hidden state representation. Uses normalized variance across feature dimensions and probe confidence. """ # Feature dispersion: standard deviation across feature dimensions feat_std = torch.std(hidden_states, dim=-1) # (B, T) # Probe prediction probe_val = torch.sigmoid(self.uncertainty_probe(hidden_states)).squeeze(-1) # (B, T) # Combined uncertainty score U_t = torch.sigmoid(feat_std * probe_val) return U_t def forward( self, hidden_states: torch.Tensor, attn_weights: Optional[torch.Tensor] = None, logits: Optional[torch.Tensor] = None, reconstruction_error: Optional[torch.Tensor] = None, hardware_state: Optional[Dict[str, float]] = None, ) -> torch.Tensor: """ Construct the full normalized 8D information state tensor z_t. Args: hidden_states: (batch, seq_len, d_model) attn_weights: (batch, n_heads, seq_len, seq_len) or None logits: (batch, seq_len, vocab_size) or None reconstruction_error: (batch, seq_len) or float or None hardware_state: dict with keys ['latency_pressure', 'memory_pressure', 'bandwidth_pressure'] Returns: z_t: (batch, seq_len, 8) normalized information state vector """ B, T, D = hidden_states.shape device = hidden_states.device # 1. S_t & A_t (Attention dispersion & importance) S_t, A_t = self.compute_attention_entropy(attn_weights, hidden_states=hidden_states) # 2. H_t (Predictive entropy) H_t = self.compute_predictive_entropy(logits, hidden_states, B, T, device) # 3. U_t (Uncertainty) U_t = self.compute_uncertainty(hidden_states) # 4. R_t (Reconstruction error) if reconstruction_error is None: R_t = torch.full((B, T), 0.1, device=device) elif isinstance(reconstruction_error, (int, float)): R_t = torch.full((B, T), float(reconstruction_error), device=device) else: R_t = reconstruction_error.expand(B, T) if reconstruction_error.dim() < 2 else reconstruction_error # 5. Hardware pressure signals: L_t, M_t, B_t hw = hardware_state or {} L_val = float(hw.get("latency_pressure", 0.3)) M_val = float(hw.get("memory_pressure", 0.3)) B_val = float(hw.get("bandwidth_pressure", 0.3)) L_t = torch.full((B, T), L_val, device=device) M_t = torch.full((B, T), M_val, device=device) B_t = torch.full((B, T), B_val, device=device) # Stack into raw vector (B, T, 8) raw_z = torch.stack([S_t, H_t, U_t, A_t, R_t, L_t, M_t, B_t], dim=-1) # Update running statistics during training if self.training: with torch.no_grad(): batch_mean = raw_z.mean(dim=(0, 1)) batch_var = raw_z.var(dim=(0, 1), unbiased=False) + 1e-6 self.running_mean.mul_(self.ema_decay).add_(batch_mean, alpha=1.0 - self.ema_decay) self.running_var.mul_(self.ema_decay).add_(batch_var, alpha=1.0 - self.ema_decay) self.total_updates.add_(1) # Z-score normalization clamped to [-3, 3], then mapped to [0, 1] via sigmoid std = torch.sqrt(self.running_var + 1e-6).to(device) mean = self.running_mean.to(device) norm_z = (raw_z - mean) / std norm_z = torch.sigmoid(norm_z) # Apply ablation mask z_t = norm_z * self.ablation_mask.to(device) return z_t def get_signal_dict(self, z_t: torch.Tensor) -> Dict[str, float]: """Convert a batch mean z_t tensor into a readable dictionary.""" mean_vals = z_t.mean(dim=(0, 1)).detach().cpu().tolist() return {name: round(val, 4) for name, val in zip(self.FEATURE_NAMES, mean_vals)}