张量维度内索引的最简实现及通用索引张量获取方法问询
这个问题戳中了很多人在处理动态维度张量时的痛点——手动重塑、拼接索引确实太繁琐了!好在不管你用PyTorch还是TensorFlow,都有更简洁的实现方式,甚至完全不需要硬编码维度信息。
一、直接用框架原生的gather函数(最优雅)
你的需求本质上是按batch维度分组,对最后一维进行索引,框架的gather函数天生就支持这种场景,完全不需要手动处理形状或生成额外索引:
以PyTorch为例
假设x的形状是(*batch_dims, C)(前k维是batch维度,最后一维是待索引的维度),y的形状是(*batch_dims, K)(每个元素是0~C-1之间的索引值),直接一行代码就能得到目标张量z:
z = x.gather(dim=-1, index=y)
比如x = torch.randn(2, 3, 4)(batch_dims为[2,3],C=4),y = torch.randint(0,4,(2,3,5))(K=5),执行后z的形状是(2,3,5),且完全满足z[i,j,k] = x[i,j,y[i,j,k]]的要求。
以TensorFlow为例
TensorFlow 2.0+支持batch_dims参数,用来指定需要按batch分组处理的维度数量(这里就是前k维的数量,即len(x.shape)-1):
z = tf.gather(x, y, axis=-1, batch_dims=len(x.shape)-1)
效果和PyTorch版本完全一致,自动匹配batch维度并完成索引。
二、动态生成类似np.indices的索引张量
如果确实需要生成前k维的索引张量(比如某些特殊场景下必须用gather_nd),也可以动态实现,无需知晓张量的具体秩:
PyTorch实现
利用torch.meshgrid动态生成各维度的索引,再拼接成最终的索引张量:
def get_batch_indices(x): # x是输入张量,前k维为batch维度 batch_dims = x.shape[:-1] # 生成每个维度的索引序列 axes = [torch.arange(d, device=x.device) for d in batch_dims] # 生成网格索引(indexing='ij'保证和np.indices行为一致) grid = torch.meshgrid(axes, indexing='ij') # 拼接成形状为(*batch_dims, k)的索引张量 return torch.stack(grid, dim=-1)
比如输入x形状为(2,3,4),返回的索引张量形状为(2,3,2),每个位置[i,j]的值为[i,j],完美对应前k维的索引。
TensorFlow实现
类似地,用tf.meshgrid实现:
def get_batch_indices(x): batch_dims = x.shape[:-1] axes = [tf.range(d) for d in batch_dims] grid = tf.meshgrid(*axes, indexing='ij') return tf.stack(grid, axis=-1)
如果要结合gather_nd使用,只需要把这个索引张量和y的扩展维度拼接即可:
# PyTorch示例 batch_indices = get_batch_indices(x) # 将y扩展为(*batch_dims, K, 1),再和batch_indices拼接成(*batch_dims, K, k+1) full_indices = torch.cat([batch_indices.unsqueeze(-2).repeat(1,1,y.shape[-1],1), y.unsqueeze(-1)], dim=-1) z = torch.gather_nd(x, full_indices)
不过还是那句话——如果只是满足你的核心需求,直接用gather要简洁得多!
内容的提问来源于stack exchange,提问作者deasmhumnha

