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

