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
相关产品推荐
相关产品推荐

