| import torch |
| import torch.nn as nn |
|
|
|
|
| class UncertaintyHead(nn.Module): |
| """ |
| Per-token retrieval-uncertainty head. |
| |
| Predicts a scalar log-variance ``um`` for every latent token from the |
| model's own x0 prediction (so ``um`` reflects "how unsure the model is |
| about the current denoised prediction"). Used by |
| :func:`retrieval_uncertainty_loss` to form the heteroscedastic loss |
| |
| L_unc = mean( exp(-um) * sg(MSE(x0_target, x0_pred)) + um ) |
| |
| Shape contract: |
| forward(x0_pred) where x0_pred: (B, C, T, H, W) -> um: (B, 1, T, H, W) |
| |
| The final conv is zero-initialised so ``um == 0`` at init (=> exp(-um)==1, |
| a no-op weighting), which keeps the existing training dynamics intact on |
| step 0 and lets the head warm up gracefully. |
| """ |
|
|
| def __init__(self, in_channels: int, hidden: int = 64, |
| um_min: float = -10.0, um_max: float = 10.0): |
| super().__init__() |
| self.in_channels = int(in_channels) |
| self.um_min = float(um_min) |
| self.um_max = float(um_max) |
| self.net = nn.Sequential( |
| nn.Conv3d(self.in_channels, hidden, kernel_size=1), |
| nn.SiLU(), |
| nn.Conv3d(hidden, 1, kernel_size=1), |
| ) |
| nn.init.zeros_(self.net[-1].weight) |
| nn.init.zeros_(self.net[-1].bias) |
|
|
| def forward(self, x0_pred: torch.Tensor) -> torch.Tensor: |
| um = self.net(x0_pred.float()) |
| |
| return um.clamp(self.um_min, self.um_max) |
|
|
|
|
| def recover_x0_flow_match(noisy_target: torch.Tensor, |
| v_pred: torch.Tensor, |
| sigma: torch.Tensor) -> torch.Tensor: |
| """ |
| Flow-matching x0 recovery. |
| |
| The FlowMatchScheduler uses x_t = (1 - sigma) * x0 + sigma * noise and |
| trains the model to predict the velocity v = noise - x0. Therefore |
| |
| x0 = x_t - sigma * v. |
| |
| Args: |
| noisy_target: x_t for the TARGET tokens, (B, C, T, H, W). |
| v_pred: model velocity prediction for the TARGET tokens, same shape. |
| sigma: scalar tensor (the sigma matching the sampled timestep). |
| """ |
| return noisy_target - sigma * v_pred |
|
|
|
|
| def per_token_mse_map(x0_pred: torch.Tensor, |
| x0_target: torch.Tensor) -> torch.Tensor: |
| """Mean-squared error averaged over the channel axis -> (B, 1, T, H, W).""" |
| return (x0_pred.float() - x0_target.float()).pow(2).mean(dim=1, keepdim=True) |
|
|
|
|
| def retrieval_uncertainty_loss(um: torch.Tensor, |
| mse_map: torch.Tensor, |
| detach_mse: bool = True) -> torch.Tensor: |
| """ |
| Heteroscedastic retrieval-uncertainty loss: |
| |
| e^(-um) * sg(MSE) + um (averaged over all tokens) |
| |
| With ``detach_mse=True`` (the ``sg`` in the spec) only the uncertainty head |
| is trained by this term: it learns to predict where the x0 reconstruction |
| is wrong, giving a per-token "retrieval confidence" signal. The denoiser is |
| untouched by this term. To additionally let the signal shape the memory |
| pathway, reweight the MAIN denoising loss with ``exp(-um.detach())`` (see |
| retrieve.md, Method A variant). |
| """ |
| if detach_mse: |
| mse_map = mse_map.detach() |
| return (torch.exp(-um) * mse_map + um).mean() |
|
|