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

如何从任意维度的PyTorch张量特定轴获取最终值?

获取任意维度PyTorch张量指定维度的最后一个值(保持维度数不变)

当你处理固定维度的PyTorch张量时,可以直接通过切片x[:, -1:, :]取指定维度的最后一个值并保持维度数,但面对任意维度的张量时,可以用以下通用方法实现:

方法一:构造通用索引元组

通过遍历张量的所有维度,为目标维度构造-1:的切片,其余维度保持全选(:),最终生成索引元组进行切片:

import torch

def get_last_element_keep_dim(x, target_dim):
    # 构造索引:非目标维度用全选,目标维度取最后一个元素且保留维度
    index = tuple(
        slice(-1, None) if i == target_dim else slice(None)
        for i in range(x.ndim)
    )
    return x[index]

验证示例

  • 对3维张量测试:
x = torch.randn([3, 4, 5])
result = get_last_element_keep_dim(x, 1)
print(result.shape)  # 输出 torch.Size([3, 1, 5]),和手动切片结果一致
  • 对4维张量测试:
x_4d = torch.randn([2, 3, 4, 5])
result_4d = get_last_element_keep_dim(x_4d, 2)
print(result_4d.shape)  # 输出 torch.Size([2, 3, 1, 5])

这种方法完全基于PyTorch的原生切片机制,不需要额外创建张量,效率高且兼容性强,适用于任意维度的张量。

方法二:结合index_select与维度保留

也可以用torch.index_select取出目标维度的最后一个元素,该方法会自动保留维度:

def get_last_element_keep_dim_v2(x, target_dim):
    # 获取目标维度最后一个元素的索引
    last_idx = torch.tensor([x.shape[target_dim] - 1], device=x.device)
    return torch.index_select(x, target_dim, last_idx)

注意:需要确保索引张量的设备与输入张量一致(比如输入在GPU上时,索引也要移到GPU)。


内容的提问来源于stack exchange,提问作者Matt Pitkin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 17:24:59