PyTorch中如何指定索引操作对应轴 实现按指定轴提取张量元素
实现方案
你要的沿指定轴提取索引元素的逻辑可以通过两种方式实现,以下是修改后的完整函数:
import torch from torch import Tensor def remove_weaklings(x: Tensor, percentage: float, axis: int) -> Tensor: all_axes = set(range(x.ndim)) - set([axis]) y = x # 对所有非目标轴求和 for a in all_axes: y = y.sum(axis=a, keepdim=True) y = y.squeeze() # 排序得到对应索引 _, idx = torch.sort(y) # 按比例截取保留的索引 idx = idx[:int(percentage * len(idx))] # --------------- 核心实现 二选一即可 --------------- # 方案1:直接调用PyTorch官方封装接口 return torch.index_select(x, dim=axis, index=idx) # 方案2:手动构造切片元组,和你示例中x[:, indices]的写法完全等价 # slicer = tuple([slice(None)] * axis + [idx] + [slice(None)] * (x.ndim - axis - 1)) # return x[slicer]
两种方案说明
torch.index_select是官方原生API,可读性和稳定性更高,适合绝大多数常规场景- 手动构造切片的方式灵活度更高,如果后续需要对其他轴增加额外的索引规则,可以直接在slicer元组里修改对应位置的逻辑
效果验证
你可以用如下测试用例验证效果:
# 测试:输入形状为(2, 3, 4)的三维张量,沿最后一个轴(axis=2)保留50%的元素 x = torch.randn(2, 3, 4) output = remove_weaklings(x, percentage=0.5, axis=2) print(output.shape) # 输出为 torch.Size([2, 3, 2]),符合预期
内容的提问来源于stack exchange,提问作者RafazZ
相关产品推荐
相关产品推荐

