PyTorch批量生成图像空间索引及条件坐标筛选方法
批量张量坐标生成与筛选实现方案
基础需求:生成B×H×W×2全量坐标网格
你最初写的arange拼接逻辑功能可运行,但写法冗余,反复调用unsqueeze+expand容易写错维度,也会增加不必要的视图维护成本,标准实现用torch.meshgrid即可,内存效率更高也更易读:
import torch B, H, W = M.shape # 生成形状为B,H,W,2的LongTensor坐标,最后一维顺序为(行索引, 列索引) coord_grid = torch.stack( torch.meshgrid( torch.arange(H, device=M.device, dtype=torch.long), torch.arange(W, device=M.device, dtype=torch.long), indexing="ij" ), dim=-1 ).unsqueeze(0).expand(B, -1, -1, -1)
实现说明:
indexing="ij"必须显式指定,保证输出坐标顺序是(行号, 列号),避免默认xy索引模式下维度顺序颠倒- 批量维度用
expand实现不会复制实际内存,比逐段拼接expand的方案内存占用低很多 - 生成结果默认就是
LongTensor类型,直接匹配需求,不需要额外做类型转换 - 所有张量初始化时指定和输入M相同的设备,避免CPU/GPU跨设备报错
进阶需求:批量按条件筛选坐标(批量版nonzero)
原生torch.nonzero无法直接输出固定形状的批量结果——因为每个batch内符合筛选条件的元素数量不一定相等,工业界通用实现是按单batch最大符合数做对齐填充,不足的位置填无效值-1,可直接复用的代码如下:
def batch_select_coords(base_tensor: torch.Tensor, cond_mask: torch.Tensor) -> torch.Tensor: """ 批量筛选符合条件的元素坐标 Args: base_tensor: 输入张量,形状为(B, H, W) cond_mask: 布尔筛选掩码,形状和base_tensor一致,True位置为需要保留的元素 Returns: 坐标张量,形状为(B, K, 2),K为单个batch内最多的符合条件元素数,无有效坐标的位置填-1 """ B, H, W = base_tensor.shape # 生成基础坐标网格 grid = torch.stack( torch.meshgrid( torch.arange(H, device=base_tensor.device, dtype=torch.long), torch.arange(W, device=base_tensor.device, dtype=torch.long), indexing="ij" ), dim=-1 ).unsqueeze(0).expand(B, -1, -1, -1) # 统计每个batch的有效元素数,计算对齐长度K per_batch_valid_num = cond_mask.flatten(1).sum(dim=1) max_k = per_batch_valid_num.max().item() # 处理全零掩码的边界情况 if max_k == 0: return torch.full((B, 0, 2), -1, dtype=torch.long, device=base_tensor.device) # 初始化输出张量,逐batch填充有效坐标 output = torch.full((B, max_k, 2), -1, dtype=torch.long, device=base_tensor.device) for batch_idx in range(B): curr_valid_num = per_batch_valid_num[batch_idx].item() curr_coords = grid[batch_idx][cond_mask[batch_idx]] output[batch_idx, :curr_valid_num] = curr_coords return output
常用场景示例(掩码质心计算)
# 阈值化得到浮点掩码的布尔掩码 bool_mask = M > 0.5 # 提取所有掩码内的坐标 coords = batch_select_coords(M, bool_mask) # 过滤填充值计算质心 valid_mask = coords[..., 0] != -1 centroids = (coords * valid_mask.unsqueeze(-1)).sum(dim=1) / valid_mask.sum(dim=1, keepdim=True).clamp(min=1) # 输出centroids形状为(B,2),对应每个batch掩码的(行方向中心, 列方向中心)
如果业务场景下能保证所有batch的有效元素数完全一致(比如提前做了固定点数采样),可以去掉填充逻辑,直接堆叠有效坐标即可得到严格B×K×2无填充值的结果。
内容的提问来源于stack exchange,提问作者ysig
相关产品推荐
相关产品推荐

