关于PyTorch 1.12中meshgrid函数输出结果的疑问
PyTorch中torch.meshgrid的行为解析
核心逻辑说明
torch.meshgrid的作用是生成网格坐标矩阵,其输出形态由indexing参数控制——PyTorch 1.12版本默认使用indexing='ij'模式,这是你得到当前结果的关键。
针对你的示例拆解
输入x = torch.tensor([1, 2, 3])(长度3),y = torch.tensor([4, 5, 6])(长度3):
- 第一个输出张量:对应
x的网格展开,按行维度重复——每一行是x的单个元素,重复y的长度次,因此得到3行(匹配x的长度)、每行3个相同元素的矩阵:tensor([[1, 1, 1], [2, 2, 2], [3, 3, 3]]) - 第二个输出张量:对应
y的网格展开,按列维度重复——每一行是完整的y数组,重复x的长度次,因此得到3行(匹配x的长度)、每行都是[4,5,6]的矩阵:tensor([[4, 5, 6], [4, 5, 6], [4, 5, 6]])
如何得到你预期的结果
你预期的第二个张量是按行重复y元素的形态,这对应indexing='xy'模式(类似NumPy早期默认行为)。只需显式指定该参数即可:
x = torch.tensor([1, 2, 3]) y = torch.tensor([4, 5, 6]) torch.meshgrid(x, y, indexing='xy')
此时返回的第二个张量为:
tensor([[4, 4, 4], [5, 5, 5], [6, 6, 6]])
关键区别总结
indexing='ij'(PyTorch默认):遵循矩阵索引逻辑,第一个输入对应行维度,第二个对应列维度。indexing='xy':遵循笛卡尔坐标系逻辑,第一个输入对应x轴/列维度,第二个对应y轴/行维度。
内容的提问来源于stack exchange,提问作者SamuelMastrelli
相关产品推荐
相关产品推荐

