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
相关产品推荐
相关产品推荐

