基于TensorFlow与LSTM的类20问系统技术实现问询
TensorFlow实现20问风格目标预测系统:核心方案解析
我之前做过类似的20问式目标分类系统,结合LSTM处理序列对话的特性,给你梳理下这三个核心问题的落地思路:
1. 训练数据生成与迭代轮次建议
数据生成方法
要生成足量的训练数据,模拟交互是最高效的方式,不需要真实用户参与就能批量造数据:
- 先搭建一个目标属性知识库:给每个目标(比如猫、飞机、椅子)定义一组二元属性(比如「是否是动物」「是否有翅膀」「是否有毛发」等)。
- 生成模拟对话序列:对每个目标,随机打乱属性顺序,依次生成对应问题+标准答案(是/否),组成一条完整的问答路径(比如「它是动物吗?→是;它有翅膀吗?→否;...」)。
- 加入噪声增强鲁棒性:随机把10%-15%的标准答案反转(模拟用户误答),或者加入「不知道」这类模糊响应,让模型适应真实场景的不确定性。
数据量与迭代轮次
- 单目标样本数:如果你的目标类别在50-200个之间,每个目标至少生成50-100条不同的问答路径(不同的提问顺序),总样本量建议在5000-20000条之间。
- 训练迭代轮次:用LSTM训练时,batch size设为32或64,epochs先从20开始跑,观察验证集准确率:如果还在上升就继续加,直到准确率稳定(一般30-50个epochs足够收敛)。
这里给个简单的数据生成伪代码:
import random import numpy as np # 示例目标属性库 target_attrs = { "cat": {"is_animal": True, "has_fur": True, "can_fly": False, "has_tail": True}, "airplane": {"is_animal": False, "has_wings": True, "can_fly": True, "has_engine": True}, "chair": {"is_animal": False, "has_legs": True, "can_sit": True, "is_electronic": False} } target_list = list(target_attrs.keys()) # 生成单条训练样本 def gen_sample(target): attrs = target_attrs[target] # 随机打乱提问顺序 attr_order = random.sample(list(attrs.keys()), len(attrs)) dialogue = [] for attr in attr_order: # 映射属性为自然语言问题 q = f"它是动物吗?" if attr == "is_animal" else f"它有{attr.split('_')[1]}吗?" # 生成回答,随机加噪声 ans = 1 if attrs[attr] else 0 if random.random() < 0.1: # 10%概率反转答案 ans = 1 - ans dialogue.append((q, ans)) return dialogue, target_list.index(target) # 批量生成数据 train_data = [] for target in target_list: for _ in range(80): # 每个目标生成80条样本 train_data.append(gen_sample(target))
2. 问题生成方案
分训练阶段和推理阶段两种场景:
训练阶段:模板化生成
训练时不需要太智能的问题生成,直接用属性-问题模板映射即可:
- 给每个属性定义固定的问题模板,比如
is_animal→「它是动物吗?」,has_wings→「它有翅膀吗?」。 - 随机选取属性生成问题,保证训练数据的多样性就行。
推理阶段:基于信息增益的最优问题生成
推理时要选最能缩小目标范围的问题,信息增益(Information Gain)是最实用的方法:
- 先让LSTM根据当前对话历史输出所有目标的概率分布。
- 对每个属性,计算如果用户回答是/否后,能减少的熵(也就是信息增益),选增益最大的属性对应的问题。
示例代码(计算信息增益并生成问题):
def calc_info_gain(probs, attr): # 按属性分组目标概率 group_true = [] group_false = [] for idx, p in enumerate(probs): if target_attrs[target_list[idx]][attr]: group_true.append(p) else: group_false.append(p) # 计算各组权重 w_true = sum(group_true) w_false = sum(group_false) # 计算熵(加小值避免log(0)) entropy_true = -sum(p * np.log2(p + 1e-8) for p in group_true) if w_true > 0 else 0 entropy_false = -sum(p * np.log2(p + 1e-8) for p in group_false) if w_false > 0 else 0 current_entropy = -sum(p * np.log2(p + 1e-8) for p in probs) # 信息增益 = 当前熵 - 加权分组熵 return current_entropy - (w_true * entropy_true + w_false * entropy_false) def gen_next_question(probs): attrs = list(target_attrs[target_list[0]].keys()) # 计算每个属性的信息增益 gain_dict = {attr: calc_info_gain(probs, attr) for attr in attrs} best_attr = max(gain_dict, key=gain_dict.get) # 映射为自然语言问题 q_map = { "is_animal": "它是动物吗?", "has_fur": "它有毛发吗?", "can_fly": "它会飞吗?", "has_legs": "它有腿吗?" } return q_map[best_attr]
如果想更智能,也可以在LSTM基础上加一个seq2seq分支,让模型自动生成问题,但对20问场景来说,信息增益法已经足够高效且可控。
3. 用户响应处理
核心是把用户的自然语言输入转换成模型能理解的序列特征,同时处理歧义:
输入标准化
先用规则匹配把用户的模糊输入映射为标准标签:
- 定义「是」「否」的关键词库,比如「是、对、没错、yes」→映射为1,「否、不是、不对、no」→映射为0。
- 如果用户输入不在关键词里,提示用户重新回答,或者把「不知道」这类输入映射为特殊标记(比如0.5),让模型学习处理模糊情况。
序列更新
把用户的标准化响应加入到对话历史序列中,更新LSTM的隐藏状态,用于下一次的目标预测或问题生成。
示例代码:
def process_response(user_input): user_input = user_input.lower().strip() yes_words = {"是", "对", "没错", "y", "yes"} no_words = {"否", "不是", "不对", "n", "no"} if any(word in user_input for word in yes_words): return 1 elif any(word in user_input for word in no_words): return 0 else: print("不好意思,我没听懂你的回答,请直接告诉我「是」或「否」哦~") return None # 推理时更新对话历史 def update_dialogue_history(history, question, response): # 把问题和响应转换成模型输入的token(比如用预训练的词向量编码) q_tokens = tokenizer.texts_to_sequences([question])[0] r_token = [response] history.extend(q_tokens + r_token) # 截断过长的历史,避免LSTM序列太长 if len(history) > max_seq_len: history = history[-max_seq_len:] return history
内容的提问来源于stack exchange,提问作者Glennismade
相关产品推荐
相关产品推荐

