You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何高效计算候选旋转矩形与目标旋转矩形的批量IoU?

批量计算旋转矩形与目标矩形的可微分IoU(PyTorch张量版)

已用Shapely实现两个旋转矩形(4个2D点组成)的IoU计算,但需要批量计算候选旋转矩形列表rec_cands(PyTorch张量,形状[N,4,2])与单个目标旋转矩形rec_tgt(形状[4,2])的IoU。当前用for循环实现效率低,且Shapely不支持张量批量操作,需求是纯张量操作、无需类型转换、平滑可微分的方案。

解决方案

直接基于PyTorch实现分离轴定理(Separating Axis Theorem, SAT)的批量IoU计算,SAT是计算凸多边形交集的经典方法,完全可微分,适合张量批量处理。


1. 从矩形点集提取参数化表示

先将4个点组成的矩形转化为参数形式(中心坐标、半长、半宽、旋转角度),方便批量处理:

import torch

def rect_points_to_params(rects):
    """
    将旋转矩形的点集转化为参数化表示
    参数:
        rects: 形状为[N,4,2](候选列表)或[4,2](目标)的张量
    返回:
        cx, cy: 中心坐标,形状[N]或标量
        a, b: 半长、半宽,形状[N]或标量
        theta: 旋转角度(弧度),形状[N]或标量
    """
    # 计算中心
    center = rects.mean(dim=-2)
    cx, cy = center[...,0], center[...,1]
    
    # 计算所有边向量并筛选长/短轴
    edges = rects[:, 1:] - rects[:, :-1]
    edges = torch.cat([edges, rects[:, [0]] - rects[:, [-1]]], dim=-2)
    edge_lengths = torch.norm(edges, dim=-1)
    long_edge_idx = torch.argmax(edge_lengths, dim=-1)
    long_edge = edges[torch.arange(edges.shape[0]), long_edge_idx]
    short_edge_idx = (long_edge_idx + 1) % 4
    short_edge = edges[torch.arange(edges.shape[0]), short_edge_idx]
    
    # 半长半宽
    a = torch.norm(long_edge, dim=-1) / 2
    b = torch.norm(short_edge, dim=-1) / 2
    
    # 旋转角度(长轴与x轴夹角)
    theta = torch.atan2(long_edge[...,1], long_edge[...,0])
    
    # 单个目标矩形调整为标量
    if rects.dim() == 2:
        cx, cy, a, b, theta = cx.squeeze(), cy.squeeze(), a.squeeze(), b.squeeze(), theta.squeeze()
    
    return cx, cy, a, b, theta

2. 批量实现分离轴定理计算交集面积

利用SAT批量判断矩形是否分离,并对相交的矩形计算精确交集面积:

