File size: 4,066 Bytes
a19671a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
83
84
85
86
87
88
89
90
91
92
93
94
95
"""
Sub-Center ArcFace with K Sub-Centers per Class.
Designed for fine-grained botanical vision to capture multi-modal phenology
(e.g., Flower vs. Leaf/Foliage vs. Fruit/Seed Pods) without intra-class noise.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F

class SubCenterArcFace(nn.Module):
    """
    Sub-Center ArcFace for 512-D L2 Embeddings.
    Maintains K sub-centers per class (default K=3).
    """
    def __init__(self, embed_dim: int, n_classes: int, k_subcenters: int = 3, m: float = 0.5, s: float = 30.0, easy_margin: bool = False):
        super().__init__()
        self.embed_dim = embed_dim
        self.n_classes = n_classes
        self.k_subcenters = k_subcenters
        self.m = m
        self.s = s
        self.easy_margin = easy_margin

        # Weight matrix: [n_classes * k_subcenters, embed_dim]
        self.W = nn.Parameter(torch.randn(n_classes * k_subcenters, embed_dim) * 0.01)
        nn.init.xavier_uniform_(self.W)

        # Precomputed angular constants
        self.cos_m = math.cos(m)
        self.sin_m = math.sin(m)
        self.th = math.cos(math.pi - m)
        self.mm = math.sin(math.pi - m) * m

    def forward(self, emb: torch.Tensor, y: torch.Tensor = None) -> torch.Tensor:
        """
        emb: [B, D] L2 normalized embedding
        y: [B] Ground-truth class labels (0 .. n_classes-1)
        Returns: [B, n_classes] Logits scaled by s
        """
        # Normalize weights and embeddings
        W_norm = F.normalize(self.W, dim=1) # [C * K, D]
        emb_norm = F.normalize(emb, dim=1)   # [B, D]

        # Compute cosine similarity across all sub-centers: [B, C * K]
        cos_all = F.linear(emb_norm, W_norm)
        B = emb.size(0)

        # Reshape to [B, C, K] and select the maximum cosine per class
        cos_sub = cos_all.view(B, self.n_classes, self.k_subcenters)
        cos_max, _ = cos_sub.max(dim=-1) # [B, C]
        cos = cos_max.clamp(-1.0 + 1e-7, 1.0 - 1e-7)

        # Inference mode: return scaled max cosine logits
        if y is None:
            return cos * self.s

        # Training mode: Apply additive angular margin to ground-truth class
        sin = torch.sqrt((1.0 - cos ** 2).clamp(0.0, 1.0))
        cos_m = cos * self.cos_m - sin * self.sin_m

        if self.easy_margin:
            cos_m = torch.where(cos > 0, cos_m, cos)
        else:
            cos_m = torch.where(cos > self.th, cos_m, cos - self.mm)

        one_hot = F.one_hot(y, num_classes=self.n_classes).bool()
        logits = torch.where(one_hot, cos_m, cos)
        return logits * self.s

    def diversity_loss(self) -> torch.Tensor:
        """

        Penalise sub-center collapse (only winning K=1 gets gradients via max-pool).

        Computes mean absolute pairwise cosine between sub-centers within each class;

        encourages K=3 centers to stay diverse (e.g., flower/leaf/fruit).

        """
        if self.k_subcenters <= 1:
            return torch.tensor(0.0, device=self.W.device, dtype=self.W.dtype)
        # Reshape to [C, K, D] and L2-normalize per sub-center
        W_ck = F.normalize(self.W.view(self.n_classes, self.k_subcenters, self.embed_dim), dim=-1)
        # Pairwise cosine within each class: [C, K, K]
        cos = torch.einsum("ckd,cjd->ckj", W_ck, W_ck)
        # Mask diagonal (self-similarity =1)
        mask = torch.eye(self.k_subcenters, device=cos.device, dtype=torch.bool).unsqueeze(0).expand(
            self.n_classes, -1, -1
        )
        cos_off = cos[~mask]
        if cos_off.numel() == 0:
            return torch.tensor(0.0, device=cos.device, dtype=cos.dtype)
        # Want low similarity → penalise absolute correlation; hinge variant: relu(cos - 0.2)
        # Use abs mean for symmetry (both positive collapse and opposite collapse are diverse enough?)
        # but positive collapse is the failure mode, so use mean of positive part.
        # Here we use abs mean scaled: keeps gradient if centers align or anti-align strongly.
        return cos_off.abs().mean()