Cccccz's picture
Upload predictor training source
20962c9 verified
Raw History Blame Contribute Delete
10.5 kB
"""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)