如何高效计算候选旋转矩形与目标旋转矩形的批量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
相关产品推荐
相关产品推荐

