Download sacflow/methods/task_space.py from sathiiii/SACFlow: direct link, hf CLI and curl.
- Browser
- Download file 1.29 kB
-
https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/methods/task_space.py
- Command line
-
hf download hf://sathiiii/SACFlow/sacflow/methods/task_space.py
-
curl -L -o task_space.py https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/methods/task_space.py
1.29 kB
| from __future__ import annotations | |
| import torch | |
| import torch.nn.functional as F | |
| def centered_classifier_basis(W: torch.Tensor, tol: float = 1e-5, foreground_only: bool = False): | |
| # W: [C,d] | |
| if foreground_only and W.shape[0] > 1: | |
| W = W[1:] | |
| Wc = W - W.mean(dim=0, keepdim=True) | |
| # U basis in feature dimension for row space of Wc | |
| U, S, Vh = torch.linalg.svd(Wc.T, full_matrices=False) | |
| if S.numel() == 0: | |
| return torch.empty(W.shape[1], 0, device=W.device, dtype=W.dtype) | |
| r = int((S > tol * S.max()).sum().item()) | |
| return U[:, :max(1, r)] | |
| def random_basis(feature_dim: int, rank: int, device, dtype): | |
| A = torch.randn(feature_dim, rank, device=device, dtype=dtype) | |
| Q, _ = torch.linalg.qr(A, mode="reduced") | |
| return Q | |
| def project_task_and_residual(Fmap: torch.Tensor, Q: torch.Tensor): | |
| # Fmap [B,d,H,W,D], Q [d,r] | |
| if Q.numel() == 0: | |
| return torch.zeros_like(Fmap), Fmap | |
| B, d, H, W, D = Fmap.shape | |
| flat = Fmap.permute(0,2,3,4,1).reshape(-1, d) | |
| task = (flat @ Q) @ Q.T | |
| task = task.reshape(B,H,W,D,d).permute(0,4,1,2,3).contiguous() | |
| residual = Fmap - task | |
| return task, residual | |
| def project_to_residual(x: torch.Tensor, Q: torch.Tensor): | |
| task, res = project_task_and_residual(x, Q) | |
| return res | |