Ultralytics YOLO27:
Get Started

Reference for ultralytics/utils/tal.py#

Improvements

This page is sourced from https://github.com/ultralytics/ultralytics/blob/main/ultralytics/utils/tal.py. Have an improvement or example to add? Open a Pull Request — thank you! 🙏


Summary

Class ultralytics.utils.tal.TaskAlignedAssigner#

TaskAlignedAssigner(
    topk: int = 13,
    num_classes: int = 80,
    alpha: float = 1.0,
    beta: float = 6.0,
    stride: list | None = None,
    eps: float = 1e-9,
    topk2: int | None = None,
)

Bases: nn.Module

A task-aligned assigner for object detection.

This class assigns ground-truth (gt) objects to anchors based on the task-aligned metric, which combines both classification and localization information.

Args

NameTypeDescriptionDefault
topkint, optionalThe number of top candidates to consider.13
num_classesint, optionalThe number of object classes.80
alphafloat, optionalThe alpha parameter for the classification component of the task-aligned metric.1.0
betafloat, optionalThe beta parameter for the localization component of the task-aligned metric.6.0
stridelist, optionalList of stride values for different feature levels.None
epsfloat, optionalA small value to prevent division by zero.1e-9
topk2int, optionalSecondary topk value for additional filtering. If None, topk is used.None

Attributes

NameTypeDescription
topkintThe number of top candidates to consider.
topk2intSecondary topk value for additional filtering, defaults to topk.
num_classesintThe number of object classes.
alphafloatThe alpha parameter for the classification component of the task-aligned metric.
betafloatThe beta parameter for the localization component of the task-aligned metric.
stridelistList of stride values for different feature levels.
stride_valintMinimum ground-truth box side in select_candidates_in_gts; smaller sides are enlarged to it.
epsfloatA small value to prevent division by zero.

Methods

NameDescription
_forwardCompute the task-aligned assignment.
forwardCompute the task-aligned assignment.
get_box_metricsCompute alignment metric given predicted and ground truth bounding boxes.
get_pos_maskGet positive mask for each ground truth box.
get_targetsCompute target labels, target bounding boxes, and target scores for the positive anchor points.
iou_calculationCalculate CIoU for horizontal bounding boxes, clamped to be non-negative.
select_candidates_in_gtsSelect positive anchor centers within ground truth bounding boxes.
select_highest_overlapsSelect anchor boxes with highest IoU when assigned to multiple ground truths.
select_topk_candidatesSelect the top-k candidates based on the given metrics.
GitHubultralytics/utils/tal.py
class TaskAlignedAssigner(nn.Module):
    """A task-aligned assigner for object detection.

    This class assigns ground-truth (gt) objects to anchors based on the task-aligned metric, which combines both
    classification and localization information.

    Attributes:
        topk (int): The number of top candidates to consider.
        topk2 (int): Secondary topk value for additional filtering, defaults to topk.
        num_classes (int): The number of object classes.
        alpha (float): The alpha parameter for the classification component of the task-aligned metric.
        beta (float): The beta parameter for the localization component of the task-aligned metric.
        stride (list): List of stride values for different feature levels.
        stride_val (int): Minimum ground-truth box side in select_candidates_in_gts; smaller sides are enlarged to it.
        eps (float): A small value to prevent division by zero.
    """

    def __init__(
        self,
        topk: int = 13,
        num_classes: int = 80,
        alpha: float = 1.0,
        beta: float = 6.0,
        stride: list | None = None,
        eps: float = 1e-9,
        topk2: int | None = None,
    ):
        """Initialize a TaskAlignedAssigner object with customizable hyperparameters.

        Args:
            topk (int, optional): The number of top candidates to consider.
            num_classes (int, optional): The number of object classes.
            alpha (float, optional): The alpha parameter for the classification component of the task-aligned metric.
            beta (float, optional): The beta parameter for the localization component of the task-aligned metric.
            stride (list, optional): List of stride values for different feature levels.
            eps (float, optional): A small value to prevent division by zero.
            topk2 (int, optional): Secondary topk value for additional filtering. If None, topk is used.
        """
        super().__init__()
        self.topk = topk
        self.topk2 = topk2 or topk
        self.num_classes = num_classes
        self.alpha = alpha
        self.beta = beta
        self.stride = stride if stride is not None else [8, 16, 32]
        self.stride_val = self.stride[1] if len(self.stride) > 1 else self.stride[0]
        self.eps = eps
        self._oom_warned = False

Method ultralytics.utils.tal.TaskAlignedAssigner._forward#

def _forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)

Compute the task-aligned assignment.

Args

NameTypeDescriptionDefault
pd_scorestorch.TensorPredicted classification scores with shape (bs, num_total_anchors, num_classes).required
pd_bboxestorch.TensorPredicted bounding boxes with shape (bs, num_total_anchors, 4).required
anc_pointstorch.TensorAnchor points with shape (num_total_anchors, 2).required
gt_labelstorch.TensorGround truth labels with shape (bs, n_max_boxes, 1).required
gt_bboxestorch.TensorGround truth boxes with shape (bs, n_max_boxes, 4).required
mask_gttorch.TensorMask for valid ground truth boxes with shape (bs, n_max_boxes, 1).required

Returns

