如何无循环并行对批量张量执行torch.meshgrid操作?
批量无循环实现PyTorch批量meshgrid
需求背景
给定两个形状均为[60, 9]的二维张量x、y,需要实现类meshgrid操作,满足:
- 输出张量
xx、yy形状均为[60, 9, 9] - 对任意批次索引
i,xx[i, :, :]、yy[i, :, :]的计算结果,与直接对一维张量x[i]、y[i]调用torch.meshgrid的输出完全一致 - 禁止使用for循环,保证GPU并行计算效率
实现方法
直接利用PyTorch的广播机制,通过维度扩展即可实现,无任何循环、无额外数据拷贝开销,GPU运行效率拉满:
import torch # 以下实现对齐 torch.meshgrid(..., indexing="ij") 的行为,也是PyTorch官方推荐的显式索引模式 batch_size, seq_len = x.shape xx = x.unsqueeze(dim=-1).expand(batch_size, seq_len, seq_len) yy = y.unsqueeze(dim=1).expand(batch_size, seq_len, seq_len)
逻辑说明
- 单样本维度下,
torch.meshgrid(x[i], y[i], indexing="ij")的本质是:把一维的x[i]转为列向量(形状[9, 1])后沿列方向复制9次得到xx[i],把一维的y[i]转为行向量(形状[1, 9])后沿行方向复制9次得到yy[i] - 批量操作时,
unsqueeze加维度、expand扩展的操作是逐批次独立生效的,完全等价于对每个批次单独做上述变换,没有计算逻辑差异 expand操作仅修改张量的步长元数据,不会复制真实显存数据,开销几乎为0
正确性验证
可以随机抽取批次索引和原生实现做对比,结果完全一致:
# 随机选一个批次测试 test_idx = torch.randint(0, 60, (1,)).item() xx_ref, yy_ref = torch.meshgrid(x[test_idx], y[test_idx], indexing="ij") print(torch.allclose(xx[test_idx], xx_ref)) # 输出 True print(torch.allclose(yy[test_idx], yy_ref)) # 输出 True
如果需要对齐旧版本PyTorch默认的indexing="xy"行为,只需要交换xx和yy的扩展维度即可。
内容的提问来源于stack exchange,提问作者lbwnb123
相关产品推荐
相关产品推荐

