File size: 3,158 Bytes
320e2b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
74
import torch
import torch.nn as nn
import torch.nn.functional as F

class PCAw_Pool(nn.Module):
    def __init__(self, kernel_size, stride=(1, 1), eps: float = 1e-4, normalize_weights: bool = True):
        super().__init__()
        self.kernel_size = kernel_size
        self.stride = stride
        self.eps = eps
        self.normalize_weights = normalize_weights

        # D = number of features per patch = F_k * T_k
        F_k, T_k = kernel_size
        D = F_k * T_k

        # Trainable weights over PCA components (columns). Shape: (D,)
        # Initialized small to avoid overpowering early training.
        self.weights = nn.Parameter(0.01 * torch.randn(D))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        B, C, Fh, Tw = x.shape
        F_k, T_k = self.kernel_size

        # Unfold per-channel
        x_bc = x.view(B * C, 1, Fh, Tw)  # (B*C, 1, F, T)
        patches = F.unfold(x_bc, kernel_size=(F_k, T_k), stride=self.stride)  # (B*C, D, NumPatches)
        D = patches.shape[1]
        NumPatches = patches.shape[2]

        # Arrange to (B, C, NumPatches, D) -> (B, C*NumPatches, D)
        patches = patches.permute(0, 2, 1).contiguous()         # (B*C, NumPatches, D)
        patches = patches.view(B, C, NumPatches, D)             # (B, C, NumPatches, D)
        X = patches.view(B, C * NumPatches, D)                  # (B, N, D) with N = C*NumPatches
        N = X.shape[1]

        # Output spatial size
        H = (Fh - F_k) // self.stride[0] + 1
        W = (Tw - T_k) // self.stride[1] + 1

        # Center features
        mean = X.mean(dim=1, keepdim=True)                     # (B, 1, D)
        Xc = X - mean                                          # (B, N, D)

        # Optional: standardize per-feature to tame scale explosions from DSConv
        var = Xc.pow(2).mean(dim=1, keepdim=True)              # (B, 1, D)
        Xc = Xc / torch.sqrt(var + 1e-6)

        # --------- Stable projection basis via SVD (more robust than eigh) ---------
        # Xc = U S V^T  -> principal components are columns of V
        # Use full_matrices=False for efficiency and stable backward
        U, S, Vh = torch.linalg.svd(Xc, full_matrices=False)   # Vh: (B, D, D)
        eigvecs = Vh.transpose(1, 2)                           # (B, D, D)

        # Detach eigenvectors to avoid backprop through SVD
        eigvecs = eigvecs.detach()

        # Weighted projection direction v = E @ w
        if self.normalize_weights:
            w = torch.softmax(self.weights, dim=0)             # (D,)
        else:
            w = self.weights

        v = torch.matmul(eigvecs, w)                           # (B, D)
        v = v / (v.norm(dim=1, keepdim=True) + 1e-8)           # normalize

        # Project samples onto v -> scalar per sample
        scores = torch.matmul(Xc, v.unsqueeze(-1)).squeeze(-1) # (B, N)

        return scores.view(B, C, H, W)

    def extra_repr(self) -> str:
        return (f"kernel_size={self.kernel_size}, stride={self.stride}, "
                f"eps={self.eps}, normalize_weights={self.normalize_weights}")