TypeDescription
target_labels (torch.Tensor)Target labels with shape (bs, num_total_anchors).
target_bboxes (torch.Tensor)Target bounding boxes with shape (bs, num_total_anchors, 4).
target_scores (torch.Tensor)Target scores with shape (bs, num_total_anchors, num_classes).
fg_mask (torch.Tensor)Foreground mask with shape (bs, num_total_anchors).
target_gt_idx (torch.Tensor)Target ground truth indices with shape (bs, num_total_anchors).
GitHubultralytics/utils/tal.py
def _forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt):
    """Compute the task-aligned assignment.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

    Returns:
        target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
        target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
        target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
        fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
        target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).
    """
    mask_pos, align_metric, overlaps = self.get_pos_mask(
        pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt
    )

    target_gt_idx, fg_mask, mask_pos = self.select_highest_overlaps(
        mask_pos, overlaps, self.n_max_boxes, align_metric
    )

    # Assigned target
    target_labels, target_bboxes, target_scores = self.get_targets(gt_labels, gt_bboxes, target_gt_idx, fg_mask)

    # Normalize
    align_metric *= mask_pos
    pos_align_metrics = align_metric.amax(dim=-1, keepdim=True)  # b, max_num_obj
    overlaps *= mask_pos
    pos_overlaps = overlaps.amax(dim=-1, keepdim=True)  # b, max_num_obj
    align_metric.mul_(pos_overlaps).div_(pos_align_metrics + self.eps)
    norm_align_metric = align_metric.amax(-2).unsqueeze(-1)
    target_scores = target_scores * norm_align_metric

    return target_labels, target_bboxes, target_scores, fg_mask.bool(), target_gt_idx

Method ultralytics.utils.tal.TaskAlignedAssigner.forward#

def forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)

Compute the task-aligned assignment.

Args

NameTypeDescriptionDefault
pd_scorestorch.TensorPredicted classification scores with shape (bs, num_total_anchors, num_classes).required
pd_bboxestorch.TensorPredicted bounding boxes with shape (bs, num_total_anchors, 4).required
anc_pointstorch.TensorAnchor points with shape (num_total_anchors, 2).required
gt_labelstorch.TensorGround truth labels with shape (bs, n_max_boxes, 1).required
gt_bboxestorch.TensorGround truth boxes with shape (bs, n_max_boxes, 4).required
mask_gttorch.TensorMask for valid ground truth boxes with shape (bs, n_max_boxes, 1).required

Returns

TypeDescription
target_labels (torch.Tensor)Target labels with shape (bs, num_total_anchors).
target_bboxes (torch.Tensor)Target bounding boxes with shape (bs, num_total_anchors, 4).
target_scores (torch.Tensor)Target scores with shape (bs, num_total_anchors, num_classes).
fg_mask (torch.Tensor)Foreground mask with shape (bs, num_total_anchors).
target_gt_idx (torch.Tensor)Target ground truth indices with shape (bs, num_total_anchors).
Notes

On a CUDA out-of-memory error the assignment is retried one image at a time.

References

GitHubultralytics/utils/tal.py
@torch.no_grad()
def forward(self, pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt):
    """Compute the task-aligned assignment.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

    Returns:
        target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
        target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
        target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
        fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
        target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).

    Notes:
        On a CUDA out-of-memory error the assignment is retried one image at a time.

    References:
        https://github.com/Nioolek/PPYOLOE_pytorch/blob/master/ppyoloe/assigner/tal_assigner.py
    """
    self.bs = pd_scores.shape[0]
    self.n_max_boxes = gt_bboxes.shape[1]
    if self.n_max_boxes == 0:
        return (
            torch.full_like(pd_scores[..., 0], self.num_classes),
            torch.zeros_like(pd_bboxes),
            torch.zeros_like(pd_scores),
            torch.zeros_like(pd_scores[..., 0], dtype=torch.bool),
            torch.zeros_like(pd_scores[..., 0]),
        )

    try:
        return self._forward(pd_scores, pd_bboxes, anc_points, gt_labels, gt_bboxes, mask_gt)
    except RuntimeError as e:
        if "out of memory" not in str(e).lower():
            raise
    # Recover outside the except block so e.__traceback__ releases the failed attempt's GPU intermediates.
    bs, n_max_boxes = self.bs, self.n_max_boxes
    if not self._oom_warned:
        LOGGER.warning(
            f"CUDA out of memory in TaskAlignedAssigner with batch_size={bs} and max_num_obj={n_max_boxes}; "
            "retrying assignment one image at a time on GPU. Model forward batch size is unchanged."
        )
        self._oom_warned = True
    last_gt_idx = (
        mask_gt.squeeze(-1)
        .bool()
        .mul(torch.arange(1, n_max_boxes + 1, device=mask_gt.device))
        .amax(1)
        .clamp_(min=1)
        .tolist()
    )
    self.bs = 1
    results = None
    try:
        for i, self.n_max_boxes in enumerate(last_gt_idx):
            result = self._forward(
                pd_scores[i : i + 1],
                pd_bboxes[i : i + 1],
                anc_points,
                gt_labels[i : i + 1, : self.n_max_boxes],
                gt_bboxes[i : i + 1, : self.n_max_boxes],
                mask_gt[i : i + 1, : self.n_max_boxes],
            )
            if results is None:
                results = tuple(x.new_empty((bs, *x.shape[1:])) for x in result)
            for output, x in zip(results, result):
                output[i] = x[0]
    finally:
        self.bs, self.n_max_boxes = bs, n_max_boxes
    return results

Method ultralytics.utils.tal.TaskAlignedAssigner.get_box_metrics#

def get_box_metrics(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, mask_gt)

