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

如何基于旋转框四坐标填充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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 07:13:24