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}")