如何基于旋转框四坐标填充PyTorch张量指定区域?
PyTorch中填充旋转框区域的简便实现方法
这里提供两种实用的实现方式,可根据你的场景选择:
方法一:借助OpenCV多边形填充(快速上手)
这种方式代码简洁直观,适合快速验证需求,无需复杂的张量运算:
import torch import cv2 import numpy as np # 初始化输入张量,示例:batch size=2,尺寸256×256 bs, _, h, w = 2, 1, 256, 256 input_tensor = torch.zeros((bs, 1, h, w)) # 每个batch样本对应的旋转框四个顶点(需保证顶点按顺时针/逆时针顺序排列) rot_boxes = [ np.array([[50,50], [150,30], [170,130], [70,150]], dtype=np.int32), np.array([[100,100], [200,80], [220,180], [120,200]], dtype=np.int32) ] for idx in range(bs): # 创建空白掩码图像 mask = np.zeros((h, w), dtype=np.uint8) # 填充旋转框区域 cv2.fillPoly(mask, [rot_boxes[idx]], 1) # 将掩码赋值给原张量对应位置 input_tensor[idx, 0] = torch.from_numpy(mask)
方法二:纯PyTorch向量化实现(GPU友好)
如果需要在GPU上批量处理、避免CPU-GPU数据传输延迟,推荐这种纯张量运算的方式,核心用射线法判断像素是否在旋转框内:
import torch def point_in_rotated_box(points, rot_boxes): # points: [batch_size, h*w, 2],所有像素的(x,y)坐标 # rot_boxes: [batch_size, 4, 2],每个旋转框的四个顶点 bs, num_points, _ = points.shape num_edges = 4 # 扩展维度实现批量运算 points_expand = points.unsqueeze(2).repeat(1, 1, num_edges, 1) edges_start = rot_boxes.unsqueeze(1).repeat(1, num_points, 1, 1) edges_end = torch.roll(rot_boxes, shifts=-1, dims=1).unsqueeze(1).repeat(1, num_points, 1, 1) # 射线法判断点是否在多边形内 y_cond = (edges_start[..., 1] > points_expand[..., 1]) != (edges_end[..., 1] > points_expand[..., 1]) x_intersect = (points_expand[..., 1] - edges_start[..., 1]) * (edges_end[..., 0] - edges_start[..., 0]) / \ (edges_end[..., 1] - edges_start[..., 1] + 1e-8) + edges_start[..., 0] x_cond = points_expand[..., 0] < x_intersect inside_flag = torch.sum((y_cond & x_cond).int(), dim=-1) % 2 == 1 return inside_flag # 初始化输入张量,支持GPU device = 'cuda' if torch.cuda.is_available() else 'cpu' bs, _, h, w = 2, 1, 256, 256 input_tensor = torch.zeros((bs, 1, h, w), device=device) # 旋转框坐标(格式:[batch_size, 4, 2],x对应宽度维度,y对应高度维度) rot_boxes = torch.tensor([ [[50,50], [150,30], [170,130], [70,150]], [[100,100], [200,80], [220,180], [120,200]] ], device=device, dtype=torch.float32) # 生成所有像素的坐标网格 y_grid, x_grid = torch.meshgrid(torch.arange(h, device=device), torch.arange(w, device=device), indexing='ij') all_pixels = torch.stack([x_grid.flatten(), y_grid.flatten()], dim=-1).unsqueeze(0).repeat(bs, 1, 1) # 计算哪些像素在旋转框内 inside_mask = point_in_rotated_box(all_pixels, rot_boxes).reshape(bs, h, w) # 填充区域为1 input_tensor[inside_mask.unsqueeze(1)] = 1
两种方法对比
- OpenCV方法:代码简单,适合小批量/CPU场景,缺点是需要遍历样本,批量效率较低,依赖第三方库。
- 纯PyTorch方法:支持GPU批量加速,无外部依赖,适合大规模训练/推理场景,代码逻辑稍复杂。
内容的提问来源于stack exchange,提问作者Kami YAN
相关产品推荐
相关产品推荐

