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

如何从已训练保存的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 16:48:54