TobiasLogic's picture
Upload ObjectModel-v1: code, README, assets (v1 checkpoint)
9b92c75 verified
Raw
History Blame Contribute Delete
9.53 kB
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