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

如何用input()测试基于bAbI数据集训练的Keras记忆网络模型

我懂你现在的困扰——好不容易把基于bAbI数据集的Keras记忆网络跑通了,想改成能实时接收用户输入聊天的版本却卡壳了,不管是用input()还是读txt文件都没成功。别着急,咱们从记忆网络的输入逻辑入手一步步解决问题。

核心原因:你的输入没匹配模型的训练格式

bAbI数据集的记忆网络模型,训练时接收的是经过编码、标准化长度的文本序列(比如故事上下文、问题都被转换成了整数编码,并且填充到固定长度)。直接用input()输入的原始文本,或者从txt里读的未处理内容,模型根本“看不懂”。

解决步骤:从训练到预测的完整适配

1. 训练阶段:保存词汇映射表(Tokenizer)

首先,在你训练模型的代码里,肯定有对文本做分词、转整数的步骤(一般用Keras的Tokenizer)。训练完必须保存这个Tokenizer,这样新输入才能用和训练时完全一样的词汇映射规则。

在训练代码末尾添加这段:

from keras.preprocessing.text import Tokenizer
import pickle

# 假设你已经用Tokenizer拟合了所有训练数据(故事、问题、答案)
# tokenizer = Tokenizer(filters='')
# tokenizer.fit_on_texts(all_stories + all_questions + all_answers)

# 保存Tokenizer到本地
with open('babi_tokenizer.pkl', 'wb') as f:
    pickle.dump(tokenizer, f)

# 同时保存训练时用的最大序列长度(非常重要!)
import numpy as np
np.save('max_lengths.npy', [max_story_len, max_question_len])

这里的max_story_len和max_question_len是你训练时为故事、问题设置的最大长度(比如用pad_sequences时指定的maxlen),一定要保存下来。

2. 预测阶段:编写输入预处理函数

写一个函数,把用户输入的原始文本,转换成模型能接受的格式:

from keras.models import load_model
from keras.preprocessing.sequence import pad_sequences
import pickle
import numpy as np

# 加载训练好的模型、Tokenizer和最大长度
model = load_model('memory_network_model.h5')  # 替换成你的模型文件名
with open('babi_tokenizer.pkl', 'rb') as f:
    tokenizer = pickle.load(f)
max_story_len, max_question_len = np.load('max_lengths.npy')

def preprocess_user_input(story_text, question_text):
    # 1. 拆分故事为句子(和训练时的拆分逻辑一致,bAbI里是按句号+空格拆分)
    story_sentences = [s.strip() for s in story_text.split('. ') if s.strip()]
    # 2. 把故事和问题转换成整数序列
    story_seq = tokenizer.texts_to_sequences(story_sentences)
    question_seq = tokenizer.texts_to_sequences([question_text])
    # 3. 拼接故事的所有词为一个长序列(如果你的模型是这样输入的,看原始代码的输入结构)
    story_flat = [word for sent in story_seq for word in sent]
    # 4. 填充到训练时的最大长度
    story_padded = pad_sequences([story_flat], maxlen=max_story_len)
    question_padded = pad_sequences(question_seq, maxlen=max_question_len)
    return story_padded, question_padded

注意:如果你的原始模型输入是多维度的(比如每个故事句子作为单独输入),那要调整预处理逻辑,和训练时的输入结构完全对齐。

3. 实现实时聊天功能(用input())

现在可以写一个聊天循环,用input()获取用户输入,预处理后喂给模型预测,再把预测结果转成文本:

def start_chat():
    print("👋 欢迎和记忆网络聊天!输入'exit'随时退出")
    while True:
        print("\n---")
        story_input = input("请输入上下文故事(句子用句号分隔):")
        if story_input.lower() == 'exit':
            print("再见!")
            break
        question_input = input("请输入你的问题:")
        if question_input.lower() == 'exit':
            print("再见!")
            break
        
        # 预处理输入
        story_padded, question_padded = preprocess_user_input(story_input, question_input)
        # 预测答案
        pred_probs = model.predict([story_padded, question_padded])
        pred_idx = pred_probs.argmax(axis=-1)[0]
        # 把整数索引转回文本
        idx2word = {v: k for k, v in tokenizer.word_index.items()}
        answer = idx2word.get(pred_idx, "抱歉,我没找到答案")
        
        print(f"\n🤖 机器人回答:{answer}")

# 启动聊天
start_chat()

4. 读取txt文件作为输入的解决方法

如果想读取本地txt文件作为上下文故事,只需要把input()替换成文件读取即可:

# 读取txt文件内容
with open('user_story.txt', 'r', encoding='utf-8') as f:
    story_text = f.read().strip()
# 然后把story_text传入preprocess_user_input函数即可
关键注意事项
  • 输入尺寸必须完全匹配:max_story_len和max_question_len一定要和训练时的数值一致,否则模型会报错输入形状不匹配。
  • 对齐预处理逻辑:如果你的原始模型对故事的处理方式是保留句子结构(比如每个句子作为一个输入序列),那预处理时不能直接拼接,要对应处理成二维序列。
  • 词汇表覆盖问题:如果用户输入的词不在训练时的词汇表里,Tokenizer会忽略它,可能导致预测不准,这是小数据集的正常限制,可以考虑扩展词汇表或者做OOV处理。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 08:13:28