File size: 2,839 Bytes
d65ae7d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
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