| 23 | # 基础工具:ADAIN 所需的统计量(保留以备需要;管线默认用 wavelet) |
| 24 | # ----------------------------- |
| 25 | def _calc_mean_std(feat: torch.Tensor, eps: float = 1e-5) -> Tuple[torch.Tensor, torch.Tensor]: |
| 26 | assert feat.dim() == 4, 'feat 必须是 (N, C, H, W)' |
| 27 | N, C = feat.shape[:2] |
| 28 | var = feat.view(N, C, -1).var(dim=2, unbiased=False) + eps |
| 29 | std = var.sqrt().view(N, C, 1, 1) |
| 30 | mean = feat.view(N, C, -1).mean(dim=2).view(N, C, 1, 1) |
| 31 | return mean, std |
| 32 | |
| 33 | |
| 34 | def _adain(content_feat: torch.Tensor, style_feat: torch.Tensor) -> torch.Tensor: |