求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
相关产品推荐
相关产品推荐

