from __future__ import annotations from collections.abc import Sequence import torch import torch.nn.functional as F from torch import Tensor, nn from .boxes import box_cxcywh_to_xyxy, generalized_box_iou from .matching import hungarian_match, hungarian_match_layers def sigmoid_focal_loss( logits: Tensor, targets: Tensor, alpha: float = 0.25, gamma: float = 2.0 ) -> Tensor: probabilities = logits.sigmoid() ce = F.binary_cross_entropy_with_logits(logits, targets, reduction="none") p_t = probabilities * targets + (1.0 - probabilities) * (1.0 - targets) loss = ce * (1.0 - p_t).pow(gamma) if alpha >= 0: alpha_t = alpha * targets + (1.0 - alpha) * (1.0 - targets) loss = alpha_t * loss return loss class ObjectModelCriterion(nn.Module): def __init__(self, config: dict) -> None: super().__init__() loss_config = config.get("loss", config) self.cost_class = float(loss_config.get("cost_class", 2.0)) self.cost_bbox = float(loss_config.get("cost_bbox", 5.0)) self.cost_giou = float(loss_config.get("cost_giou", 2.0)) self.weight_class = float(loss_config.get("weight_class", 2.0)) self.weight_bbox = float(loss_config.get("weight_bbox", 5.0)) self.weight_giou = float(loss_config.get("weight_giou", 2.0)) self.weight_dense = float(loss_config.get("weight_dense", 1.0)) self.aux_weight = float(loss_config.get("aux_weight", 1.0)) self.dense_topk = int(loss_config.get("dense_topk", 5)) self.alpha = float(loss_config.get("focal_alpha", 0.25)) self.gamma = float(loss_config.get("focal_gamma", 2.0)) def _dense_targets( self, logits: Tensor, targets: Sequence[dict[str, Tensor]], level_index: int, ) -> tuple[Tensor, Tensor, Tensor]: batch, _, height, width = logits.shape device = logits.device target_logits = torch.zeros_like(logits) target_boxes_hwc = torch.zeros(batch, height, width, 4, dtype=torch.float32, device=device) positive = torch.zeros(batch, height, width, dtype=torch.float32, device=device) offsets = torch.tensor( [ (-1, -1), (0, -1), (1, -1), (-1, 0), (0, 0), (1, 0), (-1, 1), (0, 1), (1, 1), ], dtype=torch.int64, device=device, ) distances = offsets.square().sum(dim=1) candidate_count = min(self.dense_topk, len(offsets)) nonempty = [(i, t) for i, t in enumerate(targets) if t["labels"].numel() > 0] if not nonempty: return target_logits, target_boxes_hwc.permute(0, 3, 1, 2), positive all_boxes = torch.cat([t["boxes"] for _, t in nonempty]) all_labels = torch.cat([t["labels"] for _, t in nonempty]) all_batch = torch.cat( [ torch.full((t["labels"].numel(),), i, dtype=torch.int64, device=device) for i, t in nonempty ] ) areas = all_boxes[:, 2] * all_boxes[:, 3] target_levels = torch.where(areas < 0.02, 0, torch.where(areas < 0.15, 1, 2)) level_mask = target_levels == level_index if not bool(level_mask.any()): return target_logits, target_boxes_hwc.permute(0, 3, 1, 2), positive boxes = all_boxes[level_mask] labels = all_labels[level_mask] sel_batch = all_batch[level_mask] grid = (boxes[:, :2] * boxes.new_tensor([width, height])).long() grid[:, 0].clamp_(0, width - 1) grid[:, 1].clamp_(0, height - 1) x = (grid[:, None, 0] + offsets[None, :, 0]).clamp(0, width - 1) y = (grid[:, None, 1] + offsets[None, :, 1]).clamp(0, height - 1) sort_key = distances[None] * ((width + 1) * (height + 1)) sort_key = sort_key + x * (height + 1) + y order = sort_key.argsort(dim=1, stable=True)[:, :candidate_count] x = x.gather(1, order) y = y.gather(1, order) expanded_labels = labels[:, None].expand_as(x) expanded_batch = sel_batch[:, None].expand_as(x) target_logits[expanded_batch, expanded_labels, y, x] = 1.0 flat_cells = (expanded_batch * (height * width) + y * width + x).reshape(-1) owners = torch.full((batch * height * width,), -1, dtype=torch.int64, device=device) source_owners = torch.arange(boxes.shape[0], device=device)[:, None].expand_as(x).reshape(-1) owners.scatter_reduce_(0, flat_cells, source_owners, reduce="amax", include_self=True) occupied = owners >= 0 positive.view(-1)[occupied] = 1.0 target_boxes_hwc.view(-1, 4)[occupied] = boxes[owners[occupied]] return target_logits, target_boxes_hwc.permute(0, 3, 1, 2), positive def _set_loss( self, outputs: dict[str, Tensor], targets: Sequence[dict[str, Tensor]], matches=None ) -> dict[str, Tensor]: logits = outputs["pred_logits"] boxes = outputs["pred_boxes"] if matches is None: matches = hungarian_match( outputs, targets, self.cost_class, self.cost_bbox, self.cost_giou, ) device = logits.device target_classes = torch.zeros_like(logits) normalizer = max(sum(len(target["labels"]) for target in targets), 1) nonempty = [ (batch_index, prediction_indices, target_indices) for batch_index, (prediction_indices, target_indices) in enumerate(matches) if prediction_indices.numel() > 0 ] if nonempty: batch_ids = torch.cat( [torch.full_like(pred_idx, batch_index) for batch_index, pred_idx, _ in nonempty] ) pred_idx_t = torch.cat([pred_idx for _, pred_idx, _ in nonempty]) tgt_idx_t = torch.cat([tgt_idx for _, _, tgt_idx in nonempty]) counts = torch.tensor([len(target["labels"]) for target in targets], device=device) offsets = torch.cat([counts.new_zeros(1), counts.cumsum(0)[:-1]]) global_target_idx = tgt_idx_t + offsets[batch_ids] all_target_boxes = torch.cat([target["boxes"] for target in targets]) all_target_labels = torch.cat([target["labels"] for target in targets]) labels = all_target_labels[global_target_idx] target_classes[batch_ids, pred_idx_t, labels] = 1.0 predicted = boxes[batch_ids, pred_idx_t] expected = all_target_boxes[global_target_idx] else: predicted = None expected = None class_loss = sigmoid_focal_loss(logits, target_classes, self.alpha, self.gamma).sum() class_loss = class_loss / normalizer if predicted is not None: bbox_loss = F.l1_loss(predicted, expected, reduction="sum") / normalizer giou = generalized_box_iou(box_cxcywh_to_xyxy(predicted), box_cxcywh_to_xyxy(expected)) giou_loss = (1.0 - giou.diag()).sum() / normalizer else: bbox_loss = boxes.sum() * 0.0 giou_loss = boxes.sum() * 0.0 return { "loss_class": class_loss * self.weight_class, "loss_bbox": bbox_loss * self.weight_bbox, "loss_giou": giou_loss * self.weight_giou, } def _dense_loss( self, outputs: list[dict[str, Tensor]], targets: Sequence[dict[str, Tensor]] ) -> Tensor: total = outputs[0]["logits"].sum() * 0.0 normalizer = max(sum(len(target["labels"]) for target in targets), 1) for level_index, level_output in enumerate(outputs): logits = level_output["logits"] boxes = level_output["distances"].sigmoid() target_logits, target_boxes, positive = self._dense_targets( logits, targets, level_index ) cls_loss = sigmoid_focal_loss(logits, target_logits, self.alpha, self.gamma) cls_loss = cls_loss.sum() / normalizer positive_mask = positive[:, None].expand_as(boxes) box_loss = (F.l1_loss(boxes, target_boxes, reduction="none") * positive_mask).sum() total = total + cls_loss + box_loss / normalizer return total / len(outputs) def forward( self, outputs: dict[str, Tensor], targets: Sequence[dict[str, Tensor]] ) -> dict[str, Tensor]: layer_outputs = [outputs, *outputs.get("aux_outputs", [])] layer_matches = hungarian_match_layers( layer_outputs, targets, self.cost_class, self.cost_bbox, self.cost_giou, ) primary = self._set_loss(outputs, targets, layer_matches[0]) total = sum(primary.values()) for auxiliary, matches in zip( outputs.get("aux_outputs", []), layer_matches[1:], strict=True ): auxiliary_losses = self._set_loss(auxiliary, targets, matches) total = total + self.aux_weight * sum(auxiliary_losses.values()) / max( len(outputs["aux_outputs"]), 1 ) if "dense_outputs" in outputs: dense = self._dense_loss(outputs["dense_outputs"], targets) primary["loss_dense"] = dense * self.weight_dense total = total + primary["loss_dense"] primary["loss_total"] = total return primary