File size: 3,501 Bytes
dbbceb8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
75
76
77
78
79
80
81
82
"""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)