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

PyTorch是否有np.ix_等效实现?如何实现张量高级索引?

PyTorch中布尔与数值索引组合的修改实现

完全可以实现类似NumPy的索引与修改操作,不管是CPU还是GPU张量都支持。下面是对应NumPy示例的PyTorch实现方式:

基础CPU张量示例

import torch

# 创建目标张量
x = torch.arange(12).reshape(3, 4)
print(x)
# 输出:
# tensor([[ 0,  1,  2,  3],
#         [ 4,  5,  6,  7],
#         [ 8,  9, 10, 11]])

# 定义行布尔索引和列数值索引
row_mask = torch.tensor([False, True, True])
col_indices = torch.tensor([0, 3])

# 将布尔索引转换为对应的行位置索引
row_indices = torch.where(row_mask)[0]

# 生成网格索引(对应NumPy的np.ix_逻辑)
rows, cols = torch.meshgrid(row_indices, col_indices, indexing='ij')

# 执行赋值操作
x[rows, cols] = torch.tensor([[1, 2], [3, 4]])
print(x)
# 输出:
# tensor([[ 0,  1,  2,  3],
#         [ 1,  5,  6,  2],
#         [ 3,  9, 10,  4]])

GPU张量适配

如果要处理GPU张量,只需将所有相关张量转移到GPU设备即可,逻辑和CPU完全一致:

# 检查GPU可用性
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# 将张量移至GPU
x_gpu = x.to(device)
row_mask_gpu = row_mask.to(device)
col_indices_gpu = col_indices.to(device)

# 重复索引与赋值流程
row_indices_gpu = torch.where(row_mask_gpu)[0]
rows_gpu, cols_gpu = torch.meshgrid(row_indices_gpu, col_indices_gpu, indexing='ij')
x_gpu[rows_gpu, cols_gpu] = torch.tensor([[1, 2], [3, 4]]).to(device)

print(x_gpu)
# 输出结果与CPU版本一致,仅设备为GPU

补充说明

PyTorch的torch.meshgrid配合indexing='ij'参数,和NumPy的np.ix_行为完全匹配,能够生成用于多维索引的网格坐标。这种方式既支持布尔索引与数值索引的组合,也能保证赋值操作的维度对齐,避免广播错误。

内容的提问来源于stack exchange,提问作者H.Rappeport

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 19:22:52