如何从已训练保存的PyTorch seq2seq手语识别模型获取预测输出
PyTorch Seq2Seq手语识别模型推理操作流程
你已经完成模型加载并切换到eval()模式,后续推理按以下步骤操作即可,核心要求是所有预处理、解码逻辑必须和训练阶段完全对齐,否则会出现预测结果完全错乱的问题。
1. 视频输入预处理
推理的第一优先级是对齐训练时的视频处理规则,通用处理逻辑如下,所有参数(采样帧数、分辨率、归一化数值)必须和训练时使用的参数完全一致:
- 逐帧读取待测试视频,按训练时的采样策略抽固定数量的帧(通常是均匀采样,帧数不足时循环补帧、帧数过多时等间隔抽帧)
- 对每帧做和训练一致的图像变换:统一转RGB格式、resize到指定输入分辨率、转Tensor、用数据集的均值方差做归一化
- 按模型要求的维度组装Tensor:单视频推理需要补充batch维度,最终输入形状一般为
[1, 序列长度, 通道数, 帧高度, 帧宽度] - 如果模型训练时额外输入了光流、手部关键点等特征,推理阶段需要同步生成对应特征,和视频Tensor按训练时的规则拼接
通用预处理参考代码:
import cv2 import torch from torchvision import transforms # 以下参数请替换为你训练时实际使用的数值 SAMPLE_FRAMES = 32 INPUT_IMG_SIZE = (224, 224) DATASET_MEAN = [0.485, 0.456, 0.406] DATASET_STD = [0.229, 0.224, 0.225] def process_video(video_path): cap = cv2.VideoCapture(video_path) raw_frames = [] while cap.isOpened(): ret, frame = cap.read() if not ret: break # OpenCV默认读取BGR格式,必须转成和训练一致的RGB raw_frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) cap.release() # 均匀采样固定数量的帧 frame_count = len(raw_frames) sample_indices = torch.linspace(0, frame_count-1, SAMPLE_FRAMES).long() sampled_frames = [raw_frames[i] for i in sample_indices] # 帧变换 frame_transform = transforms.Compose([ transforms.ToPILImage(), transforms.Resize(INPUT_IMG_SIZE), transforms.ToTensor(), transforms.Normalize(DATASET_MEAN, DATASET_STD) ]) frame_tensor = torch.stack([frame_transform(f) for f in sampled_frames]) # 补充batch维度,匹配模型输入形状 return frame_tensor.unsqueeze(0)
2. 模型前向推理
推理阶段必须用torch.no_grad()包裹前向过程,避免PyTorch缓存梯度占用显存、拖慢推理速度。
注意Seq2Seq结构和普通单标签分类模型不同:训练时通常会开启教师强制(Teacher Forcing)加速收敛,推理时需要关闭这个逻辑,大部分实现会在forward函数里预留参数(比如传入mode='inference'、teacher_forcing_ratio=0)切换模式,部分实现还支持传入beam_size参数开启束搜索提升解码准确率,你可以对照自己模型的forward函数定义确认传参。
推理参考代码:
# 优先用GPU推理,没有GPU自动切CPU device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = model.to(device) # 处理测试视频 input_tensor = process_video("test_video.mp4").to(device) with torch.no_grad(): # 传参请对齐你自己的模型定义,示例为常见调用方式 # pred_logits = model(input_tensor, mode='inference', beam_size=5) pred_logits = model(input_tensor)
3. 输出结果转自然语言句子
模型直接输出的是词表维度的概率logits,需要转成可读文本,步骤如下:
- 对logits取每个序列位置概率最大的token ID,形状为
[batch_size, 序列长度] - 用训练时保存的
id2word映射字典(token ID到词/字的对应关系)把ID转成实际文本 - 过滤掉训练时加入的特殊token:包括起始符
<sos>、结束符<eos>、填充符<pad>,遇到结束符时直接终止后续拼接 - 按语言规则拼接剩余token得到最终句子:英文单词之间加空格,中文直接按顺序拼接即可
结果解码参考代码:
# id2word为你训练时生成的映射字典,格式示例:{0:"<pad>", 1:"<sos>", 2:"<eos>", 3:"Hello", 4:"how"...} pred_token_ids = torch.argmax(pred_logits, dim=-1)[0].cpu().numpy() result_tokens = [] for token_id in pred_token_ids: current_word = id2word[token_id] if current_word == "<eos>": break if current_word not in ["<sos>", "<pad>"]: result_tokens.append(current_word) # 英文用空格拼接,中文可直接用"".join(result_tokens) final_sentence = " ".join(result_tokens) print("模型预测结果:", final_sentence)
常见坑点提醒:如果预测结果全是乱码或者重复无意义内容,优先检查预处理逻辑是否和训练一致、推理时是否正确关闭了教师强制、词表映射是否对应正确。
内容的提问来源于stack exchange,提问作者Israel mehari
相关产品推荐
相关产品推荐