Compute alignment metric given predicted and ground truth bounding boxes.

Args

NameTypeDescriptionDefault
pd_scorestorch.TensorPredicted classification scores with shape (bs, num_total_anchors, num_classes).required
pd_bboxestorch.TensorPredicted bounding boxes with shape (bs, num_total_anchors, 4).required
gt_labelstorch.TensorGround truth labels with shape (bs, n_max_boxes, 1).required
gt_bboxestorch.TensorGround truth boxes with shape (bs, n_max_boxes, 4).required
mask_gttorch.TensorMask for valid ground truth boxes with shape (bs, n_max_boxes, h*w).required

Returns

TypeDescription
align_metric (torch.Tensor)Alignment metric combining classification and localization with shape (bs, n_max_boxes, h*w).
overlaps (torch.Tensor)IoU overlaps between predicted and ground truth boxes with shape (bs, n_max_boxes, h*w).
GitHubultralytics/utils/tal.py
def get_box_metrics(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, mask_gt):
    """Compute alignment metric given predicted and ground truth bounding boxes.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, h*w).

    Returns:
        align_metric (torch.Tensor): Alignment metric combining classification and localization with shape (bs,
            n_max_boxes, h*w).
        overlaps (torch.Tensor): IoU overlaps between predicted and ground truth boxes with shape (bs, n_max_boxes,
            h*w).
    """
    na = pd_bboxes.shape[-2]
    mask_gt = mask_gt.bool()  # b, max_num_obj, h*w
    shape = self.bs, self.n_max_boxes, na
    indices = mask_gt.nonzero(as_tuple=True)
    bbox_scores = pd_scores[indices[0], indices[2], gt_labels[indices[0], indices[1], 0].long()]
    overlap_values = self.iou_calculation(gt_bboxes[indices[:2]], pd_bboxes[indices[0], indices[2]])
    align_values = bbox_scores.pow(self.alpha) * overlap_values.pow(self.beta)

    overlaps = torch.zeros(shape, dtype=pd_bboxes.dtype, device=pd_bboxes.device)
    align_metric = torch.zeros(shape, dtype=align_values.dtype, device=pd_scores.device)
    overlaps[indices] = overlap_values
    align_metric[indices] = align_values
    return align_metric, overlaps

Method ultralytics.utils.tal.TaskAlignedAssigner.get_pos_mask#

def get_pos_mask(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt)

Get positive mask for each ground truth box.

Args

NameTypeDescriptionDefault
pd_scorestorch.TensorPredicted classification scores with shape (bs, num_total_anchors, num_classes).required
pd_bboxestorch.TensorPredicted bounding boxes with shape (bs, num_total_anchors, 4).required
gt_labelstorch.TensorGround truth labels with shape (bs, n_max_boxes, 1).required
gt_bboxestorch.TensorGround truth boxes with shape (bs, n_max_boxes, 4).required
anc_pointstorch.TensorAnchor points with shape (num_total_anchors, 2).required
mask_gttorch.TensorMask for valid ground truth boxes with shape (bs, n_max_boxes, 1).required

Returns

TypeDescription
mask_pos (torch.Tensor)Positive mask with shape (bs, max_num_obj, h*w).
align_metric (torch.Tensor)Alignment metric with shape (bs, max_num_obj, h*w).
overlaps (torch.Tensor)Overlaps between predicted vs ground truth boxes with shape (bs, max_num_obj, h*w).
GitHubultralytics/utils/tal.py
def get_pos_mask(self, pd_scores, pd_bboxes, gt_labels, gt_bboxes, anc_points, mask_gt):
    """Get positive mask for each ground truth box.

    Args:
        pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
        pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
        gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
        gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
        anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

    Returns:
        mask_pos (torch.Tensor): Positive mask with shape (bs, max_num_obj, h*w).
        align_metric (torch.Tensor): Alignment metric with shape (bs, max_num_obj, h*w).
        overlaps (torch.Tensor): Overlaps between predicted vs ground truth boxes with shape (bs, max_num_obj, h*w).
    """
    mask_in_gts = self.select_candidates_in_gts(anc_points, gt_bboxes, mask_gt)
    # Get anchor_align metric, (b, max_num_obj, h*w)
    align_metric, overlaps = self.get_box_metrics(pd_scores, pd_bboxes, gt_labels, gt_bboxes, mask_in_gts * mask_gt)
    # Get topk_metric mask, (b, max_num_obj, h*w)
    mask_topk = self.select_topk_candidates(align_metric, topk_mask=mask_gt.expand(-1, -1, self.topk).bool())
    # Merge all mask to a final mask, (b, max_num_obj, h*w)
    mask_pos = mask_topk.mul_(mask_in_gts).mul_(mask_gt.bool())

    return mask_pos, align_metric, overlaps

Method ultralytics.utils.tal.TaskAlignedAssigner.get_targets#

def get_targets(self, gt_labels, gt_bboxes, target_gt_idx, fg_mask)

Compute target labels, target bounding boxes, and target scores for the positive anchor points.

Args

NameTypeDescriptionDefault
gt_labelstorch.TensorGround truth labels of shape (b, max_num_obj, 1), where b is the batch size and max_num_obj is the maximum number of objects.required
gt_bboxestorch.TensorGround truth bounding boxes of shape (b, max_num_obj, 4).required
target_gt_idxtorch.TensorIndices of the assigned ground truth objects for positive anchor points, with shape (b, hw), where hw is the total number of anchor points.required
fg_masktorch.TensorA boolean tensor of shape (b, h*w) indicating the positive (foreground) anchor points.required

