ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

基于序列到序列注意力模型的 TensorFlow 神经聊天机器人实战:stanford-tensorflow-tutorials 指南

基于序列到序列注意力模型的 TensorFlow 神经聊天机器人实战:stanford-tensorflow-tutorials 指南 教程深度学习【免费下载链接】stanford-tensorflow-tutorialsThis repository contains code examples for the Stanfords course: TensorFlow for Deep Learning Research.项目地址https://gitcode.com/gh_mirrors/st/stanford-tensorflow-tutorials点击查看免费下载本篇技术指南以仓库 assignments/chatbot 目录下的神经聊天机器人为核心完整讲解如何基于序列到序列Sequence-to-Sequence模型 注意力解码器构建一个可直接运行的对话机器人。文章覆盖数据准备、超参数配置、预处理流水线、模型架构、训练与交互式聊天全流程并结合 config.py、data.py、model.py、chatbot.py 的源码实现逐层剖析原理。读完本文你将掌握 seq2seq 聊天机器人的完整工程化落地方法并能独立完成从 Cornell Movie-Dialogs 语料到可用聊天机器人的端到端训练与部署。一、项目背景CS20 课程中的神经聊天机器人该聊天机器人由斯坦福大学 CS20 课程TensorFlow for Deep Learning Researchcs20.stanford.edu的讲师 Chip Huyen 创建属于课程作业assignment之一的完整可运行项目。其技术路线明确写在 README.md 中采用带注意力解码器attentional decoder的序列到序列模型sequence to sequence model模型框架借鉴自 Google 官方 TensorFlow 模型库中的机器翻译教程Google Translate Tensorflow model即tensorflow/models仓库tutorials/rnn/translate目录下的经典实现序列到序列模型的理论基础出自 Cho et al.2014的经典论文。因此这个聊天机器人本质上是一个对话领域的翻译任务把用户输入的一句话encoder 端翻译成机器人的回答decoder 端通过注意力机制在解码的每一步动态聚焦输入序列中与当前生成词最相关的部分。与机器翻译不同聊天机器人没有现成的平行语料因此项目选用康奈尔电影对话语料库Cornell Movie-Dialogs Corpus该语料包含大量电影剧本中的多轮对话天然适合构造问题—回答式的训练对。二、环境与依赖主仓库 README.md 说明课程使用Python 3.6 TensorFlow 1.4.1。聊天机器人代码大量使用了 TensorFlow 1.x 时代的经典 APItf.contrib.rnn.GRUCell/tf.contrib.rnn.MultiRNNCelltf.contrib.legacy_seq2seq.embedding_attention_seq2seq与model_with_bucketstf.compat.as_str其中contrib与legacy_seq2seq在 TensorFlow 2.x 中已被移除因此本项目只能在 TensorFlow 1.x推荐 1.4.1环境下运行这是复现实验的先决条件。其余依赖可参考 setup/requirements.txt核心为tensorflow1.4.1另有scipy、scikit-learn、matplotlib、xlrd、Pillow等课程通用依赖设置步骤详见 setup/setup_instruction.md。三、四步快速上手依据 README.md 的 Usage 章节完整的运行流程分为四步Step 1准备数据。在项目目录下创建data文件夹下载并解压Cornell Movie-Dialogs Corpus电影对话语料库解压后应包含movie_lines.txt与movie_conversations.txt两个核心文件。注意目录名默认是cornell movie-dialogs corpus含空格与配置文件中的默认DATA_PATH保持一致。Step 2修改配置。编辑 config.py将DATA_PATH改成你实际存放语料的路径例如DATA_PATH data/cornell movie-dialogs corpusStep 3数据预处理。在assignments/chatbot目录下执行python3 data.py该命令会完成 Cornell 语料的全部预处理详见第五节并在processed目录下生成模型可直接读取的 id 序列文件。Step 4训练 / 聊天。执行# 训练模式 python3 chatbot.py --mode train # 聊天模式 python3 chatbot.py --mode chat--mode仅接受train或chat两个取值默认是 train见 chatbot.py 的参数解析逻辑train 模式默认会恢复 checkpoints 文件夹中已有的训练权重并继续训练若想从零开始请删除 checkpoints 文件夹中的所有 checkpoint 文件chat 模式进入与机器人交互的命令行模式默认情况下你与机器人的所有对话都会被追加写入processed/output_convo.txt。此外入口代码还有一个隐性的自动流程首次运行时若发现processed目录不存在会自动依次执行prepare_raw_data()与process_data()见 chatbot.py无需手工重复 Step 3。四、核心超参数配置详解config.py 集中了全部可调超参数理解这些参数是调优与复现的基础。4.1 数据与路径参数参数默认值作用DATA_PATHdata/cornell movie-dialogs corpusCornell 语料所在目录CONVO_FILEmovie_conversations.txt对话记录文件LINE_FILEmovie_lines.txt台词文件OUTPUT_FILEoutput_convo.txt聊天记录输出文件PROCESSED_PATHprocessed预处理产物目录CPT_PATHcheckpoints模型权重保存目录THRESHOLD2词频过滤阈值出现次数低于该值的词将被丢弃TESTSET_SIZE25000从问答对中随机抽取的测试集规模4.2 特殊符号 ID符号ID含义PAD_ID0填充符padUNK_ID1未登录词unkSTART_ID2解码起始符sEOS_ID3句尾符\s这四个 ID 与词汇表文件vocab.enc/vocab.dec的前四行一一对应见 data.py是序列转换的基础约定。4.3 桶Bucket配置BUCKETS [(19, 19), (28, 28), (33, 33), (40, 43), (50, 53), (60, 63)]每个桶是(encoder_max_len, decoder_max_len)二元组训练/解码时按句长就近放入最紧凑的桶避免为超长句做全长度 padding从而显著提升批量计算效率。从源码看encoder 侧最大长度BUCKETS[-1][0] 60直接决定了 model.py 中encoder_inputs占位符的数量decoder 侧最大长度BUCKETS[-1][1] 1 64对应decoder_inputs与decoder_masks占位符数量多出的 1 是 GO 符号位chat 模式下单条输入的最大长度即为config.BUCKETS[-1][0]60 个 token超长输入会被拒绝见 chatbot.py。仓库 2017 版 2017/assignments/chatbot/config.py 中保留了调桶过程中的分布观察注释语料中 encoder 句长分布集中在较短的区间作者曾尝试(6,8)~(39,44)的 9 桶方案、(8,10)~(39,43)的 5 桶方案等最终采用的 6 桶方案在训练样本分布如 [19 530 / 17 449 / 17 585 / 23 444 / 22 884 / 16 435] 量级上表现最优——这说明桶的划分应根据实际语料长度分布来调整而非固定不变。4.4 缩略语替换规则CONTRACTIONS [(i m , i m ), ( d , d ), ( s , s ), (don t , do nt ), ...]在预处理中用于把分词产生的形如i m的碎片重新粘合为i m等规范形式改善词表质量。4.5 模型与训练参数参数默认值说明NUM_LAYERS3多层 RNN 的层数GRU 堆叠层数HIDDEN_SIZE256隐层维度同时也是词嵌入维度BATCH_SIZE64训练批大小chat 模式固定为 1LR0.5优化器学习率SGDMAX_GRAD_NORM5.0全局梯度裁剪范数上限NUM_SAMPLES512采样 softmax 的采样数0 表示关闭五、数据预处理流水线data.pydata.py 承担从原始语料到模型输入的全部预处理流程如下5.1 原始数据解析get_lines()逐行读取movie_lines.txt按 $ 分隔符解析出台词ID - 台词文本的映射字段数须为 5get_convos()读取movie_conversations.txt解析出每条对话包含的台词 ID 序列question_answers()把每段对话切分为连续的(前一句, 后一句)问答对构成训练/测试数据prepare_dataset()随机抽取TESTSET_SIZE25000个问答对作为测试集其余写入train.enc/train.dec/test.enc/test.dec四个文件enc 为问题、dec 为回答。5.2 分词与词表构建basic_tokenizer()实现了基础分词器统一转小写、去除u/u与[]标记、按([.,!?-:;)(])正则切分标点并将数字统一替换为#normalize_digitsTrue从而把123与456归并为同一个 token缓解数字稀疏问题。build_vocab()统计词频后首先固定写入 4 个特殊符号pad、unk、s、\s索引 0~3按词频降序写入出现次数不低于THRESHOLD的词自动把词表大小以ENC_VOCAB N/DEC_VOCAB N的形式追加写回 config.pydata.py供模型构建时读取——这是本项目一个值得注意的自动化设计词表规模由数据决定并回写配置。5.3 序列化与分桶token2id()将文本转为 ID 序列decoder 端序列需在开头加s、末尾加\s而 encoder 端不加。load_data()按BUCKETS把每个问答对归入满足len(enc) enc_max and len(dec) dec_max的最小桶。get_batch()负责批量生成对 encoder 输入做padding 逆序list(reversed(...))这是 seq2seq 的经典技巧缩短信息传播路径decoder 输入做 padding但不逆序生成decoder_masks对应目标为 PAD 或最后一个位置时 mask 置 0从而在损失计算中屏蔽填充位。六、模型架构model.pymodel.py 定义了ChatBotModel类构造函数接受两个关键参数forward_only是否只构建前向传播聊天/评估时为 True不创建反向传播路径batch_size批大小训练 64聊天 1。build_graph()依次构建四部分6.1 占位符Placeholders按桶最大长度创建encoder_inputs、decoder_inputs、decoder_masks三个列表目标序列targets decoder_inputs[1:]即跳过 GO 符号。6.2 推理单元Inference当0 NUM_SAMPLES DEC_VOCAB时创建输出投影矩阵proj_w[HIDDEN_SIZE, DEC_VOCAB]与偏置proj_b并使用tf.nn.sampled_softmax_loss做采样 softmax——当词表很大时这能大幅降低 softmax 的计算开销基元单元为GRUCell(HIDDEN_SIZE)用MultiRNNCell堆叠NUM_LAYERS3层。6.3 损失与解码Loss Decode通过tf.contrib.legacy_seq2seq.embedding_attention_seq2seq构建带注意力的 seq2seq配合model_with_buckets为每个桶独立展开计算图训练时feed_previousFalse使用教师强制teacher forcing逐词监督聊天/评估时feed_previousTrue模型自回归解码把上一步输出作为下一步输入若启用了输出投影解码输出还需经过matmul(output, w) b投影回词表空间。该阶段建图耗时较长源码中专门打印了might take a couple of minutes的提示。6.4 优化器使用GradientDescentOptimizer(config.LR)SGD学习率 0.5对每个桶独立执行tf.clip_by_global_norm梯度裁剪上限MAX_GRAD_NORM 5.0缓解梯度爆炸维护不可训练的global_step变量供 checkpoint 命名与训练进度追踪使用。七、训练与聊天的执行逻辑chatbot.pychatbot.py 是入口与驱动模块核心机制如下7.1 单步执行 run_steprun_step()通过占位符的.name构造input_feed字典喂入模型并依据forward_only决定输出训练返回(梯度范数, 损失, None)推理返回(None, 损失, 各时间步 logits)。喂入前还会通过_assert_lengths()校验各序列长度与桶定义一致。7.2 桶采样与训练循环_get_buckets()统计每个桶的样本数并计算累积比例train_buckets_scale_get_random_bucket()依据随机数落在的区间选择桶实现按样本量加权的随机分桶采样_get_skip_step()决定保存间隔迭代数 100 时每 30 步保存一次之后每 100 步保存一次每10 * skip_step步调用_eval_test_set()在测试集上评估各桶损失每次保存调用saver.save(sess, os.path.join(config.CPT_PATH, chatbot), global_step...)权重落在checkpoints/目录_check_restore_parameters()通过tf.train.get_checkpoint_state检测既有 checkpoint 并自动恢复这就是 README 所说默认续训的实现位置。7.3 聊天模式chat()流程加载processed/vocab.enc与processed/vocab.decload_vocab同时返回词表列表与 词-ID 反向映射以batch_size1构建只含前向路径的模型恢复权重打印欢迎语Welcome to TensorBro. Say something. Enter to exit.并提示最大输入长度为 60对每行输入sentence2id()分词并映射为 ID未登录词用unk按长度选桶_find_right_bucketget_batch生成单样本批run_step前向解码_construct_response()对每个时间步取logits 的 argmax贪心解码若出现EOS_ID则截断最终把 ID 序列还原为单词并拼接成回答每条HUMAN .../BOT ...对话都追加写入processed/output_convo.txt退出时写入分隔线...。八、真实对话效果参考README.md 中给出了训练后机器人的示例对话output_convo.txt 则保留了更多轮次的真实交互记录。以下摘录一段有代表性的对话HUMAN hi BOT hi . what s your name ? HUMAN im chip BOT please . HUMAN where are you from? BOT california . HUMAN do you like california? BOT i m okay . HUMAN what do you like to do in your free time? BOT and what about the gun ? HUMAN do you have a gun? BOT yes . HUMAN youre dangerous! BOT i m not .从记录可以看出该模型的真实行为特征能学会打招呼、回应来自哪里等高频模式但不具备真实世界知识where are you from? 回答 california而问年龄回答 eighteen 也是语料中高频出现的答案面对无法回答的问题会退回到语料中出现频率很高的兜底句如i don t know what to say .、let s talk about something else .、i m fine .回答中的分词格式如i m是基本分词器在小写化与标点切分后的直接产物属于预期行为。这些记录同时说明seq2seq 聊天机器人本质是基于语料分布的文本生成模型其回答质量受限于训练语料覆盖度与模型容量适合作为教学示例理解对话生成机制而非可直接商用的产品级助手。九、运行注意事项与调优提示必须使用 TensorFlow 1.xcontrib/legacy_seq2seqAPI 在 TF 2.x 中已移除请严格按 requirements.txt 安装tensorflow1.4.1首次训练前确保 processed 已生成若processed目录不存在chatbot.py会自动触发预处理但更推荐显式执行python3 data.py以观察每一步输出从零训练删除checkpoints/中全部文件程序只检查 checkpoint 状态文件是否存在词表大小由数据自动决定build_vocab会把ENC_VOCAB/DEC_VOCAB回写到config.py如果重复执行预处理会产生重复追加行注意清理调参入口关注THRESHOLD词表规模、BUCKETS句长分布适配、NUM_LAYERS/HIDDEN_SIZE容量、LR/MAX_GRAD_NORM优化稳定性与NUM_SAMPLES大词表下的 softmax 加速对话记录聊天输出默认追加到processed/output_convo.txt可用于观察模型行为、收集 bad case。十、源码导览文件职责assignments/chatbot/README.md项目说明、使用步骤、示例对话assignments/chatbot/config.py全部超参数与路径配置assignments/chatbot/data.py语料解析、分词、词表构建、分桶与批生成assignments/chatbot/model.pyChatBotModelseq2seq 注意力 多桶训练/解码assignments/chatbot/chatbot.py训练循环、评估、命令行交互入口assignments/chatbot/output_convo.txt多轮真实对话记录2017/assignments/chatbot/config.py2017 版配置含调桶分布笔记结语本仓库的聊天机器人是一个教科书级别的 seq2seq 注意力实现数据侧涵盖解析、分词、词表、分桶、padding/mask 的完整工程细节模型侧涵盖多桶展开、采样 softmax、教师强制与自回归解码、梯度裁剪的经典技巧工程侧涵盖 checkpoint 续训、随机分桶采样、命令行双模式与对话日志落盘。它既适合作为学习序列到序列对话系统的起点也适合作为后续向 Transformer、BERT 等现代架构迁移的基线参照。按照本文的四步流程你即可在自己的环境上完整复现一个可交互的神经网络聊天机器人。赞分享教程深度学习【免费下载链接】stanford-tensorflow-tutorialsThis repository contains code examples for the Stanfords course: TensorFlow for Deep Learning Research.项目地址https://gitcode.com/gh_mirrors/st/stanford-tensorflow-tutorials点击查看免费下载相关推荐BladeOne缓存机制揭秘MODE_AUTO、MODE_SLOW、MODE_FAST三种模式详解BladeOne缓存机制揭秘MODE_AUTO、MODE_SLOW、MODE_FAST三种模式详解 BladeOne作为一款高性能的PHP模板引擎其 缓存机stanford-tensorflow-tutorials循环神经网络状态管理动态序列长度处理stanford tensorflow tutorials循环神经网络状态管理动态序列长度处理 在自然语言处理、时间序列预测等领域输入数据往往具有可变长度的教程深度学习homophonous_logography/neural基于注意力序列到序列模型的书素度Logography神经度量训练与评测指南homophonous_logography/neural基于注意力序列到序列模型的书素度Logography神经度量训练与评测指南 本指南系统介绍 ho人工智能深度学习NLP计算机视觉强化学习创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进