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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 19:36:06