Returns

TypeDescription
target_labels (torch.Tensor)Target labels for positive anchor points with shape (b, h*w).
target_bboxes (torch.Tensor)Target bounding boxes for positive anchor points with shape (b, h*w, 4).
target_scores (torch.Tensor)Target scores for positive anchor points with shape (b, h*w, num_classes).
GitHubultralytics/utils/tal.py
def get_targets(self, gt_labels, gt_bboxes, target_gt_idx, fg_mask):
    """Compute target labels, target bounding boxes, and target scores for the positive anchor points.

    Args:
        gt_labels (torch.Tensor): Ground truth labels of shape (b, max_num_obj, 1), where b is the batch size and
            max_num_obj is the maximum number of objects.
        gt_bboxes (torch.Tensor): Ground truth bounding boxes of shape (b, max_num_obj, 4).
        target_gt_idx (torch.Tensor): Indices of the assigned ground truth objects for positive anchor points, with
            shape (b, h*w), where h*w is the total number of anchor points.
        fg_mask (torch.Tensor): A boolean tensor of shape (b, h*w) indicating the positive (foreground) anchor
            points.

    Returns:
        target_labels (torch.Tensor): Target labels for positive anchor points with shape (b, h*w).
        target_bboxes (torch.Tensor): Target bounding boxes for positive anchor points with shape (b, h*w, 4).
        target_scores (torch.Tensor): Target scores for positive anchor points with shape (b, h*w, num_classes).
    """
    # Assigned target labels, (b, 1)
    batch_ind = torch.arange(end=self.bs, dtype=torch.int64, device=gt_labels.device)[..., None]
    target_gt_idx = target_gt_idx + batch_ind * self.n_max_boxes  # (b, h*w)
    target_labels = gt_labels.long().flatten()[target_gt_idx]  # (b, h*w)

    # Assigned target boxes, (b, max_num_obj, 4) -> (b, h*w, 4)
    target_bboxes = gt_bboxes.view(-1, gt_bboxes.shape[-1])[target_gt_idx]

    # Assigned target scores
    target_labels.clamp_(0)

    # 10x faster than F.one_hot()
    target_scores = torch.zeros(
        (target_labels.shape[0], target_labels.shape[1], self.num_classes),
        dtype=torch.int8,
        device=target_labels.device,
    )  # (b, h*w, 80)
    target_scores.scatter_(2, target_labels.unsqueeze(-1), 1)

    target_scores = target_scores * (fg_mask[:, :, None] > 0)

    return target_labels, target_bboxes, target_scores

Method ultralytics.utils.tal.TaskAlignedAssigner.iou_calculation#

def iou_calculation(self, gt_bboxes, pd_bboxes)

Calculate CIoU for horizontal bounding boxes, clamped to be non-negative.

Args

NameTypeDescriptionDefault
gt_bboxestorch.TensorGround truth boxes in xyxy format with shape (N, 4).required
pd_bboxestorch.TensorPredicted boxes in xyxy format with shape (N, 4).required

Returns

TypeDescription
torch.TensorCIoU values between each pair of boxes with shape (N,).
GitHubultralytics/utils/tal.py
def iou_calculation(self, gt_bboxes, pd_bboxes):
    """Calculate CIoU for horizontal bounding boxes, clamped to be non-negative.

    Args:
        gt_bboxes (torch.Tensor): Ground truth boxes in xyxy format with shape (N, 4).
        pd_bboxes (torch.Tensor): Predicted boxes in xyxy format with shape (N, 4).

    Returns:
        (torch.Tensor): CIoU values between each pair of boxes with shape (N,).
    """
    return bbox_iou(gt_bboxes, pd_bboxes, xywh=False, CIoU=True).squeeze(-1).clamp_(0)

Method ultralytics.utils.tal.TaskAlignedAssigner.select_candidates_in_gts#

def select_candidates_in_gts(self, xy_centers, gt_bboxes, mask_gt, eps=1e-9)

Select positive anchor centers within ground truth bounding boxes.

Args

NameTypeDescriptionDefault
xy_centerstorch.TensorAnchor center coordinates, shape (h*w, 2).required
gt_bboxestorch.TensorGround truth bounding boxes, shape (b, n_boxes, 4).required
mask_gttorch.TensorMask for valid ground truth boxes, shape (b, n_boxes, 1).required
epsfloat, optionalSmall value for numerical stability.1e-9

Returns

TypeDescription
torch.TensorBoolean mask of positive anchors, shape (b, n_boxes, h*w).
Notes
  • b: batch size, n_boxes: number of ground truth boxes, h: height, w: width.
  • Bounding box format: [x_min, y_min, x_max, y_max].
  • Valid boxes with a side smaller than stride_val are enlarged to stride_val about their center.
