import torch import torch.nn.functional as F def convert_keypoints_to_( gt: torch.Tensor, H: int = 768, W: int = 1024 ) -> torch.Tensor: B, N, _ = gt.shape xy = gt.round().long() x = xy[..., 0].clamp(0, W - 1) y = xy[..., 1].clamp(0, H - 1) keypoint_map = torch.zeros((B, H, W), device=gt.device, dtype=torch.float32) b = torch.arange(B, device=gt.device).unsqueeze(1).expand(B, N) keypoint_map[b, y, x] = 1.0 return keypoint_map def detector_loss( keypoint_map: torch.Tensor, # (B,H,W) or (B,1,H,W), binary/bool logits: torch.Tensor, # (B,65,Hc,Wc), raw convPb output valid_mask: torch.Tensor | None = None, # (B,H,W) or (B,1,H,W) grid_size: int = 8, eps: float = 1e-8, ) -> torch.Tensor: # --- labels: space_to_depth + dustbin + random tie-break + argmax --- if keypoint_map.ndim == 3: keypoint_map = keypoint_map.unsqueeze(1) # (B,1,H,W) keypoint_map = keypoint_map.float() # TF NHWC space_to_depth -> PyTorch NCHW pixel_unshuffle labels = F.pixel_unshuffle(keypoint_map, downscale_factor=grid_size) # (B,64,Hc,Wc) dustbin = torch.ones_like(labels[:, :1]) # (B,1,Hc,Wc) labels = torch.cat([2.0 * labels, dustbin], dim=1) # (B,65,Hc,Wc) # same tie-break idea as TF random_uniform(..., 0, 0.1) labels = torch.argmax( labels + 0.1 * torch.rand_like(labels), dim=1 ) # (B,Hc,Wc), long # --- valid mask path --- if valid_mask is None: valid_mask = torch.ones_like(keypoint_map) elif valid_mask.ndim == 3: valid_mask = valid_mask.unsqueeze(1) valid_mask = valid_mask.float() valid_mask = F.pixel_unshuffle( valid_mask, downscale_factor=grid_size ) # (B,64,Hc,Wc) valid_mask = torch.prod(valid_mask, dim=1) # (B,Hc,Wc) # --- sparse softmax cross entropy with weights --- per_cell = F.cross_entropy(logits, labels, reduction="none") # (B,Hc,Wc) weighted = per_cell * valid_mask loss = weighted.sum() / valid_mask.sum().clamp_min(1.0 + eps) return loss def detector_loss_simple( gt_logits: torch.Tensor, pred_logits: torch.Tensor, # (B,65,Hc,Wc), raw convPb output ) -> torch.Tensor: gt_labels = torch.argmax( gt_logits + 0.1 * torch.rand_like(gt_logits), dim=1 ) # (B,Hc,Wc), long return F.cross_entropy(pred_logits, gt_labels, reduction="mean") def detector_kd_kl(teacher_logits, student_logits, T=2.0, valid_mask=None, eps=1e-8): # student_logits, teacher_logits: (B,65,Hc,Wc) log_p_s = F.log_softmax(student_logits / T, dim=1) p_t = F.softmax(teacher_logits / T, dim=1) # KL per cell: (B,Hc,Wc) kl_map = F.kl_div(log_p_s, p_t, reduction="none").sum(dim=1) if valid_mask is None: return (T * T) * kl_map.mean() # valid_mask expected (B,H,W) or (B,1,H,W), convert to cell mask (B,Hc,Wc) if valid_mask.ndim == 3: valid_mask = valid_mask.unsqueeze(1) vm = F.pixel_unshuffle(valid_mask.float(), downscale_factor=8) # (B,64,Hc,Wc) vm = torch.prod(vm, dim=1) # (B,Hc,Wc) return (T * T) * (kl_map * vm).sum() / vm.sum().clamp_min(1.0 + eps) def descriptor_loss_simple( pred_desc: torch.Tensor, # (B, D, Hc, Wc) gt_desc: torch.Tensor, # (B, D, Hc, Wc) valid_mask: torch.Tensor | None = None, # (B,H,W) or (B,1,H,W) grid_size: int = 8, eps: float = 1e-8, ) -> torch.Tensor: pred = F.normalize(pred_desc, p=2, dim=1) gt = F.normalize(gt_desc, p=2, dim=1) # cosine distance per cell per_cell = 1.0 - (pred * gt).sum(dim=1) # (B, Hc, Wc) if valid_mask is None: return per_cell.mean() if valid_mask.ndim == 3: valid_mask = valid_mask.unsqueeze(1) # (B,1,H,W) vm = F.pixel_unshuffle( valid_mask.float(), downscale_factor=grid_size ) # (B,64,Hc,Wc) vm = torch.prod(vm, dim=1) # (B,Hc,Wc) return (per_cell * vm).sum() / vm.sum().clamp_min(1.0 + eps) def descriptor_loss( descriptors: torch.Tensor, # (B, D, Hc, Wc), student target_descriptors: torch.Tensor, # (B, D, Hc, Wc), teacher/GT valid_mask: torch.Tensor | None = None, # (B,H,W) or (B,1,H,W) grid_size: int = 8, positive_margin: float = 1.0, negative_margin: float = 0.2, lambda_d: float = 0.05, eps: float = 1e-8, ) -> torch.Tensor: B, D, Hc, Wc = descriptors.shape HW = Hc * Wc # L2 normalize descriptors desc = F.normalize(descriptors, p=2, dim=1) # (B,D,Hc,Wc) tgt = F.normalize(target_descriptors, p=2, dim=1) # (B,D,Hc,Wc) # Flatten spatial dims desc = desc.flatten(2).transpose(1, 2) # (B,HW,D) tgt = tgt.flatten(2).transpose(1, 2) # (B,HW,D) # Pairwise dot products: (B,HW,HW) dot = torch.bmm(desc, tgt.transpose(1, 2)) dot = F.relu(dot) # TF does double normalization over pairwise axes dot = F.normalize(dot, p=2, dim=2) dot = F.normalize(dot, p=2, dim=1) # Identity correspondence mask s (diagonal) eye = torch.eye(HW, device=dot.device, dtype=dot.dtype).unsqueeze(0) # (1,HW,HW) s = eye.expand(B, -1, -1) positive_dist = F.relu(positive_margin - dot) negative_dist = F.relu(dot - negative_margin) pairwise_loss = ( lambda_d * s * positive_dist + (1.0 - s) * negative_dist ) # (B,HW,HW) # valid mask: same logic as TF space_to_depth + reduce_prod if valid_mask is None: vm = torch.ones( (B, Hc * grid_size, Wc * grid_size), device=dot.device, dtype=dot.dtype ) else: vm = valid_mask if vm.ndim == 3: vm = vm.unsqueeze(1) # (B,1,H,W) vm = vm.float() vm = F.pixel_unshuffle(vm, downscale_factor=grid_size) # (B,grid^2,Hc,Wc) vm = torch.prod(vm, dim=1) # (B,Hc,Wc) vm = vm.reshape(B, HW) # valid target cells vm = vm[:, None, :] # (B,1,HW), broadcast to (B,HW,HW) normalization = vm.sum() * float(HW) + eps loss = (pairwise_loss * vm).sum() / normalization return loss