def sat_intersection_area(cx1, cy1, a1, b1, theta1, cx2, cy2, a2, b2, theta2):
    """
    用分离轴定理批量计算旋转矩形的交集面积
    参数:
        cx1, cy1, a1, b1, theta1: 候选矩形的参数,形状[N]
        cx2, cy2, a2, b2, theta2: 目标矩形的参数,标量
    返回:
        intersection_area: 每个候选矩形与目标的交集面积,形状[N]
    """
    # 生成所有需要测试的分离轴(两个矩形的四条边法向量)
    theta1_plus_90 = theta1 + torch.pi/2
    axes1 = torch.stack([
        torch.cos(theta1), torch.sin(theta1),
        torch.cos(theta1_plus_90), torch.sin(theta1_plus_90)
    ], dim=-1).reshape(-1, 2, 2)  # [N,2,2]
    
    theta2_plus_90 = theta2 + torch.pi/2
    axes2 = torch.tensor([
        [torch.cos(theta2), torch.sin(theta2)],
        [torch.cos(theta2_plus_90), torch.sin(theta2_plus_90)]
    ]).unsqueeze(0).repeat(cx1.shape[0], 1, 1)  # [N,2,2]
    
    axes = torch.cat([axes1, axes2], dim=1)  # [N,4,2]
    
    # 计算矩形在轴上的投影区间
    def project(cx, cy, a, b, theta, axes):
        corners_rel = torch.tensor([[a,b], [-a,b], [-a,-b], [a,-b]]).unsqueeze(0).repeat(cx.shape[0],1,1)
        rot_mat = torch.stack([torch.cos(theta), -torch.sin(theta), torch.sin(theta), torch.cos(theta)], dim=-1).reshape(-1,2,2)
        corners_rel_rot = torch.bmm(corners_rel, rot_mat)
        projections = torch.bmm(corners_rel_rot, axes.transpose(1,2))
        proj_min = projections.min(dim=1)[0]
        proj_max = projections.max(dim=1)[0]
        center_proj = cx.unsqueeze(1)*axes[...,0] + cy.unsqueeze(1)*axes[...,1]
        proj_min += center_proj
        proj_max += center_proj
        return proj_min, proj_max
    
    proj1_min, proj1_max = project(cx1, cy1, a1, b1, theta1, axes)
    proj2_min, proj2_max = project(
        cx2.unsqueeze(0).repeat(cx1.shape[0]), 
        cy2.unsqueeze(0).repeat(cx1.shape[0]), 
        a2.unsqueeze(0).repeat(cx1.shape[0]), 
        b2.unsqueeze(0).repeat(cx1.shape[0]), 
        theta2.unsqueeze(0).repeat(cx1.shape[0]), 
        axes
    )
    
    # 检查是否存在分离轴
    overlap_min = torch.max(proj1_min, proj2_min)
    overlap_max = torch.min(proj1_max, proj2_max)
    has_separating_axis = torch.any(overlap_min > overlap_max, dim=1)
    intersection_area = torch.zeros_like(cx1)
    
    # 对无分离轴的矩形计算交集面积(批量Sutherland-Hodgman裁剪)
    def get_rect_corners(cx, cy, a, b, theta):
        corners_rel = torch.tensor([[a,b], [-a,b], [-a,-b], [a,-b]]).unsqueeze(0).repeat(cx.shape[0],1,1)
        rot_mat = torch.stack([torch.cos(theta), -torch.sin(theta), torch.sin(theta), torch.cos(theta)], dim=-1).reshape(-1,2,2)
        corners = torch.bmm(corners_rel, rot_mat) + torch.stack([cx, cy], dim=-1).unsqueeze(1)
        return corners
    
    corners1 = get_rect_corners(cx1, cy1, a1, b1, theta1)
    corners2 = get_rect_corners(
        cx2.unsqueeze(0).repeat(cx1.shape[0]), 
        cy2.unsqueeze(0).repeat(cx1.shape[0]), 
        a2.unsqueeze(0).repeat(cx1.shape[0]), 
        b2.unsqueeze(0).repeat(cx1.shape[0]), 
        theta2.unsqueeze(0).repeat(cx1.shape[0])
    )
    
    def polygon_intersection_area(poly1, poly2):
        def clip_polygon(subject, clip):
            def inside(p, edge_start, edge_end):
                edge_vec = edge_end - edge_start
                p_vec = p - edge_start
                cross = edge_vec[...,0]*p_vec[...,1] - edge_vec[...,1]*p_vec[...,0]
                return cross >= 0
            
            def intersect(p1, p2, edge_start, edge_end):
                den = (edge_end[...,0]-edge_start[...,0])*(p2[...,1]-p1[...,1]) - (edge_end[...,1]-edge_start[...,1])*(p2[...,0]-p1[...,0])
                t_num = (edge_start[...,0]-p1[...,0])*(p2[...,1]-p1[...,1]) - (edge_start[...,1]-p1[...,1])*(p2[...,0]-p1[...,0])
                t = t_num / den
                x = p1[...,0] + t*(p2[...,0]-p1[...,0])
                y = p1[...,1] + t*(p2[...,1]-p1[...,1])
                return torch.stack([x,y], dim=-1)
            
            output = subject
            for i in range(clip.shape[1]):
                edge_start = clip[:, i]
                edge_end = clip[:, (i+1)%clip.shape[1]]
                input_poly = output
                output = []
                s = input_poly[:, -1]
                for j in range(input_poly.shape[1]):
                    e = input_poly[:, j]
                    if inside(e, edge_start, edge_end):
                        if not inside(s, edge_start, edge_end):
                            output.append(intersect(s, e, edge_start, edge_end))
                        output.append(e)
                    elif inside(s, edge_start, edge_end):
                        output.append(intersect(s, e, edge_start, edge_end))
                    s = e
                if not output:
                    return torch.zeros_like(input_poly[:,0])
                output = torch.stack(output, dim=1)
            if output.shape[1] < 3:
                return torch.zeros(output.shape[0])
            x = output[...,0]
            y = output[...,1]
            area = 0.5 * torch.abs(torch.sum(x[:,:-1]*y[:,1:] - x[:,1:]*y[:,:-1], dim=1))
            return area
        
        area1 = clip_polygon(poly1, poly2)
        area2 = clip_polygon(poly2, poly1)
        return torch.max(area1, area2)
    
    valid_mask = ~has_separating_axis
    if torch.any(valid_mask):
        intersection_area[valid_mask] = polygon_intersection_area(corners1[valid_mask], corners2[valid_mask])
    
    return intersection_area

3. 批量计算IoU

组合上述模块,完成最终的IoU批量计算:

def batch_rotated_rect_iou(rec_cands, rec_tgt):
    """
    批量计算候选旋转矩形列表与目标旋转矩形的IoU
    参数:
        rec_cands: 候选矩形列表,形状[N,4,2]的PyTorch张量
        rec_tgt: 目标矩形,形状[4,2]的PyTorch张量
    返回:
        iou: 每个候选矩形与目标的IoU,形状[N]的PyTorch张量
    """
    # 提取参数
    cx1, cy1, a1, b1, theta1 = rect_points_to_params(rec_cands)
    cx2, cy2, a2, b2, theta2 = rect_points_to_params(rec_tgt)
    
    # 计算各自面积
    area1 = 4 * a1 * b1
    area2 = 4 * a2 * b2
    
    # 计算交集面积
    inter_area = sat_intersection_area(cx1, cy1, a1, b1, theta1, cx2, cy2, a2, b2, theta2)
    
    # 计算IoU并处理数值稳定性
    union_area = area1 + area2 - inter_area
    iou = inter_area / union_area
    iou = torch.where(union_area == 0, torch.zeros_like(iou), iou)
    
    return iou

核心特性

  • 纯张量操作:全程基于PyTorch张量,支持GPU加速,无需类型转换。
  • 可微分:所有运算采用PyTorch原生可微分算子,可直接用于反向传播训练。
  • 高效批量处理:针对候选列表的批量维度[N]优化,彻底避免for循环。
  • 数值稳定性:处理了矩形分离、面积为0等边界情况。

内容的提问来源于stack exchange,提问作者mr MR

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 22:17:11