GitHubultralytics/utils/tal.py
def select_candidates_in_gts(self, xy_centers, gt_bboxes, mask_gt, eps=1e-9):
    """Select positive anchor centers within ground truth bounding boxes.

    Args:
        xy_centers (torch.Tensor): Anchor center coordinates, shape (h*w, 2).
        gt_bboxes (torch.Tensor): Ground truth bounding boxes, shape (b, n_boxes, 4).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes, shape (b, n_boxes, 1).
        eps (float, optional): Small value for numerical stability.

    Returns:
        (torch.Tensor): Boolean mask of positive anchors, shape (b, n_boxes, h*w).

    Notes:
        - b: batch size, n_boxes: number of ground truth boxes, h: height, w: width.
        - Bounding box format: [x_min, y_min, x_max, y_max].
        - Valid boxes with a side smaller than stride_val are enlarged to stride_val about their center.
    """
    gt_bboxes_xywh = xyxy2xywh(gt_bboxes)
    wh_mask = gt_bboxes_xywh[..., 2:] < self.stride_val  # floor tiny sides so the pool grows monotonically
    gt_bboxes_xywh[..., 2:] = torch.where(
        (wh_mask * mask_gt).bool(),
        torch.tensor(self.stride_val, dtype=gt_bboxes_xywh.dtype, device=gt_bboxes_xywh.device),
        gt_bboxes_xywh[..., 2:],
    )
    gt_bboxes = xywh2xyxy(gt_bboxes_xywh)

    lt, rb = gt_bboxes.unsqueeze(2).chunk(2, 3)  # (b, n_boxes, 1, 2) left-top, right-bottom
    mask = xy_centers[:, 0] - lt[..., 0] > eps
    mask &= xy_centers[:, 1] - lt[..., 1] > eps
    mask &= rb[..., 0] - xy_centers[:, 0] > eps
    mask &= rb[..., 1] - xy_centers[:, 1] > eps
    return mask

Method ultralytics.utils.tal.TaskAlignedAssigner.select_highest_overlaps#

def select_highest_overlaps(self, mask_pos, overlaps, n_max_boxes, align_metric)

Select anchor boxes with highest IoU when assigned to multiple ground truths.

Args

NameTypeDescriptionDefault
mask_postorch.TensorPositive mask, shape (b, n_max_boxes, h*w).required
overlapstorch.TensorIoU overlaps, shape (b, n_max_boxes, h*w).required
n_max_boxesintMaximum number of ground truth boxes.required
align_metrictorch.TensorAlignment metric, shape (b, n_max_boxes, h*w), used for the topk2 filtering.required

Returns

TypeDescription
target_gt_idx (torch.Tensor)Indices of assigned ground truths, shape (b, h*w).
fg_mask (torch.Tensor)Foreground mask, shape (b, h*w).
mask_pos (torch.Tensor)Updated positive mask, shape (b, n_max_boxes, h*w).
GitHubultralytics/utils/tal.py
def select_highest_overlaps(self, mask_pos, overlaps, n_max_boxes, align_metric):
    """Select anchor boxes with highest IoU when assigned to multiple ground truths.

    Args:
        mask_pos (torch.Tensor): Positive mask, shape (b, n_max_boxes, h*w).
        overlaps (torch.Tensor): IoU overlaps, shape (b, n_max_boxes, h*w).
        n_max_boxes (int): Maximum number of ground truth boxes.
        align_metric (torch.Tensor): Alignment metric, shape (b, n_max_boxes, h*w), used for the topk2 filtering.

    Returns:
        target_gt_idx (torch.Tensor): Indices of assigned ground truths, shape (b, h*w).
        fg_mask (torch.Tensor): Foreground mask, shape (b, h*w).
        mask_pos (torch.Tensor): Updated positive mask, shape (b, n_max_boxes, h*w).
    """
    # Convert (b, n_max_boxes, h*w) -> (b, h*w)
    fg_mask = mask_pos.sum(-2)
    # Anchors assigned to multiple gt_bboxes keep the highest overlap; a no-op when there are none, without a sync
    mask_multi_gts = (fg_mask.unsqueeze(1) > 1).expand(-1, n_max_boxes, -1)  # (b, n_max_boxes, h*w)
    max_overlaps_idx = overlaps.max(1).indices  # (b, h*w)
    is_max_overlaps = torch.zeros(mask_pos.shape, dtype=mask_pos.dtype, device=mask_pos.device)
    is_max_overlaps.scatter_(1, max_overlaps_idx.unsqueeze(1), 1)
    mask_pos = torch.where(mask_multi_gts, is_max_overlaps, mask_pos)  # (b, n_max_boxes, h*w)

    if self.topk2 != self.topk:
        align_metric = align_metric * mask_pos  # update overlaps
        # (b, n_max_boxes, topk2)
        max_overlaps_idx = torch.topk(align_metric, self.topk2, dim=-1, largest=True).indices
        topk_idx = torch.zeros(mask_pos.shape, dtype=mask_pos.dtype, device=mask_pos.device)  # update mask_pos
        topk_idx.scatter_(-1, max_overlaps_idx, 1.0)
        mask_pos *= topk_idx
    # Each anchor now serves at most one gt, so the column max is both its foreground flag and its gt index
    fg_mask, target_gt_idx = mask_pos.max(-2)  # (b, h*w)
    return target_gt_idx, fg_mask, mask_pos

Method ultralytics.utils.tal.TaskAlignedAssigner.select_topk_candidates#

def select_topk_candidates(self, metrics, topk_mask=None)

Select the top-k candidates based on the given metrics.

Args

NameTypeDescriptionDefault
metricstorch.TensorA tensor of shape (b, max_num_obj, hw), where b is the batch size, max_num_obj is the maximum number of objects, and hw represents the total number of anchor points.required
topk_masktorch.Tensor, optionalAn optional boolean tensor of shape (b, max_num_obj, topk), where topk is the number of top candidates to consider. If not provided, it is derived from whether each row's largest metric exceeds eps.None

Returns

