"""Cofiber Threshold with dimension selection: 768→20→80 classification. The bottleneck dimension K=20 was selected from SVD analysis of the pruned prototype matrix, where rank 20 captures 72% of the energy. This is the information bottleneck variant applied to detection: how few feature dimensions does the backbone need to expose for 80-class detection? ~20K total params. """ import math import torch import torch.nn as nn import torch.nn.functional as F def cofiber_decompose(f, n_scales): cofibers = [] residual = f for _ in range(n_scales - 1): omega = F.avg_pool2d(residual, 2) sigma_omega = F.interpolate(omega, size=residual.shape[2:], mode="bilinear", align_corners=False) cofibers.append(residual - sigma_omega) residual = omega cofibers.append(residual) return cofibers class CofiberThresholdDim20(nn.Module): """Cofiber decomposition + 768→20 projection + 20→80 classification. ~20K params.""" name = "cofiber_threshold_dim20" needs_intermediates = False def __init__(self, feat_dim=768, bottleneck_dim=20, num_classes=80, n_scales=3, reg_hidden=16): super().__init__() self.n_scales = n_scales self.scale_norms = nn.ModuleList([nn.LayerNorm(feat_dim) for _ in range(n_scales)]) # Bottleneck projection self.project = nn.Linear(feat_dim, bottleneck_dim, bias=False) # Classification from bottleneck self.cls_weight = nn.Parameter(torch.randn(num_classes, bottleneck_dim) * 0.01) self.cls_bias = nn.Parameter(torch.zeros(num_classes)) # Box regression from bottleneck (small hidden layer) self.reg_hidden = nn.Linear(bottleneck_dim, reg_hidden) self.reg_act = nn.GELU() self.reg_out = nn.Linear(reg_hidden, 4) # Centerness from bottleneck self.ctr_weight = nn.Parameter(torch.randn(1, bottleneck_dim) * 0.01) self.ctr_bias = nn.Parameter(torch.zeros(1)) self.scale_params = nn.Parameter(torch.ones(n_scales)) def forward(self, spatial, inter=None): cofibers = cofiber_decompose(spatial, self.n_scales) cls_l, reg_l, ctr_l = [], [], [] for i, cof in enumerate(cofibers): B, C, H, W = cof.shape f = self.scale_norms[i](cof.permute(0, 2, 3, 1).reshape(-1, C)) z = self.project(f) # (N, 20) cls = (z @ self.cls_weight.T + self.cls_bias).reshape(B, H, W, -1).permute(0, 3, 1, 2) reg_raw = (self.reg_out(self.reg_act(self.reg_hidden(z))) * self.scale_params[i]).clamp(-10, 10) reg = torch.exp(reg_raw).reshape(B, H, W, 4).permute(0, 3, 1, 2) ctr = (z @ self.ctr_weight.T + self.ctr_bias).reshape(B, H, W, 1).permute(0, 3, 1, 2) cls_l.append(cls) reg_l.append(reg) ctr_l.append(ctr) return cls_l, reg_l, ctr_l def loss(self, preds, locs, boxes_b, labels_b): from losses.fcos import fcos_loss return fcos_loss(*preds, locs, boxes_b, labels_b) def decode(self, preds, locs, **kw): from utils.decode import decode_fcos return decode_fcos(*preds, locs, **kw) def get_locs(self, spatial): from utils.decode import make_locations dummy = cofiber_decompose(spatial[:1], self.n_scales) sizes = [(c.shape[2], c.shape[3]) for c in dummy] strides = [16 * (2 ** i) for i in range(self.n_scales)] return make_locations(sizes, strides, spatial.device)