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

求PyTorch中在有效掩码内高效生成固定数量连续均匀随机2D点的方案

在有效掩码内高效生成固定数量连续型均匀点(PyTorch适配)

你的当前方案通过生成冗余点再过滤的方式虽然可行,但在有效区域占比低时会浪费大量计算资源。下面是针对PyTorch训练循环优化的高效方案,直接在有效区域内采样,无需后续过滤:

核心思路

  • 提取掩码中所有有效像素的坐标
  • 从有效像素中随机选择目标数量的样本
  • 为每个样本添加亚像素级随机偏移,得到连续型坐标
  • 转换为PyTorch grid_sample 要求的[-1, 1]坐标范围

实现代码

import torch
import cv2
import numpy as np
import torch.nn.functional as F

def sample_valid_points(valid_mask_tensor, num_points):
    """
    在有效掩码内生成固定数量的连续型均匀分布点
    Args:
        valid_mask_tensor: 二值掩码张量,形状(H, W),有效区域值>0.5
        num_points: 目标生成点数量
    Returns:
        points: 形状(num_points, 2)的张量,坐标范围[-1, 1],适配grid_sample
    """
    # 获取所有有效像素的(y, x)坐标
    valid_yx = torch.nonzero(valid_mask_tensor > 0.5, as_tuple=False)
    
    # 随机选择num_points个有效像素
    selected_idx = torch.randint(0, valid_yx.shape[0], (num_points,))
    selected_yx = valid_yx[selected_idx].float()
    
    # 添加亚像素偏移:让点落在像素内的连续区域(范围[-0.5, 0.5])
    subpixel_offset = torch.rand(num_points, 2) - 0.5
    continuous_yx = selected_yx + subpixel_offset
    
    # 转换为(x, y)顺序,并归一化到[-1, 1]
    H, W = valid_mask_tensor.shape
    continuous_xy = continuous_yx.flip(1)  # 交换y和x,适配grid_sample的坐标顺序
    points = (continuous_xy / torch.tensor([W-1, H-1], dtype=torch.float32)) * 2 - 1
    
    return points

# 测试示例
if __name__ == "__main__":
    # 创建旋转后的示例掩码
    valid_mask = np.ones((300, 400), dtype=np.uint8) * 255
    center = (valid_mask.shape[1]//2, valid_mask.shape[0]//2)
    rotation_matrix = cv2.getRotationMatrix2D(center, 30, 1.0)
    rotated_mask = cv2.warpAffine(valid_mask, rotation_matrix, (valid_mask.shape[1], valid_mask.shape[0]), flags=cv2.INTER_NEAREST)
    
    # 转为PyTorch张量
    mask_tensor = torch.tensor(rotated_mask, dtype=torch.float32) / 255
    
    # 生成100个有效点
    target_points = 100
    points = sample_valid_points(mask_tensor, target_points)
    
    # 验证所有点都在有效区域(可选)
    sampled = F.grid_sample(mask_tensor.unsqueeze(0).unsqueeze(0), points.unsqueeze(0).unsqueeze(0), 
                           mode='bilinear', padding_mode='zeros', align_corners=False)
    assert torch.all(sampled.squeeze() > 0.5), "生成的点存在无效区域"
    
    # 可视化结果
    mask_bgr = cv2.cvtColor(rotated_mask, cv2.COLOR_GRAY2BGR)
    for pt in points.numpy():
        x, y = pt
        img_x = int((x + 1) * mask_bgr.shape[1] / 2)
        img_y = int((y + 1) * mask_bgr.shape[0] / 2)
        cv2.circle(mask_bgr, (img_x, img_y), 2, (0, 255, 0), -1)
    
    cv2.imshow('Valid Points', mask_bgr)
    cv2.waitKey(0)
    cv2.destroyAllWindows()

方案优势

  • 零冗余计算:直接从有效区域采样,避免生成大量无效点再过滤的开销
  • GPU友好:全程使用PyTorch张量操作,可无缝迁移到GPU运行,适配训练循环
  • 均匀分布保证:有效像素选择和亚像素偏移均为均匀随机,生成的点在有效区域内均匀分布

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 20:14:51