TypeDescription
torch.TensorAn int8 tensor of shape (b, max_num_obj, h*w) that is 1 for the selected top-k candidates and 0 elsewhere.
GitHubultralytics/utils/tal.py
def select_topk_candidates(self, metrics, topk_mask=None):
    """Select the top-k candidates based on the given metrics.

    Args:
        metrics (torch.Tensor): A tensor of shape (b, max_num_obj, h*w), where b is the batch size, max_num_obj is
            the maximum number of objects, and h*w represents the total number of anchor points.
        topk_mask (torch.Tensor, optional): An optional boolean tensor of shape (b, max_num_obj, topk), where topk
            is the number of top candidates to consider. If not provided, it is derived from whether each row's
            largest metric exceeds eps.

    Returns:
        (torch.Tensor): An int8 tensor of shape (b, max_num_obj, h*w) that is 1 for the selected top-k candidates
            and 0 elsewhere.
    """
    # (b, max_num_obj, topk)
    topk_metrics, topk_idxs = torch.topk(metrics, self.topk, dim=-1, largest=True)
    if topk_mask is None:
        topk_mask = (topk_metrics.max(-1, keepdim=True)[0] > self.eps).expand_as(topk_idxs)
    # (b, max_num_obj, topk)
    topk_idxs.masked_fill_(~topk_mask, 0)

    # Count how many of the topk lists select each anchor; scatter_add_ accumulates duplicate indices in one pass
    count_tensor = torch.zeros(metrics.shape, dtype=torch.int8, device=topk_idxs.device)
    count_tensor.scatter_add_(-1, topk_idxs, torch.ones_like(topk_idxs, dtype=torch.int8))
    # Filter invalid bboxes
    count_tensor.masked_fill_(count_tensor > 1, 0)

    return count_tensor





Class ultralytics.utils.tal.RotatedTaskAlignedAssigner#

RotatedTaskAlignedAssigner(
    topk: int = 13,
    num_classes: int = 80,
    alpha: float = 1.0,
    beta: float = 6.0,
    stride: list | None = None,
    eps: float = 1e-9,
    topk2: int | None = None,
)

Bases: TaskAlignedAssigner

Assigns ground-truth objects to rotated bounding boxes using a task-aligned metric.

Args

NameTypeDescriptionDefault
topkint, optionalThe number of top candidates to consider.13
num_classesint, optionalThe number of object classes.80
alphafloat, optionalThe alpha parameter for the classification component of the task-aligned metric.1.0
betafloat, optionalThe beta parameter for the localization component of the task-aligned metric.6.0
stridelist, optionalList of stride values for different feature levels.None
epsfloat, optionalA small value to prevent division by zero.1e-9
topk2int, optionalSecondary topk value for additional filtering. If None, topk is used.None

Methods

NameDescription
iou_calculationCalculate probabilistic IoU (ProbIoU) for rotated bounding boxes, clamped to be non-negative.
select_candidates_in_gtsSelect positive anchor centers within rotated ground truth bounding boxes.
GitHubultralytics/utils/tal.py
class RotatedTaskAlignedAssigner(TaskAlignedAssigner):
    """Assigns ground-truth objects to rotated bounding boxes using a task-aligned metric."""

Method ultralytics.utils.tal.RotatedTaskAlignedAssigner.iou_calculation#

def iou_calculation(self, gt_bboxes, pd_bboxes)

Calculate probabilistic IoU (ProbIoU) for rotated bounding boxes, clamped to be non-negative.

GitHubultralytics/utils/tal.py
def iou_calculation(self, gt_bboxes, pd_bboxes):
    """Calculate probabilistic IoU (ProbIoU) for rotated bounding boxes, clamped to be non-negative."""
    return probiou(gt_bboxes, pd_bboxes).squeeze(-1).clamp_(0)

Method ultralytics.utils.tal.RotatedTaskAlignedAssigner.select_candidates_in_gts#

def select_candidates_in_gts(self, xy_centers, gt_bboxes, mask_gt)

Select positive anchor centers within rotated ground truth bounding boxes.

Args

NameTypeDescriptionDefault
xy_centerstorch.TensorAnchor center coordinates with shape (h*w, 2).required
gt_bboxestorch.TensorGround truth bounding boxes in xywhr format with shape (b, n_boxes, 5).required
mask_gttorch.TensorMask for valid ground truth boxes with shape (b, n_boxes, 1).required

Returns

TypeDescription
torch.TensorBoolean mask of positive anchors with shape (b, n_boxes, h*w).
GitHubultralytics/utils/tal.py
def select_candidates_in_gts(self, xy_centers, gt_bboxes, mask_gt):
    """Select positive anchor centers within rotated ground truth bounding boxes.

    Args:
        xy_centers (torch.Tensor): Anchor center coordinates with shape (h*w, 2).
        gt_bboxes (torch.Tensor): Ground truth bounding boxes in xywhr format with shape (b, n_boxes, 5).
        mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (b, n_boxes, 1).

    Returns:
        (torch.Tensor): Boolean mask of positive anchors with shape (b, n_boxes, h*w).
    """
    gt_bboxes_clone = gt_bboxes.clone()
    wh_mask = gt_bboxes_clone[..., 2:4] < self.stride_val
    gt_bboxes_clone[..., 2:4] = torch.where(
        (wh_mask * mask_gt).bool(),
        torch.tensor(self.stride_val, dtype=gt_bboxes_clone.dtype, device=gt_bboxes_clone.device),
        gt_bboxes_clone[..., 2:4],
    )

    # (b, n_boxes, 5) --> (b, n_boxes, 4, 2)
    corners = xywhr2xyxyxyxy(gt_bboxes_clone)
    # (b, n_boxes, 1, 2)
    a, b, _, d = corners.split(1, dim=-2)
    ab = b - a
    ad = d - a

    # (b, n_boxes, h*w) per coordinate
    apx = xy_centers[:, 0] - a[..., 0]
    apy = xy_centers[:, 1] - a[..., 1]
    norm_ab = (ab * ab).sum(dim=-1)
    norm_ad = (ad * ad).sum(dim=-1)
    ap_dot_ab = apx * ab[..., 0] + apy * ab[..., 1]
    ap_dot_ad = apx * ad[..., 0] + apy * ad[..., 1]
    return (ap_dot_ab >= 0) & (ap_dot_ab <= norm_ab) & (ap_dot_ad >= 0) & (ap_dot_ad <= norm_ad)  # is_in_box





