"""SC³: Sub-Cell Center Classification for precise bbox localization. Core idea from three cross-domain works: - SimCC (ECCV 2022, pose estimation): coordinate regression → coordinate classification - FIP-GDE (CVPR 2025, tiny obj detection): size-adaptive Gaussian position distribution - D-FINE (arXiv 2024, DETR detection): iterative distribution refinement For each anchor point, predicts a k×k distribution over sub-cell center positions. Replaces center derivation from (l,t,r,b) edges with explicit classification. Classification is easier to optimize than regression for fine-grained values. Soft-argmax gives sub-bin precision. At k=8, stride=8 (P3): each sub-cell = 1 pixel → 1px center precision. """ import torch import torch.nn as nn import torch.nn.functional as F import numpy as np class SC3Head(nn.Module): """Lightweight auxiliary head for sub-cell center classification. Takes FPN features (P3, P4, P5) and predicts a k×k distribution over sub-cell center positions for each anchor. Total added params < 5%. """ def __init__(self, in_channels, k=8): """ Args: in_channels: list of ints, input channels per FPN level [C3, C4, C5] k: sub-cell grid size (k×k bins per anchor cell) """ super().__init__() self.k = k self.nl = len(in_channels) # Per-level: lightweight conv to predict k×k logits self.convs = nn.ModuleList() for c_in in in_channels: mid = max(c_in // 4, 32) self.convs.append(nn.Sequential( nn.Conv2d(c_in, mid, 3, padding=1), nn.BatchNorm2d(mid), nn.SiLU(inplace=True), nn.Conv2d(mid, k * k, 1), )) self._init_weights() def _init_weights(self): for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') if m.bias is not None: # Initialize bias so initial distribution is centered (uniform-ish) nn.init.constant_(m.bias, 0.0) def forward(self, features): """features: list of [B, C_i, H_i, W_i] FPN features (P3, P4, P5) Returns: list of [B, k*k, H_i, W_i] logits per level """ return [conv(f) for conv, f in zip(self.convs, features)] def decode_sc3_center(sc3_logits, anchors, strides, k): """Decode SC³ logits to refined center coordinates. Args: sc3_logits: list of [B, k*k, H, W] per level anchors: list of [N_i, 2] anchor grid points per level strides: list of ints, stride per level k: sub-cell grid size Returns: centers: [N_total, 2] refined center (x, y) in image coords """ centers = [] for i, logits in enumerate(sc3_logits): B, _, H, W = logits.shape # Reshape: [B, k, k, H, W] → softmax over k×k dims logits_2d = logits.view(B, k, k, H, W) prob = F.softmax(logits_2d.reshape(B, k * k, H, W), dim=1) prob = prob.view(B, k, k, H, W) # Create sub-cell coordinate grid gy, gx = torch.meshgrid( torch.arange(k, device=logits.device, dtype=torch.float32), torch.arange(k, device=logits.device, dtype=torch.float32), indexing='ij' ) # Normalize to [0, 1) within the cell gx_norm = (gx + 0.5) / k gy_norm = (gy + 0.5) / k # Soft-argmax: expected sub-cell offset dx = (prob * gx_norm[None, :, :, None, None]).sum(dim=(1, 2)) # [B, H, W] dy = (prob * gy_norm[None, :, :, None, None]).sum(dim=(1, 2)) # [B, H, W] # Add anchor cell position anchor = anchors[i] # [H*W, 2] in pixel coords ax = anchor[:, 0].view(H, W) ay = anchor[:, 1].view(H, W) # Center in image coords cx = ax[None] + dx * strides[i] cy = ay[None] + dy * strides[i] # Stack as [B, H*W, 2] centers.append(torch.stack([cx.reshape(B, -1), cy.reshape(B, -1)], dim=-1)) return torch.cat(centers, dim=1) # [B, N_total, 2] def make_gaussian_target(cx_sub, cy_sub, k, sigma): """Create 2D Gaussian classification target on k×k grid. Args: cx_sub, cy_sub: GT center in sub-cell coords [0, k), [N] tensors k: grid size sigma: Gaussian width per sample, [N] tensor Returns: target: [N, k, k] normalized Gaussian distributions """ gy, gx = torch.meshgrid( torch.arange(k, device=cx_sub.device, dtype=torch.float32), torch.arange(k, device=cx_sub.device, dtype=torch.float32), indexing='ij' ) # Grid centers at (i + 0.5) gx_ctr = gx + 0.5 gy_ctr = gy + 0.5 # Distance from GT to each grid point: [N, k, k] dx = gx_ctr[None] - cx_sub[:, None, None] dy = gy_ctr[None] - cy_sub[:, None, None] d2 = dx ** 2 + dy ** 2 # Gaussian: exp(-d² / (2σ²)) sigma = sigma[:, None, None].clamp(min=0.3) gauss = torch.exp(-d2 / (2 * sigma ** 2)) # Normalize to sum to 1 gauss = gauss / gauss.sum(dim=(1, 2), keepdim=True).clamp(min=1e-8) return gauss def compute_sc3_loss(sc3_logits, batch, model, k=8): """Compute SC³ center classification loss. Uses simple center-based assignment: an anchor is responsible for a GT if the GT center falls within the anchor's cell (with 1-cell margin). Args: sc3_logits: list of [B, k*k, H_i, W_i] per FPN level batch: training batch dict with 'bboxes', 'cls', 'batch_idx' model: YOLO model (for stride info) k: sub-cell grid size Returns: scalar loss (weighted CE on assigned anchors) """ device = sc3_logits[0].device dtype = sc3_logits[0].dtype gt_bboxes = batch['bboxes'] # [N_gt, 4] xywh normalized gt_cls = batch['cls'] # [N_gt, 1] batch_idx = batch['batch_idx'] # [N_gt] bs = sc3_logits[0].shape[0] if len(gt_bboxes) == 0: return torch.tensor(0.0, device=device, requires_grad=True) # Convert GT to xyxy gt_xywh = gt_bboxes.clone() gt_cx = gt_xywh[:, 0] gt_cy = gt_xywh[:, 1] gt_w = gt_xywh[:, 2] gt_h = gt_xywh[:, 3] total_loss = torch.tensor(0.0, device=device) total_assigned = 0 detect = model.model.model[-1] strides = detect.stride for i, logits in enumerate(sc3_logits): B, _, H, W = logits.shape s = strides[i].item() logits_2d = logits.view(B, k, k, H, W) # For each image in batch for b in range(bs): # Get GTs for this image mask = batch_idx == b if mask.sum() == 0: continue img_gt_cx = gt_cx[mask] img_gt_cy = gt_cy[mask] img_gt_w = gt_w[mask] img_gt_h = gt_h[mask] # Anchor centers for this level (pixel coords, normalized) # Each cell (j, i) corresponds to anchor center at (i*s+s/2, j*s+s/2) pixels # In normalized coords: ((i+0.5)*s/W, (j+0.5)*s/H) # We'll work in pixel coords and convert GT to pixel coords # For simplicity, work in grid cell coords # GT center in cell coords for this level: # cell_x = cx_gt * img_w_px / s, cell_y = cy_gt * img_h_px / s # But we don't know img_w_px here... we can use the model input size # Use the imgsz from model or just work with normalized coords # In normalized coords, cell width = s / imgsz # Sub-cell position within cell at (i,j): # sub_x = (cx_gt_normalized * imgsz / s - i) = cx_gt_pixels / s - i # But we only have normalized coords... # Simplest: work in normalized space # Cell i in normalized coords: from i*s/imgsz to (i+1)*s/imgsz # GT center in cell units: cx_gt * imgsz / s # We'll use the imgsz from the model imgsz = 1536 # HARDCODED for now, should get from model for j in range(len(img_gt_cx)): # GT center in pixel coords (approx, assuming square imgsz) gt_cx_px = img_gt_cx[j].item() * imgsz gt_cy_px = img_gt_cy[j].item() * imgsz gt_sz = max(img_gt_w[j].item(), img_gt_h[j].item()) * imgsz / s # cells # Which cell does this GT center fall into? cell_x = int(gt_cx_px / s) cell_y = int(gt_cy_px / s) # Check if within this level's grid if 0 <= cell_x < W and 0 <= cell_y < H: # Sub-cell position sub_x = (gt_cx_px / s - cell_x) * k sub_y = (gt_cy_px / s - cell_y) * k # Clip to valid range sub_x = max(0.0, min(float(k) - 1e-6, sub_x)) sub_y = max(0.0, min(float(k) - 1e-6, sub_y)) # Adaptive sigma: smaller objects → sharper target sigma = 0.5 + 0.5 * gt_sz # [0.5, ~3.0] sigma = max(0.4, min(3.0, sigma)) # Create target target = make_gaussian_target( torch.tensor([sub_x], device=device), torch.tensor([sub_y], device=device), k, torch.tensor([sigma], device=device) ) # [1, k, k] # Predicted logits at this cell pred_logits = logits_2d[b, :, :, cell_y, cell_x] # [k, k] # Cross-entropy loss loss = -(target.squeeze(0) * F.log_softmax(pred_logits.view(1, k * k), dim=1).view(k, k)).sum() # Weight: smaller objects get higher loss weight weight = 1.0 + 2.0 * max(0, 1.0 - gt_sz / 6.0) total_loss += loss * weight total_assigned += 1 if total_assigned == 0: return torch.tensor(0.0, device=device, requires_grad=True) return total_loss / max(total_assigned, 1) # ── Model patching ────────────────────────────────────────── def inject_sc3(model, k=8, sc3_loss_weight=0.5): """Inject SC³ head into a YOLO model and patch the loss. Uses a forward hook on the Detect module to capture FPN features, then passes them through SC3Head to get sub-cell distributions. Adds SC³ loss to the total training loss. """ detect = model.model.model[-1] # Get input channels from detect head's first conv in_channels = [] for i in range(detect.nl): c_in = detect.cv2[i][0].conv.in_channels in_channels.append(c_in) # Create SC3 head sc3 = SC3Head(in_channels, k=k).to(model.device) # Store features from Detect forward pass _features = [] def hook_fn(module, input, output): _features.clear() _features.append(input[0]) # input[0] is x (list of FPN features) handle = detect.register_forward_hook(hook_fn) # Patch loss _orig_loss = model.loss def loss_with_sc3(batch, preds=None): _features.clear() loss_result = _orig_loss(batch, preds) if _features: feats = _features[0] sc3_logits = sc3(feats) sc3_loss = compute_sc3_loss(sc3_logits, batch, model, k=k) loss_result[0] = loss_result[0] + sc3_loss * sc3_loss_weight else: sc3_loss = torch.tensor(0.0, device=model.device) return loss_result import types model.loss = types.MethodType(loss_with_sc3, model) # Store for inference use model.sc3_head = sc3 model._sc3_handle = handle model._sc3_k = k print(f'[SC³] Injected: k={k}, in_channels={in_channels}, loss_weight={sc3_loss_weight}') return model