PyTorch如何实现批量gather操作且无需低效反广播
PyTorch无反广播开销实现time维度批量Gather
需求回顾
目标是对形状为[batch, time, feature]的输入张量x,沿time维度用形状为[batch, new_time]的批量索引i采样,得到形状为[batch, new_time, feature]的输出y,取值规则为:
y[b, t', f] = x[b, i[b, t'], f]
TensorFlow中可直接通过带batch_dims参数的tf.gather实现:
y = tf.gather(x, i, axis=1, batch_dims=1)
PyTorch原生接口不需要做低效的反广播/张量复制,即可实现同等性能的操作。很多人误以为需要反广播,是受早期错误教程影响,用了repeat这类真实复制数据的操作构造索引。
最优实现方案
通用兼容写法(支持所有PyTorch 1.0+版本)
直接调用torch.gather,仅需对索引张量做一次最后一维的升维,利用算子内置的广播机制完成计算,全程无真实数据复制,无额外大显存占用:
import torch y = torch.gather( x, dim=1, index=i.unsqueeze(-1) # 索引形状变为[batch, new_time, 1],自动广播到feature维度 )
运行后输出y的形状即为要求的[batch, new_time, feature],计算结果和TensorFlow实现完全对齐。
语义化写法(支持PyTorch 1.9+版本)
高版本PyTorch提供了更直白的沿指定维度采样接口torch.take_along_dim,调用逻辑和上述一致,可读性更强:
y = torch.take_along_dim(x, i.unsqueeze(-1), dim=1)
常见误区说明
之前网上流传的实现之所以效率低,核心是错误构造了索引或输入张量,引入了不必要的开销:
- 用
repeat把索引张量真实复制feature份,会产生batch*new_time*feature大小的额外显存占用,大维度下速度极慢 - 调用
torch.index_select前手动重构输入张量维度,引入不必要的转置、形状重塑开销 - 误用
torch.nn.functional.embedding做批量采样,其底层依赖index_select,同样存在上述重构开销
上述给出的实现仅对索引做一次长度为1的维度扩充,算子内部自动完成广播,索引张量始终保持极小的显存占用,性能和TensorFlow的batch_dims实现基本无差异。
注意:该实现要求索引
i的取值范围在[0, time)区间内,和TensorFlow的索引越界判定逻辑一致。
内容的提问来源于stack exchange,提问作者Frithjof
相关产品推荐
相关产品推荐