Function ultralytics.utils.tal.make_anchors#

def make_anchors(feats, strides, grid_cell_offset=0.5)

Generate anchor points and stride tensors from feature maps.

Args

NameTypeDescriptionDefault
featslist[torch.Tensor] | torch.TensorFeature maps with shape (b, c, h, w) per level, or a tensor of per-level (h, w) sizes.required
stridestorch.Tensor | listStride of each feature level.required
grid_cell_offsetfloatOffset added to grid cell indices, 0.5 for cell centers.0.5

Returns

TypeDescription
anchor_points (torch.Tensor)Anchor points in grid units with shape (N, 2), where N is the sum of h*w over all levels.
stride_tensor (torch.Tensor)Stride of each anchor point with shape (N, 1).
GitHubultralytics/utils/tal.py
def make_anchors(feats, strides, grid_cell_offset=0.5):
    """Generate anchor points and stride tensors from feature maps.

    Args:
        feats (list[torch.Tensor] | torch.Tensor): Feature maps with shape (b, c, h, w) per level, or a tensor of
            per-level (h, w) sizes.
        strides (torch.Tensor | list): Stride of each feature level.
        grid_cell_offset (float): Offset added to grid cell indices, 0.5 for cell centers.

    Returns:
        anchor_points (torch.Tensor): Anchor points in grid units with shape (N, 2), where N is the sum of h*w over all
            levels.
        stride_tensor (torch.Tensor): Stride of each anchor point with shape (N, 1).
    """
    anchor_points, stride_tensor = [], []
    assert feats is not None
    dtype = feats[0].dtype
    for i in range(len(feats)):  # use len(feats) to avoid TracerWarning from iterating over strides tensor
        stride = strides[i]
        h, w = feats[i].shape[2:] if isinstance(feats, list) else (int(feats[i][0]), int(feats[i][1]))
        # no cumsum (nondeterministic on CUDA), no device= (baked into traces), no out= (does not convert to CoreML)
        sx = torch.arange(w).type_as(feats[0]) + grid_cell_offset  # shift x
        sy = torch.arange(h).type_as(feats[0]) + grid_cell_offset  # shift y
        sy, sx = torch.meshgrid(sy, sx, indexing="ij") if TORCH_1_11 else torch.meshgrid(sy, sx)
        anchor_points.append(torch.stack((sx, sy), -1).view(-1, 2))
        stride_tensor.append(feats[0].new_full((h * w, 1), stride, dtype=dtype))
    return torch.cat(anchor_points), torch.cat(stride_tensor)





Function ultralytics.utils.tal.dist2bbox#

def dist2bbox(distance, anchor_points, xywh=True, dim=-1)

Transform distance (ltrb) to box (xywh or xyxy).

Args

NameTypeDescriptionDefault
distancetorch.TensorLeft, top, right, bottom distances from the anchor points with size 4 along dim.required
anchor_pointstorch.TensorAnchor points with size 2 along dim.required
xywhboolWhether to return boxes in xywh format (True) or xyxy format (False).True
dimintDimension along which to split and concatenate.-1

Returns

TypeDescription
torch.TensorDecoded bounding boxes.
GitHubultralytics/utils/tal.py
def dist2bbox(distance, anchor_points, xywh=True, dim=-1):
    """Transform distance (ltrb) to box (xywh or xyxy).

    Args:
        distance (torch.Tensor): Left, top, right, bottom distances from the anchor points with size 4 along dim.
        anchor_points (torch.Tensor): Anchor points with size 2 along dim.
        xywh (bool): Whether to return boxes in xywh format (True) or xyxy format (False).
        dim (int): Dimension along which to split and concatenate.

    Returns:
        (torch.Tensor): Decoded bounding boxes.
    """
    lt, rb = distance.chunk(2, dim)
    x1y1 = anchor_points - lt
    x2y2 = anchor_points + rb
    if xywh:
        c_xy = (x1y1 + x2y2) / 2
        wh = x2y2 - x1y1
        return torch.cat([c_xy, wh], dim)  # xywh bbox
    return torch.cat((x1y1, x2y2), dim)  # xyxy bbox





Function ultralytics.utils.tal.bbox2dist#

def bbox2dist(anchor_points: torch.Tensor, bbox: torch.Tensor, reg_max: int | None = None) -> torch.Tensor

Transform bbox (xyxy) to distance (ltrb).

Args

NameTypeDescriptionDefault
anchor_pointstorch.TensorAnchor points with shape (..., 2).required
bboxtorch.TensorBounding boxes in xyxy format with shape (..., 4).required
reg_maxint, optionalIf provided, distances are clamped to [0, reg_max - 0.01].None

