Download predictor_training/confidence.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/predictor_training/confidence.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/predictor_training/confidence.py
-
curl -L -o confidence.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/predictor_training/confidence.py
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), | |
| ) | |
| 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), | |
| ) | |
| 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) | |