oskarkuuse-pruna's picture
commit repo mirror
6979012
Raw History Blame Contribute Delete
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