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

如何将tf.gather_nd(batch_dims=1)转换为纯PyTorch算子?

纯PyTorch实现tf.gather_nd(batch_dims=1)的功能

针对你的需求,这里提供两种纯PyTorch算子实现的方案,均不依赖Python高级索引,适合TensorRT转换:

方案一:维度展平+torch.gather

该方案通过将H、W维度展平为线性索引,再用torch.gather选取元素,逻辑简洁高效:

import torch

# 假设offsets是shape=(1,1,320,256,2)的PyTorch张量
# kpt_inds是shape=(1,k,2)的张量,k为变量
W = offsets.shape[3]  # 获取W维度大小256

# 计算H、W对应的线性索引:h * W + w
linear_inds = kpt_inds[:, :, 0] * W + kpt_inds[:, :, 1]  # shape=(1,k)

# 扩展索引维度以匹配offsets展平后的形状
linear_inds_expanded = linear_inds.unsqueeze(1).unsqueeze(-1)  # shape=(1,1,k,1)

# 将offsets的H、W维度展平为一个维度
offsets_flat = offsets.flatten(2, 3)  # shape=(1,1,320*256,2)

# 使用torch.gather选取元素
gathered_offsets = torch.gather(offsets_flat, dim=2, index=linear_inds_expanded.repeat(1,1,1,2))

# 若需要和原tf结果一致的shape(1,k,2),可挤压多余维度
gathered_offsets = gathered_offsets.squeeze(1)  # shape=(1,k,2)

方案二:两次torch.gather分步处理H、W维度

该方案直接针对H、W维度分步选取,更贴合原tf.gather_nd的索引逻辑:

import torch

# 提取H、W维度的索引
h_inds = kpt_inds[:, :, 0]  # shape=(1,k)
w_inds = kpt_inds[:, :, 1]  # shape=(1,k)

# 扩展H索引维度,匹配offsets的维度结构
h_inds_expanded = h_inds.unsqueeze(1).unsqueeze(-1).unsqueeze(-1)  # shape=(1,1,k,1,1)
# 在H维度上选取元素
gathered_h = torch.gather(offsets, dim=2, index=h_inds_expanded.repeat(1,1,1,256,2))  # shape=(1,1,k,256,2)

# 扩展W索引维度,匹配gathered_h的维度结构
w_inds_expanded = w_inds.unsqueeze(1).unsqueeze(2).unsqueeze(-1)  # shape=(1,1,1,k,1)
# 在W维度上选取元素
gathered_offsets = torch.gather(gathered_h, dim=3, index=w_inds_expanded.repeat(1,1,1,1,2))  # shape=(1,1,k,k,2)

# 清理多余维度,得到目标shape(1,k,2)
gathered_offsets = gathered_offsets[:, :, :, 0, :].squeeze(1)

两种方案均使用纯PyTorch内置算子,避免了Python高级索引的不确定性,可安全转换为TensorRT引擎。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 23:18:42