(x: torch.Tensor, levels: int = 5)
| 64 | |
| 65 | |
| 66 | def _wavelet_decompose(x: torch.Tensor, levels: int = 5) -> Tuple[torch.Tensor, torch.Tensor]: |
| 67 | assert x.dim() == 4, 'x 必须是 (N, C, H, W)' |
| 68 | high = torch.zeros_like(x) |
| 69 | low = x |
| 70 | for i in range(levels): |
| 71 | radius = 2 ** i |
| 72 | blurred = _wavelet_blur(low, radius) |
| 73 | high = high + (low - blurred) |
| 74 | low = blurred |
| 75 | return high, low |
| 76 | |
| 77 | |
| 78 | def _wavelet_reconstruct(content: torch.Tensor, style: torch.Tensor, levels: int = 5) -> torch.Tensor: |
no test coverage detected