"""Confidence head for estimating frozen Predictor hidden-space error.""" from __future__ import annotations import torch from torch import nn class PredictorConfidenceHead(nn.Module): """Predict log hidden nRMSE from the Predictor block output. The learned head estimates local approximation error. Chunk-dependent downstream amplification stays explicit through :func:`impact_risk`, so a beta sweep does not require retraining the head. """ def __init__( self, dim: int = 1536, token_hidden_dim: int = 512, token_output_dim: int = 256, context_dim: int = 64, dropout: float = 0.1, num_steps: int = 2, ) -> None: super().__init__() if num_steps < 1: raise ValueError("num_steps must be positive") self.num_steps = int(num_steps) self.norm = nn.LayerNorm(dim, eps=1e-6) self.token_encoder = nn.Sequential( nn.Linear(dim, token_hidden_dim), nn.SiLU(), nn.Linear(token_hidden_dim, token_output_dim), nn.SiLU(), ) self.token_score = nn.Linear(token_output_dim, 1) self.chunk_mlp = nn.Sequential( nn.Linear(1, context_dim), nn.SiLU(), nn.Linear(context_dim, context_dim), ) self.step_embedding = nn.Embedding(self.num_steps, context_dim) output_input_dim = 2 * token_output_dim + 4 + 2 * context_dim self.output = nn.Sequential( nn.Linear(output_input_dim, token_output_dim), nn.SiLU(), nn.Dropout(dropout), nn.Linear(token_output_dim, context_dim), nn.SiLU(), nn.Linear(context_dim, 1), ) @staticmethod def _rms(value: torch.Tensor) -> torch.Tensor: return value.float().square().mean(dim=(1, 2)).sqrt() def forward( self, *, transformed_hidden: torch.Tensor, pred_hidden: torch.Tensor, anchor_hidden: torch.Tensor, chunk_position: torch.Tensor, step_id: torch.Tensor, ) -> torch.Tensor: if transformed_hidden.ndim != 3: raise ValueError("transformed_hidden must have shape [B, tokens, dim]") if pred_hidden.shape != transformed_hidden.shape: raise ValueError("pred_hidden and transformed_hidden shapes differ") if anchor_hidden.shape != transformed_hidden.shape: raise ValueError("anchor_hidden and transformed_hidden shapes differ") if torch.any((step_id < 1) | (step_id > self.num_steps)): raise ValueError( f"Confidence head supports step_id 1..{self.num_steps}" ) token_feature = self.token_encoder(self.norm(transformed_hidden.detach())) mean_feature = token_feature.mean(dim=1) attention = torch.softmax( self.token_score(token_feature).float(), dim=1 ).to(dtype=token_feature.dtype) attention_feature = (attention * token_feature).sum(dim=1) pred_detached = pred_hidden.detach() anchor_detached = anchor_hidden.detach() residual = pred_detached - anchor_detached scalar_stats = torch.stack( [ self._rms(residual), residual.float().abs().mean(dim=(1, 2)), self._rms(pred_detached), self._rms(anchor_detached), ], dim=-1, ).to(dtype=token_feature.dtype) chunk_feature = self.chunk_mlp( chunk_position.to(dtype=token_feature.dtype).unsqueeze(-1) ) step_feature = self.step_embedding(step_id - 1) combined = torch.cat( [ mean_feature, attention_feature, scalar_stats, chunk_feature, step_feature, ], dim=-1, ) return self.output(combined).squeeze(-1).float() class ConfidenceTokenHead(nn.Module): """Predict log hidden nRMSE with a query-only Confidence token. Predictor tokens are projected to 512 dimensions once. A learned query, conditioned on chunk, denoising step, and scalar hidden statistics, attends to all projected tokens. Keeping the Confidence token as the only query avoids the quadratic memory cost of full token self-attention. """ def __init__( self, dim: int = 1536, token_dim: int = 512, context_dim: int = 64, num_heads: int = 8, ffn_dim: int = 2048, dropout: float = 0.1, num_steps: int = 3, ) -> None: super().__init__() if num_steps < 1: raise ValueError("num_steps must be positive") if token_dim % num_heads != 0: raise ValueError("token_dim must be divisible by num_heads") self.num_steps = int(num_steps) self.token_dim = int(token_dim) self.norm = nn.LayerNorm(dim, eps=1e-6) self.token_projection = nn.Sequential( nn.Linear(dim, token_dim), nn.SiLU(), ) self.chunk_mlp = nn.Sequential( nn.Linear(1, context_dim), nn.SiLU(), nn.Linear(context_dim, context_dim), ) self.step_embedding = nn.Embedding(self.num_steps, context_dim) context_input_dim = 2 * context_dim + 4 self.query_context = nn.Linear(context_input_dim, token_dim) self.confidence_token = nn.Parameter(torch.zeros(1, 1, token_dim)) nn.init.normal_(self.confidence_token, std=0.02) self.query_norm = nn.LayerNorm(token_dim, eps=1e-6) self.key_value_norm = nn.LayerNorm(token_dim, eps=1e-6) self.cross_attention = nn.MultiheadAttention( embed_dim=token_dim, num_heads=num_heads, dropout=0.0, batch_first=True, ) self.confidence_norm = nn.LayerNorm(token_dim, eps=1e-6) self.confidence_ffn = nn.Sequential( nn.Linear(token_dim, ffn_dim), nn.SiLU(), nn.Dropout(dropout), nn.Linear(ffn_dim, token_dim), ) output_input_dim = 2 * token_dim + 2 * context_dim + 4 self.output = nn.Sequential( nn.Linear(output_input_dim, token_dim), nn.SiLU(), nn.Dropout(dropout), nn.Linear(token_dim, token_dim), nn.SiLU(), nn.Linear(token_dim, 1), ) @staticmethod def _rms(value: torch.Tensor) -> torch.Tensor: return value.float().square().mean(dim=(1, 2)).sqrt() def forward( self, *, transformed_hidden: torch.Tensor, pred_hidden: torch.Tensor, anchor_hidden: torch.Tensor, chunk_position: torch.Tensor, step_id: torch.Tensor, ) -> torch.Tensor: if transformed_hidden.ndim != 3: raise ValueError("transformed_hidden must have shape [B, tokens, dim]") if pred_hidden.shape != transformed_hidden.shape: raise ValueError("pred_hidden and transformed_hidden shapes differ") if anchor_hidden.shape != transformed_hidden.shape: raise ValueError("anchor_hidden and transformed_hidden shapes differ") if chunk_position.shape != (transformed_hidden.shape[0],): raise ValueError("chunk_position must have shape [B]") if step_id.shape != (transformed_hidden.shape[0],): raise ValueError("step_id must have shape [B]") if torch.any((step_id < 1) | (step_id > self.num_steps)): raise ValueError( f"Confidence head supports step_id 1..{self.num_steps}" ) token_feature = self.token_projection( self.norm(transformed_hidden.detach()) ) mean_feature = token_feature.mean(dim=1) pred_detached = pred_hidden.detach() anchor_detached = anchor_hidden.detach() residual = pred_detached - anchor_detached scalar_stats = torch.stack( [ self._rms(residual), residual.float().abs().mean(dim=(1, 2)), self._rms(pred_detached), self._rms(anchor_detached), ], dim=-1, ).to(dtype=token_feature.dtype) chunk_feature = self.chunk_mlp( chunk_position.to(dtype=token_feature.dtype).unsqueeze(-1) ).to(dtype=token_feature.dtype) step_feature = self.step_embedding(step_id - 1).to( dtype=token_feature.dtype ) query_context = torch.cat( [scalar_stats, chunk_feature, step_feature], dim=-1 ) query = self.confidence_token.to(dtype=token_feature.dtype).expand( token_feature.shape[0], -1, -1 ) query = query + self.query_context(query_context).unsqueeze(1) key_value = self.key_value_norm(token_feature) attended, _ = self.cross_attention( self.query_norm(query), key_value, key_value, need_weights=False, ) confidence_feature = query + attended confidence_feature = confidence_feature + self.confidence_ffn( self.confidence_norm(confidence_feature) ) confidence_feature = confidence_feature.squeeze(1) combined = torch.cat( [ mean_feature, confidence_feature, scalar_stats, chunk_feature, step_feature, ], dim=-1, ) return self.output(combined).squeeze(-1).float() def eligible_chunk_alpha( chunk_idx: torch.Tensor, *, first_predictor_chunk: int = 1, last_chunk: int = 6, ) -> torch.Tensor: """Linear propagation prior normalized over Predictor-eligible chunks.""" if last_chunk <= first_predictor_chunk: return torch.zeros_like(chunk_idx, dtype=torch.float32) clamped = chunk_idx.float().clamp(first_predictor_chunk, last_chunk) return (last_chunk - clamped) / (last_chunk - first_predictor_chunk) def impact_risk( predicted_log_error: torch.Tensor, chunk_idx: torch.Tensor, beta: float, ) -> torch.Tensor: if beta < 0: raise ValueError("beta must be non-negative") alpha = eligible_chunk_alpha(chunk_idx).to(predicted_log_error.device) return predicted_log_error.exp() * (1.0 + float(beta) * alpha)