Returns

TypeDescription
torch.TensorLeft, top, right, bottom distances with shape (..., 4).
GitHubultralytics/utils/tal.py
def bbox2dist(anchor_points: torch.Tensor, bbox: torch.Tensor, reg_max: int | None = None) -> torch.Tensor:
    """Transform bbox (xyxy) to distance (ltrb).

    Args:
        anchor_points (torch.Tensor): Anchor points with shape (..., 2).
        bbox (torch.Tensor): Bounding boxes in xyxy format with shape (..., 4).
        reg_max (int, optional): If provided, distances are clamped to [0, reg_max - 0.01].

    Returns:
        (torch.Tensor): Left, top, right, bottom distances with shape (..., 4).
    """
    x1y1, x2y2 = bbox.chunk(2, -1)
    dist = torch.cat((anchor_points - x1y1, x2y2 - anchor_points), -1)
    if reg_max is not None:
        dist = dist.clamp_(0, reg_max - 0.01)  # dist (lt, rb)
    return dist





Function ultralytics.utils.tal.dist2rbox#

def dist2rbox(pred_dist, pred_angle, anchor_points, dim=-1)

Decode predicted rotated bounding box coordinates from anchor points and distribution.

Args

NameTypeDescriptionDefault
pred_disttorch.TensorPredicted left, top, right, bottom distances with shape (bs, h*w, 4).required
pred_angletorch.TensorPredicted angle with shape (bs, h*w, 1).required
anchor_pointstorch.TensorAnchor points with shape (h*w, 2).required
dimint, optionalDimension along which to split.-1

Returns

TypeDescription
torch.TensorPredicted rotated bounding boxes in xywh format (angle excluded) with shape (bs, h*w, 4).
GitHubultralytics/utils/tal.py
def dist2rbox(pred_dist, pred_angle, anchor_points, dim=-1):
    """Decode predicted rotated bounding box coordinates from anchor points and distribution.

    Args:
        pred_dist (torch.Tensor): Predicted left, top, right, bottom distances with shape (bs, h*w, 4).
        pred_angle (torch.Tensor): Predicted angle with shape (bs, h*w, 1).
        anchor_points (torch.Tensor): Anchor points with shape (h*w, 2).
        dim (int, optional): Dimension along which to split.

    Returns:
        (torch.Tensor): Predicted rotated bounding boxes in xywh format (angle excluded) with shape (bs, h*w, 4).
    """
    lt, rb = pred_dist.split(2, dim=dim)
    cos, sin = torch.cos(pred_angle), torch.sin(pred_angle)
    # (bs, h*w, 1)
    xf, yf = ((rb - lt) / 2).split(1, dim=dim)
    x, y = xf * cos - yf * sin, xf * sin + yf * cos
    xy = torch.cat([x, y], dim=dim) + anchor_points
    return torch.cat([xy, lt + rb], dim=dim)





Function ultralytics.utils.tal.rbox2dist#

def rbox2dist(
    target_bboxes: torch.Tensor,
    anchor_points: torch.Tensor,
    target_angle: torch.Tensor,
    dim: int = -1,
    reg_max: int | None = None,
)

Transform rotated bounding box (xywh) to distance (ltrb). This is the inverse of dist2rbox.

Args

NameTypeDescriptionDefault
target_bboxestorch.TensorTarget rotated bounding boxes with shape (bs, h*w, 4), format [x, y, w, h].required
anchor_pointstorch.TensorAnchor points with shape (h*w, 2).required
target_angletorch.TensorTarget angle with shape (bs, h*w, 1).required
dimint, optionalDimension along which to split.-1
reg_maxint, optionalMaximum regression value for clamping.None

Returns

TypeDescription
torch.TensorRotated distance with shape (bs, h*w, 4), format [l, t, r, b].
GitHubultralytics/utils/tal.py
def rbox2dist(
    target_bboxes: torch.Tensor,
    anchor_points: torch.Tensor,
    target_angle: torch.Tensor,
    dim: int = -1,
    reg_max: int | None = None,
):
    """Transform rotated bounding box (xywh) to distance (ltrb). This is the inverse of dist2rbox.

    Args:
        target_bboxes (torch.Tensor): Target rotated bounding boxes with shape (bs, h*w, 4), format [x, y, w, h].
        anchor_points (torch.Tensor): Anchor points with shape (h*w, 2).
        target_angle (torch.Tensor): Target angle with shape (bs, h*w, 1).
        dim (int, optional): Dimension along which to split.
        reg_max (int, optional): Maximum regression value for clamping.

    Returns:
        (torch.Tensor): Rotated distance with shape (bs, h*w, 4), format [l, t, r, b].
    """
    xy, wh = target_bboxes.split(2, dim=dim)
    offset = xy - anchor_points  # (bs, h*w, 2)
    offset_x, offset_y = offset.split(1, dim=dim)
    cos, sin = torch.cos(target_angle), torch.sin(target_angle)
    xf = offset_x * cos + offset_y * sin
    yf = -offset_x * sin + offset_y * cos

    w, h = wh.split(1, dim=dim)
    target_l = w / 2 - xf
    target_t = h / 2 - yf
    target_r = w / 2 + xf
    target_b = h / 2 + yf

    dist = torch.cat([target_l, target_t, target_r, target_b], dim=dim)
    if reg_max is not None:
        dist = dist.clamp_(0, reg_max - 0.01)

    return dist