如何用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
相关产品推荐
相关产品推荐

