Q-TensorFormer / src /information_state.py
Premchandyadav369
feat(research): Turn Q-TensorFormer into a verified research system with LFS figure assets
0e4850b
Raw History Blame Contribute Delete
9.16 kB
"""
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)}