
深度学习NLP【免费下载链接】seq2seqA general-purpose encoder-decoder framework for Tensorflow项目地址https://gitcode.com/gh_mirrors/seq2seq1/seq2seq点击查看免费下载导读本文聚焦 Google seq2seqtf-seq2seqTensorFlow 通用 encoder-decoder 框架的推理inference环节系统讲解bin/infer.py中推理任务Inference Task的完整机制与用法。你将掌握DecodeText的输出定制、基于注意力得分与词典映射两种 UNK 替换方案、DumpAttention注意力对齐可视化、DumpBeams束搜索调试信息导出以及如何组合多个任务完成端到端解码。读完即可在自己训练的模型上直接复现这些实战命令。Inference Task 机制一切推理功能都是 SessionRunHook调用推理脚本bin/infer.py时必须通过--tasks参数提供一个任务列表最基本的任务是DecodeText它只负责把模型预测打印到标准输出。通过追加更多任务可以扩展额外能力例如存储调试信息或可视化注意力得分。在底层每个InferenceTask都被实现为一个 TensorFlow 的 SessionRunHooksix.add_metaclass(abc.ABCMeta) class InferenceTask(tf.train.SessionRunHook, Configurable): Abstract base class for inference tasks. Params: model_class: The model class to instantiate. If undefined, re-uses the class used during training. model_params: Model hyperparameters. Specified hyperparameters will overwrite those used during training. def begin(self): self._predictions graph_utils.get_dict_from_collection(predictions)begin()从 TensorFlow 图集合graph collection的predictions中取回模型输出字典具体的任务子类再通过before_run()声明需要 fetch 的张量、在after_run()中处理每个 batch 的结果。这也解释了为什么各任务与MonitoredSession配合即可工作——bin/infer.py 中正是把任务列表当作hooks传给MonitoredSessionwith tf.train.MonitoredSession( session_creatorsession_creator, hookshooks) as sess: # Run until the inputs are exhausted while not sess.should_stop(): sess.run([])推理脚本支持的主要 FLAG 见 bin/infer.pytasksYAML/JSON 字符串格式的任务列表、model_params覆盖模型推理参数、config_pathYAML 配置文件、input_pipeline推理数据读取定义、model_dir模型 checkpoint 目录、checkpoint_path指定 checkpoint缺省用model_dir中最新一个、batch_size默认 32。DecodeText把预测结果打印到标准输出DecodeText读取模型预测并把预测打印到标准输出其参数如下参数默认值说明delimiter空格拼接模型预测 token 时使用的分隔字符串unk_replaceFalse设为True时基于注意力得分执行未知 token 替换详见下文unk_mappingNone若设为某个词典文件的路径则用该映射执行未知 token 替换详见下文postproc_fn可选的句子后处理函数完整限定名如seq2seq.data.postproc.decode_sentencepiece这些默认值定义在 seq2seq/tasks/decode_text.py 的default_params()中。从after_run()的实现seq2seq/tasks/decode_text.py可以看到几个关键细节若使用了 beam searchpredicted_tokens维度大于 1代码只取第一个 beampredicted_tokens[:, 0]作为最终输出输出句按delimiter拼接后再以SEQUENCE_END截断去掉结束符及其后续内容若配置了postproc_fn会通过pydoc.locate加载该函数并作用于输出句子。最基础的推理命令如下仅打印预测不包含任何额外任务export PRED_DIR${MODEL_DIR}/pred mkdir -p ${PRED_DIR} python -m bin.infer \ --tasks - class: DecodeText \ --model_dir $MODEL_DIR \ --input_pipeline class: ParallelTextInputPipeline params: source_files: - $DEV_SOURCES \ ${PRED_DIR}/predictions.txt其中--model_dir指向训练时设置的output_dir存放 checkpoint 的目录--input_pipeline的格式与训练时使用的输入管道定义一致。该命令的完整上下文可参考 docs/nmt.md。用 Copy 机制替换 UNK token地名、人名等稀有词经常不在目标词表target vocabulary中导致预测输出里出现UNKtoken。一个简单有效的策略是把每个UNKtoken 替换为源序列中与它对齐得最好的那个词。对齐通常通过注意力机制计算注意力机制为每个目标 token 产生一组对齐得分。如果你训练的模型能够产生这类注意力得分例如AttentionSeq2Seq其解码器默认是AttentionDecoder就可以通过打开unk_replace参数来执行 UNK 替换mkdir -p ${DATA_PATH}/pred python -m bin.infer \ --tasks - class: DecodeText params: unk_replace: True从实现上看替换逻辑在 seq2seq/tasks/decode_text.py 的_unk_replace()中遍历每个预测 token若等于UNK用np.argmax取注意力得分最高的源位置取该位置的源 token 作为替换词另外after_run()中会把注意力得分按source_len - 1切片seq2seq/tasks/decode_text.py避免意外地把UNK替换成SEQUENCE_ENDtoken。注意unk_replace依赖模型输出中存在attention_scores。从 seq2seq/models/seq2seq_model.py 的_create_predictions()和 seq2seq/decoders/attention_decoder.py 可见AttentionDecoderOutput这个 namedtuple 显式包含attention_scores字段而DecodeText的before_run()也只在 predictions 中存在attention_scores时才 fetch 它seq2seq/tasks/decode_text.py。用映射表替换 UNK tokenCopy 机制只在“源词本身就是要输出的目标词”时有效。例如英语的 Munich 在德语中通常译为 München即使注意力完美对齐直接复制 Munich 也永远得不到正确翻译。更精细的做法是使用词典映射。一种常用策略是用 fast_align 基于条件概率 p(target | source) 生成映射。完整的生成流程如下# 1. 下载并编译 fast_align git clone https://github.com/clab/fast_align.git mkdir fast_align/build cd fast_align/build cmake ../ make # 2. 把数据转换成 fast_align 认识的格式source ||| target paste \ $HOME/nmt_data/toy_reverse/train/sources.txt \ $HOME/nmt_data/toy_reverse/train/targets.txt \ | sed s/$(printf \t)/ ||| /g $HOME/nmt_data/toy_reverse/train/source_targets.fastalign # 3. 学习对齐 ./fast_align \ -i $HOME/nmt_data/toy_reverse/train/source_targets.fastalign \ -v -p $HOME/nmt_data/toy_reverse/train/source_targets.cond \ $HOME/nmt_data/toy_reverse/train/source_targets.align # 4. 为每个源词找出最可能的译文写入词典文件 sort -k1,1 -k3,3gr $HOME/nmt_data/toy_reverse/train/source_targets.cond \ | sort -k1,1 -u \ $HOME/nmt_data/toy_reverse/train/source_targets.cond.dict-p参数指定的输出文件包含p(target | source)条件概率格式为source\ttarget\tprob即源词、目标词、概率三个字段以制表符分隔。随后把该词典路径传给unk_mapping即可实现更智能的 UNK 替换mkdir -p ${DATA_PATH}/pred python -m bin.infer \ --tasks - class: DecodeText params: unk_replace: True unk_mapping: $HOME/nmt_data/toy_reverse/train/source_targets.cond.dict \ --model_dir $MODEL_DIR \ --input_pipeline class: ParallelTextInputPipeline params: source_files: - $DEV_SOURCES从源码看映射文件的加载实现在 seq2seq/tasks/decode_text.py 的_get_unk_mapping()逐行读取按制表符切分取前两个字段构成source - target的字典并去除两端空白。替换时seq2seq/tasks/decode_text.py先按注意力得分选出源 token再查映射表得到最终目标词若源 token 不在映射表中则回退为直接复制源词。也就是说unk_mapping是在unk_replace基础上的增强——映射文件只负责把“选中的源 token”翻译成目标词。用 DumpAttention 可视化注意力对齐如果你使用AttentionDecoder训练模型可以在推理时用DumpAttention任务导出原始注意力得分并生成对齐可视化图python -m bin.infer \ --tasks - class: DecodeText - class: DumpAttention params: output_dir: $HOME/attention \ --model_dir $MODEL_DIR \ --input_pipeline class: ParallelTextInputPipeline params: source_files: - $DEV_SOURCESoutput_dir为必填参数seq2seq/tasks/dump_attention.py 中若未指定会直接抛出ValueErrorbegin()中会自动创建该目录。默认情况下该任务会为每个样本生成一个注意力得分数组文件和一张注意力图得分数组实际写入的文件是attention_scores.npz由np.savez生成见 seq2seq/tasks/dump_attention.py可用numpy.load加载其中包含一组形状为[target_length, source_length]的数组——第 i 行对应第 i 个目标 token 对所有源 token 的对齐得分_get_scores()会按源/目标实际长度切片seq2seq/tasks/dump_attention.py注意力图每个样本输出一张00000.png、00001.png…按{:05d}.png命名横轴为源词、纵轴为目标词用plt.cm.Blues色图渲染得分矩阵seq2seq/tasks/dump_attention.py。如果你只需要原始得分数据而不需要绘图可以设置dump_plots: False。需要说明的是当前仓库源码中控制绘图的实际参数名为dump_plots默认True见 seq2seq/tasks/dump_attention.py文档早期版本提到的dump_attention_no_plot是更早的参数命名以当前源码为准。上述行为在 seq2seq/test/pipeline_test.py 中有端到端验证推理结束后断言attention_scores.npz与00002.png存在并逐一检查每个数组的第二维源长度与测试数据一致。用 DumpBeams 导出 Beam Search 调试信息如果在解码时启用了 beam search可以使用DumpBeams任务把束搜索调试信息写入磁盘之后既可以用 numpy 直接检查数据也可以用仓库提供的可视化脚本生成可视化结果python -m bin.infer \ --tasks - class: DecodeText - class: DumpBeams params: file: ${TMPDIR:-/tmp}/wmt_16_en_de/newstest2014.pred.beams.npz \ --model_params inference.beam_search.beam_width: 5 \ --model_dir $MODEL_DIR \ --input_pipeline class: ParallelTextInputPipeline params: source_files: - $DEV_SOURCESfile参数为必填项未指定会抛出ValueError。导出的.npz文件包含四个带名字典项见 seq2seq/tasks/dump_beams.pykey含义predicted_ids每个 beam 的预测 token idbeam_parent_ids每个 beam 的父节点 id用于还原搜索树结构scores各 beam 的得分log_probs各 beam 的对数概率要启用 beam search需要覆盖模型参数inference.beam_search.beam_width。该参数定义在 seq2seq/models/seq2seq_model.py 中默认值为0小于等于1即禁用束搜索use_beam_search属性返回beam_width 1见 seq2seq/models/seq2seq_model.py。相关参数还有inference.beam_search.length_penalty_weight默认0.0束搜索假设的长度惩罚系数与inference.beam_search.choose_successors_fn默认choose_top_k选择后继的方式。束搜索的构造逻辑见 seq2seq/models/seq2seq_model.py。另外注意启用 beam search 后推理的 batch size 会被强制设为 1见 seq2seq/inference/inference.py且解码耗时明显更长。对导出的数据可以直接用 numpy 加载检查或运行可视化脚本bin/tools/generate_beam_viz.py生成 beam search 树形可视化该脚本基于 networkx 构建搜索树并输出 HTML因此需要先安装networkx详见 docs/tools.mdpython -m bin.tools.generate_beam_viz \ -o ${TMPDIR:-/tmp}/beam_visualizations \ -d ${TMPDIR:-/tmp}/beams.npz \ -v $HOME/nmt_data/toy_reverse/train/vocab.targets.txt-d指定DumpBeams导出的数据文件-o指定输出目录-v可选地指定词表文件脚本会把 token id 还原为可读词见 bin/tools/generate_beam_viz.py。组合多个任务与指定 checkpoint推理任务是可以自由组合的。例如在 docs/nmt.md 的神经机器翻译教程中同一个推理命令同时运行DecodeText打印译文和DumpBeams导出束搜索信息并用--model_params打开 beam searchpython -m bin.infer \ --tasks - class: DecodeText - class: DumpBeams params: file: ${PRED_DIR}/beams.npz \ --model_dir $MODEL_DIR \ --model_params inference.beam_search.beam_width: 5 \ --input_pipeline class: ParallelTextInputPipeline params: source_files: - $DEV_SOURCES \ ${PRED_DIR}/predictions.txt--model_params在推理时仅覆盖本次运行的模型参数训练时保存的参数会被递归合并覆盖见 bin/infer.py不会改动磁盘上的训练配置。关于 checkpoint训练脚本会在训练过程中保存多个 checkpoint默认情况下推理脚本使用model_dir中最新的 checkpoint若想评估某个特定的 checkpoint传入checkpoint_path标志即可bin/infer.py--checkpoint_path $MODEL_DIR/model.ckpt-50这类用法在 seq2seq/test/pipeline_test.py 的集成测试中也有体现。小结本文围绕bin/infer.py的推理任务体系从底层SessionRunHook机制讲起完整覆盖了四个核心任务的配置、命令与产物DecodeText基础输出任务支持delimiter、unk_replace、unk_mapping、postproc_fnUNK 替换Copy 机制基于注意力得分对齐源词与 fast_align 词典映射两条路径DumpAttention导出attention_scores.npz与逐样本注意力图DumpBeams导出predicted_ids、beam_parent_ids、scores、log_probs配合bin/tools/generate_beam_viz.py可视化。对应的源码实现集中在 seq2seq/tasks/inference_task.py、decode_text.py、dump_attention.py、dump_beams.py入口脚本为 bin/infer.py端到端行为可由 seq2seq/test/pipeline_test.py 验证模型侧参数如inference.beam_search.*、inference.max_decode_length见 docs/models.md。掌握了这套任务组合机制你就可以按需为推理流程叠加打印、替换、可视化与调试能力。赞分享深度学习NLP【免费下载链接】seq2seqA general-purpose encoder-decoder framework for Tensorflow项目地址https://gitcode.com/gh_mirrors/seq2seq1/seq2seq点击查看免费下载相关推荐seq2seq 推断任务实战指南DecodeText 解码、注意力可视化与 Beam Search 调试seq2seq 推断任务实战指南DecodeText 解码、注意力可视化与 Beam Search 调试 本文围绕 TensorFlow 通用 seq2seq深度学习NLP终极指南seq2seq推理任务完全解析 - 文本解码、注意力可视化与束搜索调试终极指南seq2seq推理任务完全解析 文本解码、注意力可视化与束搜索调试 seq2seq序列到序列模型是自然语言处理领域的核心技术广泛应用于机器翻译、深度学习NLPseq2seq 项目工具链实战词汇表生成与 Beam Search 可视化seq2seq 项目工具链实战词汇表生成与 Beam Search 可视化 seq2seq 是一个通用的 encoder decoder 框架适用于机器翻译深度学习NLP上一篇React Native应用发布指南generator-rn-toolbox的App Store和Google Play部署下一篇Cellpose项目中OpenCV图像缩放错误的深度解析与解决方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考