PyTorch经transforms预处理后如何关联原始图像实现结果可视化
视频人体关键点追踪可视化的标准实现方案
核心逻辑不需要逆变换还原预处理后的图像,也不需要修改Dataset默认的返回结构,训练链路完全保持原有实现即可,仅需在推理可视化阶段单独加载原始图像,将模型输出的关键点坐标按预处理缩放规则反向映射到原图坐标系即可。
具体实现步骤
1. 保留原有训练链路逻辑不变
你当前写的transform配置、Dataset定义、DataLoader构建、训练集验证集拆分的代码全部不需要修改,完全不会影响训练流程和模型精度。
2. 推理阶段单独读取原始图像,做坐标反向映射
推理时不要使用经过transform处理的tensor图像做可视化,直接根据样本路径读取原始画质的RGB图像,再把模型输出的关键点坐标映射回原图坐标系即可,映射规则和你在__getitem__里做坐标缩放的逻辑完全对应:
- 若你设置了
normalized=True,模型输出的是0~1区间的归一化坐标,映射公式为:原图x坐标 = 预测x值 * 原始图像宽度原图y坐标 = 预测y值 * 原始图像高度 - 若你设置了
normalized=False,模型输出的是Resize后尺寸下的像素坐标,映射公式为:横向缩放比例 = 原始图像宽度 / Resize目标宽度纵向缩放比例 = 原始图像高度 / Resize目标高度原图x坐标 = 预测x值 * 横向缩放比例原图y坐标 = 预测y值 * 纵向缩放比例
注意:如果你后续在transform里加了保持宽高比的padding操作,映射时需要先减去坐标对应的padding偏移量,再乘缩放比例,你当前的Resize是直接拉伸到固定尺寸,不需要额外处理偏移。
3. 对应推理可视化代码示例
import numpy as np import cv2 import torch from PIL import Image model.eval() with torch.no_grad(): for batch_id, (input_tensor, _) in enumerate(validation_loader): input_tensor = input_tensor.to(device) pred_coords = model(input_tensor) # 模型输出和训练时y_coords格式对齐 # 计算当前batch对应验证集的样本索引范围 batch_start = batch_id * validation_loader.batch_size batch_end = min(batch_start + validation_loader.batch_size, len(validation_set)) # random_split生成的子集通过.indices属性获取对应原始Dataset的索引 batch_original_indices = validation_set.indices[batch_start:batch_end] # 逐样本处理可视化 for sample_idx in range(len(pred_coords)): # 从原始Dataset读取图片路径,加载无损原始图像 img_path, _ = dataset.annotations.iloc[batch_original_indices[sample_idx]].values raw_img = Image.open(img_path).convert("RGB") ori_w, ori_h = raw_img.size target_h, target_w = img_size # 你定义的Resize目标尺寸 # 把预测坐标还原为(关键点数量, 2)的(x,y)格式 kpts = pred_coords[sample_idx].cpu().numpy().reshape(-1, 2) # 坐标映射到原图坐标系 if normalize: kpts[:, 0] = kpts[:, 0] * ori_w kpts[:, 1] = kpts[:, 1] * ori_h else: scale_x = ori_w / target_w scale_y = ori_h / target_h kpts[:, 0] = kpts[:, 0] * scale_x kpts[:, 1] = kpts[:, 1] * scale_y # 在原始图像上绘制关键点 canvas = np.array(raw_img) for (x, y) in kpts.astype(np.int32): cv2.circle(canvas, (x, y), radius=3, color=(0, 255, 0), thickness=-1) # 后续可直接展示、保存canvas,为100%原始画质
原有两种思路的问题说明
- 逆变换还原图像的思路没有实际价值:Resize的插值损失是不可逆的,逆变换得到的图像画质远低于直接读取的原始图,还额外增加逆归一化、逆插值的计算开销。
- 修改
__getitem__返回值的思路不符合PyTorch接口规范:PyTorch默认DataLoader的collate逻辑预期返回(图像, 标签)二元组,新增返回值容易引发兼容问题,实际上通过random_split子集的.indices属性就能直接拿到样本对应原始Dataset的索引,完全不需要修改Dataset返回结构。
内容的提问来源于stack exchange,提问作者Mrofsnart
相关产品推荐
相关产品推荐

