"""Auto-extracted from detection_arena.py.""" import math import torch import torch.nn as nn import torch.nn.functional as F from torch import Tensor from typing import List from losses.fcos import fcos_loss, focal_loss, NUM_CLASSES from losses.centernet import centernet_targets, centernet_loss from utils.decode import make_locations, decode_fcos, decode_centernet, FPN_STRIDES N_PREFIX = 5 class CofiberLinear(nn.Module): """Adjoint scale decomposition. Cofibers isolate per-scale content. ~65K params.""" name = "Q_adjoint_scale" needs_intermediates = False def __init__(self, feat_dim=768, num_classes=80, n_scales=3): super().__init__() self.n_scales = n_scales self.cls_head = nn.Conv2d(feat_dim, num_classes, 1) self.reg_head = nn.Conv2d(feat_dim, 4, 1) self.ctr_head = nn.Conv2d(feat_dim, 1, 1) self.scale_params = nn.Parameter(torch.ones(n_scales)) nn.init.constant_(self.cls_head.bias, -math.log(99)) @staticmethod def _cofiber_decompose(f, n_scales): """Compute cofibers via the adjoint pair (Σ ⊣ Ω). Σ = bilinear upsample 2x, Ω = avg pool 2x. cofiber_k = f_k - Σ(Ω(f_k)): information at scale k absent from scale k+1. The decomposition is exact: sum of cofibers + residual = original.""" 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) cofiber = residual - sigma_omega cofibers.append(cofiber) residual = omega cofibers.append(residual) return cofibers def forward(self, spatial, inter=None): cofibers = self._cofiber_decompose(spatial, self.n_scales) cls_l, reg_l, ctr_l = [], [], [] for i, cof in enumerate(cofibers): cls_l.append(self.cls_head(cof)) raw = (self.reg_head(cof) * self.scale_params[i]).clamp(-10, 10) reg_l.append(torch.exp(raw)) ctr_l.append(self.ctr_head(cof)) return cls_l, reg_l, ctr_l def loss(self, preds, locs, boxes_b, labels_b): return fcos_loss(*preds, locs, boxes_b, labels_b) def decode(self, preds, locs, **kw): return decode_fcos(*preds, locs, **kw) def get_locs(self, spatial): cofibers = self._cofiber_decompose(spatial[:1], self.n_scales) sizes = [(c.shape[2], c.shape[3]) for c in cofibers] strides = [16 * (2 ** i) for i in range(self.n_scales)] return make_locations(sizes, strides, spatial.device)