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

如何用PyTorch搭建的LSTM模型实现单词级下一词预测?

单词级LSTM下一词预测实现方案

一、预测前的预处理

首先要将输入的自然语言句子转换成模型可接受的张量格式,核心要和训练时的预处理逻辑保持一致:

  • 按训练时的规则分词(比如简单空格分割、小写转换、标点处理等)
  • 用训练阶段生成的word2idx字典将单词映射为索引,未登录词(OOV)映射到<UNK>的索引
  • 调整张量形状为模型训练时接受的格式(示例以(seq_len, batch_size)为例,batch_size设为1)
  • 转换为LongTensor并移动到模型所在设备(CPU/GPU)

示例代码:

import torch

# 假设你已具备训练好的词汇映射和模型
word2idx = {"<PAD>":0, "<UNK>":1, "How":2, "are":3, "you":4, "today":5, "doing":6}
idx2word = {v:k for k,v in word2idx.items()}
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

def preprocess_input(sentence, word2idx, device):
    # 分词逻辑需与训练时完全一致
    words = sentence.strip().split()
    # 转换为索引,OOV用<UNK>兜底
    indices = [word2idx.get(word, word2idx["<UNK>"]) for word in words]
    # 调整为(seq_len, 1)的张量格式
    input_tensor = torch.LongTensor(indices).unsqueeze(1).to(device)
    return input_tensor

二、单步下一词预测

给定输入序列,生成最可能的下一个词:

  • 将模型切换为评估模式,关闭梯度计算以节省资源
  • 传入预处理后的张量,获取模型输出
  • 取最后一个时间步的输出,通过argmax得到概率最高的单词索引,再映射回单词

示例代码:

def predict_next_word(model, input_sentence, word2idx, idx2word, device):
    model.eval()
    with torch.no_grad():
        input_tensor = preprocess_input(input_sentence, word2idx, device)
        # 假设模型输出为(outputs, hidden),outputs形状为(seq_len, batch_size, vocab_size)
        outputs, hidden = model(input_tensor)
        # 提取最后一个时间步的预测结果
        last_step_output = outputs[-1, :, :]
        pred_idx = torch.argmax(last_step_output, dim=1).item()
        pred_word = idx2word[pred_idx]
        return f"{input_sentence} {pred_word}"

三、多步连续预测(可选)

如果需要连续生成多个后续词,可循环调用单步预测逻辑,每次将新生成的词追加到输入序列中:

def predict_multiple_words(model, input_sentence, word2idx, idx2word, device, num_words=3):
    model.eval()
    current_sentence = input_sentence
    with torch.no_grad():
        for _ in range(num_words):
            input_tensor = preprocess_input(current_sentence, word2idx, device)
            outputs, hidden = model(input_tensor)
            last_step_output = outputs[-1, :, :]
            pred_idx = torch.argmax(last_step_output, dim=1).item()
            pred_word = idx2word[pred_idx]
            current_sentence += f" {pred_word}"
            # 若训练时定义了终止标记,可在此判断提前停止
            if pred_word == "<EOS>":
                break
        return current_sentence

四、关键注意事项

  • 分词一致性:预测时的分词、大小写、标点处理必须和训练阶段完全一致,否则会出现大量未登录词,影响预测效果
  • 随机性控制(可选):如果不想每次都选择概率最高的词(避免生成内容重复),可加入温度参数调整概率分布:
    temperature = 0.8
    last_step_output = outputs[-1, :, :] / temperature
    probabilities = torch.softmax(last_step_output, dim=1)
    pred_idx = torch.multinomial(probabilities, num_samples=1).item()
    
    温度越高,生成的随机性越强;温度越低,越倾向于选择高概率词
  • 输入维度匹配:确保预测时的张量形状和模型训练时接受的输入形状一致,若训练时用(batch_size, seq_len),则需调整预处理时的unsqueeze位置

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 04:25:16