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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 01:51:21