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

PyTorch矩阵逐行索引单个元素的无Python层迭代实现方案

可行实现方案(无Python层迭代开销,全部在C层执行)

方案1:使用torch.gather(推荐)

torch.gather是PyTorch原生提供的沿指定维度按索引取数的操作,完全在后端执行没有Python层开销,实现方式如下:

# 索引需要和输入张量维度数对齐,取数后去掉多余的维度即可
daily_expense = price.gather(
    dim=1,
    index=purchased_product_by_day.unsqueeze(-1)
).squeeze(-1)

原理说明:gather要求index的维度数和输入张量一致,因此我们先给形状为(num_days,)的购买索引新增最后一维变成(num_days, 1),取数完成后再去掉多余维度,就能得到形状为(num_days,)的每日支出张量,结果和你原有遍历实现完全一致。


方案2:使用张量形式的高级索引

你原有的写法触发Python层遍历的原因是用了Python列表list(range(num_days)),把行索引换成PyTorch原生张量即可规避:

# 生成和price同设备的行索引,全程在C层执行
row_indices = torch.arange(price.size(0), device=price.device)
daily_expense = price[row_indices, purchased_product_by_day]

该写法和你原有逻辑完全等价,只是把Python列表替换为了PyTorch张量索引,不会产生Python层的循环开销。

两种方案性能差异很小,可根据编码习惯选择,你可以通过torch.allclose验证返回结果和原有实现的一致性。

内容的提问来源于stack exchange,提问作者user118967

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 22:27:00