Download src/superpoint_pruning/distillation/losses.py from PrunaAI/PrunaSuperPoint: direct link, hf CLI and curl.
- Browser
- Download file 5.97 kB
-
https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/losses.py
- Command line
-
hf download hf://PrunaAI/PrunaSuperPoint/src/superpoint_pruning/distillation/losses.py
-
curl -L -o losses.py https://huggingface.co/PrunaAI/PrunaSuperPoint/resolve/main/src/superpoint_pruning/distillation/losses.py
5.97 kB
| 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 | |