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

PyTorch中如何交换检测输出张量的x与y坐标顺序

PyTorch重排检测输出张量列顺序实现

需求说明

  • 输入:形状为100*6的模型输出张量,列顺序为xmin、ymin、xmax、ymax、conf(置信度)、class(类别)
  • 输出:同形状张量,列顺序调整为ymin、xmin、ymax、xmax、conf、class
  • 转换示例:
# 输入
x = [[1,2,3,4,5,6],
     [7,8,9,10,11,12]]
# 期望输出
y = [[2,1,4,3,5,6],
     [8,7,10,9,11,12]]

实现方案

直接通过PyTorch的维度索引对列维度按目标顺序选取即可,操作高效且不会改变张量形状,代码如下:

import torch

# x为原始输入张量,形状支持(N, 6),N可以是任意值(本题中为100)
# 原始列索引映射:0:xmin, 1:ymin, 2:xmax, 3:ymax, 4:conf, 5:class
# 按目标顺序指定列索引即可完成重排
y = x[:, [1, 0, 3, 2, 4, 5]]

索引规则说明

  • 索引中第一个:表示保留行维度的所有元素,不对行做筛选
  • 列维度传入的索引列表[1, 0, 3, 2, 4, 5]对应新张量每一列取原始张量的列位置:
    • 新第0列取原始第1列(ymin)
    • 新第1列取原始第0列(xmin)
    • 新第2列取原始第3列(ymax)
    • 新第3列取原始第2列(xmax)
    • 新第4、5列直接取原始对应位置的conf、class列,顺序不变

示例验证

用题目给出的样例输入运行代码:

x = torch.tensor([[1,2,3,4,5,6],
                  [7,8,9,10,11,12]])
y = x[:, [1, 0, 3, 2, 4, 5]]
print(y)

得到的输出和期望结果完全一致:

tensor([[ 2,  1,  4,  3,  5,  6],
        [ 8,  7, 10,  9, 11, 12]])

扩展兼容写法

如果你的张量带batch维度(比如形状为(batch_size, 100, 6)),可以用省略号自动匹配前面的维度,不需要修改索引逻辑:

# 兼容带batch维度、不带batch维度的所有输入
y = x[..., [1, 0, 3, 2, 4, 5]]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.09 16:15:42