echo / code /diffsynth /models /memory /uncertainty.py
amonshano's picture
Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team)
00c7b31 verified
Raw
History Blame Contribute Delete
3.32 kB
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())
# Clamp keeps exp(-um) finite under bf16/fp16 autocast.
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()