PyTorch如何实现类似NumPy的按索引为张量赋值操作
PyTorch张量的索引赋值方法
PyTorch的基础索引赋值逻辑和NumPy高度一致,你熟悉的NumPy索引写法几乎可以无缝用到PyTorch张量上。
直接对齐NumPy写法的基础实现
用torch.zeros创建全零张量后,直接通过和NumPy完全相同的索引规则定位位置,传入对应长度的赋值序列即可完成填充,参考代码如下:
import torch # 创建和NumPy示例形状、类型一致的全零张量 tensor = torch.zeros((10, 8, 3), dtype=torch.float32) for n in range(10): for k in range(4): # 按索引定位后直接赋值,和NumPy写法无本质区别 tensor[n, k, :] = torch.tensor([x, y, -2]) # x、y为每次循环生成的不同值 tensor[n, 4 + k, :] = torch.tensor([x, y, 0.4])
使用时注意几个细节:
- 赋值时右侧的元素数量必须和索引选中区域的最后一维长度匹配,上面示例中最后一维长度为3,因此传入3个数值即可;直接传Python列表
[x,y,-2]也可运行,显式转为同dtype的张量可以避免隐式类型转换带来的警告或精度问题。 - 这种切片式的基础索引返回的是原张量的视图,赋值操作会直接修改原张量的存储值,不会生成新的张量,行为和NumPy完全一致。
更高效的向量化赋值写法
和NumPy一样,PyTorch更推荐尽量避免Python层的嵌套循环,用向量化操作完成批量赋值,运行效率会有明显提升,示例如下:
import torch tensor = torch.zeros((10, 8, 3), dtype=torch.float32) # 假设x、y是提前计算好的、形状为(10,4)的张量,对应所有n、k组合下的取值 x = torch.randn(10, 4, dtype=torch.float32) y = torch.randn(10, 4, dtype=torch.float32) # 批量填充k从0到3的位置 tensor[:, :4, 0] = x tensor[:, :4, 1] = y tensor[:, :4, 2] = -2 # 批量填充k从4到7的位置 tensor[:, 4:, 0] = x tensor[:, 4:, 1] = y tensor[:, 4:, 2] = 0.4
踩坑提示:如果使用高级索引(传入整数索引张量、布尔掩码做筛选),索引取出的内容是原张量的副本,不要先把选中的片段赋值给单独变量再修改,那样不会改动原张量;但直接对索引位置赋值是可以正常生效的,比如
tensor[tensor < 0] = 0这类掩码赋值操作,可以正常把张量里所有负值替换为0。
内容的提问来源于stack exchange,提问作者xc-2021
相关产品推荐
相关产品推荐

