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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 08:22:14