SACFlow / sacflow /methods /source_memory.py
sathiiii's picture
Add SACFlow source code
d65ae7d verified
Raw
History Blame Contribute Delete
2.84 kB
from __future__ import annotations
from pathlib import Path
from typing import Dict
import torch
import torch.nn.functional as F
@torch.no_grad()
def class_moments(residual: torch.Tensor, probs: torch.Tensor, eps: float = 1e-6):
"""Class-conditional channel moments without materializing [B,C,d,N].
residual: [B,d,H,W,D]
probs: [B,C,H,W,D], soft or one-hot class weights at the same resolution.
Returns:
mu: [C,d]
std: [C,d]
weights: [C]
This implementation uses first and second moments directly and is much more
memory efficient than computing (residual - mu)^2 for every class/voxel.
"""
B, d, H, W, D = residual.shape
C = probs.shape[1]
r = residual.reshape(B, d, -1)
p = probs.reshape(B, C, -1).to(dtype=residual.dtype)
weights = p.sum(dim=(0, 2)).clamp_min(eps) # [C]
mu = torch.einsum("bcn,bdn->cd", p, r) / weights[:, None]
second = torch.einsum("bcn,bdn->cd", p, r.pow(2)) / weights[:, None]
var = (second - mu.pow(2)).clamp_min(eps)
std = torch.sqrt(var)
return mu, std, weights
def hard_onehot(label: torch.Tensor, num_classes: int):
return F.one_hot(label.long().clamp(0, num_classes-1), num_classes).permute(0,4,1,2,3).float()
def save_source_memory(path: str | Path, memory: Dict):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(memory, path)
def load_source_memory(path: str | Path, device=None):
mem = torch.load(path, map_location=device or "cpu")
return mem
def moment_transport_residual(Rt: torch.Tensor, probs: torch.Tensor, source_mu: torch.Tensor, source_std: torch.Tensor,
target_mu: torch.Tensor | None = None, target_std: torch.Tensor | None = None,
eps: float = 1e-5):
"""Class-gated diagonal moment transport from target residual stats to source residual stats.
Computes R0(u) = sum_c p_c(u) [mu_s_c + sigma_s_c / sigma_t_c * (Rt(u)-mu_t_c)].
The implementation intentionally avoids stacking all class-wise transported
residuals, because [B,C,d,H,W,D] can be very large for 3D volumes.
"""
B, d, H, W, D = Rt.shape
C = probs.shape[1]
if target_mu is None or target_std is None:
target_mu, target_std, _ = class_moments(Rt, probs, eps=eps)
source_mu = source_mu.to(Rt.device, Rt.dtype)
source_std = source_std.to(Rt.device, Rt.dtype)
target_mu = target_mu.to(Rt.device, Rt.dtype)
target_std = target_std.to(Rt.device, Rt.dtype)
R0 = torch.zeros_like(Rt)
for c in range(C):
tr = source_mu[c].view(1,d,1,1,1) + (source_std[c].view(1,d,1,1,1) / (target_std[c].view(1,d,1,1,1) + eps)) * (Rt - target_mu[c].view(1,d,1,1,1))
R0 = R0 + probs[:, c:c+1].to(Rt.dtype) * tr
return R0, target_mu, target_std