mode/models 中的 Cross-View Training(CVT)半监督序列建模实现详解:从数据准备到一致性训练源码解析

发布时间:2026/9/7 19:18:54
mode/models 中的 Cross-View Training(CVT)半监督序列建模实现详解:从数据准备到一致性训练源码解析 mode/models 中的 Cross-View TrainingCVT半监督序列建模实现详解从数据准备到一致性训练源码解析【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models本文基于 research/cvt_text/README.md 及其配套源码系统讲解 mode/models 仓库中 CVTCross-View Training交叉视角训练半监督序列建模方案的完整落地流程如何准备 GloVe 词向量、1B 词无标注语料与 CoNLL-2000 分块数据如何配置与训练 CNN-BiLSTM 模型以及“多视角 教师软标签”的一致性训练机制在 TensorFlow 1.x 源码中是如何实现的。读完本文你可以独立复现 CVT 在 chunking文本分块任务上的训练/评估流程并理解每个关键超参数的底层作用。CVT 方案定位这个仓库实现了什么research/cvt_text/README.md 明确指出该目录包含论文Semi-Supervised Sequence Modeling with Cross-View TrainingEMNLP 2018Clark 等的官方实现代码当前支持两类序列任务序列标注sequence tagging以 CoNLL-2000 文本分块chunking为默认任务依存句法分析dependency parsing数据标签形如index_of_head-relation例如0-root。CVT 的核心思想是半监督学习模型同时用少量带标签数据和大量无标注数据训练。对无标注数据模型用一组“视角”view——即不同结构的预测头——互相学习教师模型输出的软概率作为伪标签让各个视角的预测保持一致consistency。运行环境要求以 README 为准TensorFlow 1.10.1、Numpy 1.14.5其他版本未经验证由于入口与配置代码使用了kwargs.iteritems()见 base/configure.py该实现面向Python 2在 Python 3 下需自行适配。仓库目录结构与各模块职责围绕 README 中的三个操作取数、训练、评估代码组织如下research/cvt_text/ ├── cvt.py # 训练/评估统一入口--modetrain/eval ├── fetch_data.sh # 下载 GloVe 词向量、1B 词语料、CoNLL-2000 分块数据 ├── preprocessing.py # 构建词表/词向量、划分 dev 集、写标签映射 ├── base/ │ ├── configure.py # 全部超参数与路径配置Config 类 │ ├── embeddings.py # 预训练词向量加载 │ └── utils.py ├── corpus_processing/ │ ├── unlabeled_data.py # 无标注语料流式读取 │ ├── example.py / minibatching.py │ └── scorer.py ├── model/ │ ├── encoder.py # CNN-BiLSTM 句子编码器 │ ├── multitask_model.py # 多任务/半监督模型、优化器、EMA │ ├── shared_inputs.py # 共享占位符与训练态 dropout 控制 │ └── task_module.py # 监督/半监督任务模块基类 ├── task_specific/ │ ├── word_level/ │ │ ├── tagging_module.py # 序列标注primary 视图 5 个一致性视图 │ │ ├── depparse_module.py # 依存句法分析模块 │ │ └── depparse_scorer.py / tagging_scorers.py │ └── task_definitions.py └── training/ ├── trainer.py # 训练主循环有标注/无标注交替 └── training_progress.py # 断点保存、dev 最优模型选择数据准备fetch_data.sh 与 preprocessing.py一键下载三类数据README 给出的第一步是运行fetch_data.sh位于 research/cvt_text/fetch_data.sh。脚本会在./data/raw_data/下创建三个子目录并下载数据bash fetch_data.sh数据目标目录下载源脚本内 curl用途GloVe 词向量glove.6B.zipdata/raw_data/pretrained_embeddings/nlp.stanford.edu预训练词向量300 维1 Billion Word LM Benchmark 语料data/raw_data/unlabeled_data/statmt.org无标注半监督数据CoNLL-2000 分块数据train.txt.gz、test.txt.gzdata/raw_data/chunk/clips.uantwerpen.be带标签训练/测试集脚本用unzip/tar xzf/gunzip解压后结束。README 同时说明论文中使用的其他数据集没有免费公开渠道因此未包含在仓库中。自定义数据的格式约定若要应用 CVT 到其他数据集README 规定数据必须放在data/raw_data/task_name/(train|dev|test).txt序列标注每行一个“词 空格 该词的标签”句与句之间用空行分隔依存句法分析每个标签形如index_of_head-relation例如0-root。预处理的三件事数据下载完成后运行python preprocessing.pyresearch/cvt_text/preprocessing.py。从其源码看它完成三件事构建词表与词向量调用PretrainedEmbeddingLoader(config).build()把glove.6B.300d.txt处理成词表与词向量矩阵供后续Trainer初始化 embedding 矩阵使用构造 dev 集CoNLL-2000 分块数据官方不带 dev 划分脚本将训练句随机打乱后取前 1500 句写成dev.txt其余写回train_subset.txtrandom.seed(0)保证可复现写标签映射以 BIOES 编码扫描训练集生成标签→编号的映射并持久化供任务模块确定输出类数。超参数配置base/configure.py 中的完整参数表README 指出“要更换训练任务或修改模型超参数请修改 base/configure.py”。Config 类 集中定义了全部超参数且每次启动时通过config.write()把当前配置序列化为data/models/model_name/config.json留档。以下是源码中的默认值及作用模式与任务参数默认值说明modetrain只能为train或evaltask_names[chunk]任务列表多于一个即训练多任务模型is_semisupTrue是否使用 CVTFalse 则退化为纯监督训练pretrained_embeddingsglove.6B.300d.txt使用的预训练词向量文件word_embedding_size300词向量维度编码器CNN-BiLSTM参数默认值说明use_charsTrue是否加入字符级 CNNchar_embedding_size50字符 embedding 维度char_cnn_filter_widths[2, 3, 4]字符 CNN 卷积核宽度char_cnn_n_filters100每种宽度对应的滤波器数量unidirectional_sizes[1024]第一层“双单向 LSTM”的隐层大小bidirectional_sizes[512]第二层双向 LSTM 的隐层大小projection_size512LSTM 与隐层的投影维度depparse_projection_size128依存句法双线性分类器表示维度标注任务tagging参数默认值说明label_encodingBIOES实体级标签编码方案label_smoothing0.1有标注训练时的标签平滑率优化参数默认值说明lr0.5基础学习率momentum0.9Momentum 优化器动量grad_clip1.0全局梯度范数裁剪上限warm_up_steps5000.0学习率线性 warm-up 步数lr_decay0.005学习率逐步衰减系数EMA、正则与规模参数默认值说明ema_decay0.998模型权重 EMA 平滑系数ema_testTrue测试时是否使用 EMA 权重ema_teacherFalse教师模型是否使用 EMA 权重labeled_keep_prob0.5有标注样本的 1 - dropoutunlabeled_keep_prob0.8无标注样本的 1 - dropoutmax_sentence_length100无标注句子的最大长度max_word_length20字符 CNN 的最大词长train_batch_size/test_batch_size64训练/测试批大小buckets[(0,15),(15,40),(40,1000)]按句长分桶训练节奏参数默认值说明print_every25打印训练进度的间隔步eval_dev_every500在 dev 集上评估的间隔eval_train_every2000在 train 集上评估的间隔save_model_every1000模型 checkpoint 间隔与 README“每 1000 步自动存档”一致train_set_percent100使用训练集的百分比配置对象对未知参数会直接抛ValueError路径方面则自动拼出原始数据目录data/raw_data/、预处理目录data/preprocessed_data/含word_vocabulary.pkl、word_embeddings.pkl以及模型目录data/models/model_name/内含checkpoints/、best_model_checkpoints/、summaries/、history.pkl等见 configure.py 第 106–132 行。模型架构CNN-BiLSTM 编码器编码器实现在 model/encoder.py对应论文中经典的 CNN-BiLSTM 结构Encoder.__init__按顺序构建四部分表示词表示word_reprs以预训练 GloVe 矩阵初始化word_embedding_matrixencoder.py 第 39–46 行做 dropout 后乘一个可学习标量emb_scale当use_charsTrue时再对字符序列查表得到字符 embedding用宽度[2,3,4]、各 100 个滤波器的tf.layers.conv1d做卷积ReLU 后按时间维reduce_max做最大池化最后与词向量拼接双单向 LSTM 表示uni_reprs源码用一对tf.nn.bidirectional_dynamic_rnn的前/后向输出uni_fw、uni_bw作为两个单向表示二者拼接得到uni_reprs——这正是 CVT “前向视角/后向视角”的物理来源encoder.py 第 72–92 行双向 LSTM 表示bi_reprs在uni_reprs之上再叠一层bidirectional_sizes[512]的标准双向 LSTM得到全双向表示bi_reprs各层输出经model_helpers.lstm_cell/multi_lstm_cell做隐层投影projection_size。编码器把uni_reprs、bi_reprs、uni_fw、uni_bw四种表示全部暴露给任务模块供不同“视角”选择性地使用。Cross-View Training 的核心机制源码级解析五个预测视角与一致性损失以序列标注为例task_specific/word_level/tagging_module.py 中的TaggingModule构建了 1 个主视图加 5 个辅助视图primary PredictionModule(primary, ([encoder.uni_reprs, encoder.bi_reprs])) ps [ PredictionModule(full, ([encoder.uni_reprs, encoder.bi_reprs]), activateFalse), PredictionModule(forwards, [encoder.uni_fw]), PredictionModule(backwards, [encoder.uni_bw]), PredictionModule(future, [encoder.uni_fw], roll_direction1), PredictionModule(past, [encoder.uni_bw], roll_direction-1), ] self.unsupervised_loss sum(p.loss for p in ps) # 5 个辅助视图的损失之和 self.supervised_loss primary.loss # 仅有标注任务用主视图各视角的输入表示各不相同full用全双向表示但不做 ReLU 激活forwards/backwards分别只用前向/后向单向 LSTMfuture/past在此基础上再对标签序列做 ±1 位移roll_direction模拟“用未来/过去上下文预测”。每个视图的输出经masked_ce_loss带 mask 的交叉熵计算损失并有标注训练时启用label_smoothing0.1shared_inputs.py 第 42–43 行平滑只作用于有标注训练无标注与测试时为 0。教师软标签的喂入方式无标注 batch 上没有真实标签TaggingModule.update_feed_dicttagging_module.py 第 70–76 行的处理逻辑是若 batch 带有teacher_predictions则直接把教师模型的 softmax 概率float16 存储作为“标签”喂入交叉熵否则才把真实 BIOES 标签 one-hot 化。教师模型与 EMAmultitask_model.py 的Model在同一个变量作用域tf.AUTO_REUSE内构建了trainer、tester、teacher三份推理第 45–67 行当ema_teacherFalse默认时教师就是与训练网络同构、但run_teacher以is_trainingFalse运行dropout 关闭的当前模型当ema_testTrue默认时用tf.train.ExponentialMovingAverage(ema_decay0.998)对model/作用域内的可训练变量做指数滑动平均并通过scope.set_custom_getter让tester读取 EMA 权重——即测试时使用 EMA 平滑后的参数若ema_teacherTrue教师也切换为 EMA 版本。训练循环有标注/无标注交替training/trainer.py 的Trainer.train与_get_training_mbs实现了 README 所述训练流程有标注数据集按sqrt(数据集大小)加权随机采样多任务时用于平衡各任务数据量每个循环先 yield 一个有标注 batch再当is_semisupTrue时yield 一个无标注 batch对无标注 batch先run_teacher教师对无标注句产生各任务概率再train_unlabeled用教师软标签训练 5 个辅助视图的一致性损失有标注 batch 只训练对应任务的主视图损失每eval_dev_every500步在 dev 集评估并save_if_best_dev_model保留最优 dev 模型每 2000 步评估 train 集每 1000 步写 checkpoint对应save_model_every1000history.pkl记录 dev 指标随步数的变化。无标注数据由 corpus_processing/unlabeled_data.py 的UnlabeledDataReader流式读取循环遍历 1B 词语料分片文件跳过长度 ≥max_sentence_length100的句子每积累 10000 句构造一批 Exampleendless_minibatches提供无限 batch 流文件读尽后从头再来。训练与评估命令、断点与预期效果训练python cvt.py --modetrain --model_namechunking_model该命令在 cvt.py 中定义了两个 FLAG--modetrain/eval默认train与--model_name模型名默认default_model。main()的调用链为构造Config并write()配置 → 建图Trainer内部构建多任务模型→ 初始化tf.train.Saver(max_to_keep1)与训练进度对象 → 进入model_trainer.train。两个值得注意的实现细节checkpoint 只保留 1 份max_to_keep1cvt.py 第 44 行配合progress.pkl记录训练进度含无标注语料读取位置因此 README 说“训练中断后重启会从最新 checkpoint 继续”dev 指标最优时的参数另存于best_model_checkpoints/供评估使用。评估python cvt.py --modeeval --model_namechunking_modeleval模式下cvt.py 第 57–61 行 会用best_model_saver.restore从best_model_checkpoints恢复dev 集最优模型然后对所有任务逐 minibatch 前向打分Trainer._evaluate_task调用任务 scorer 累计 F1/Acc 等指标。预期结果README 原文给出的基准在 chunking 数据上训练 200k 步的 CVT 模型dev 集 F1 应不低于97.1test 集 F1 应不低于96.6。多任务与依存句法分析把task_names配置成多个任务如[chunk, depparse]即训练多任务模型任务模块由 task_definitions.py 按任务名分发依存分析使用 depparse_module.py双线性分类器表示维度为depparse_projection_size128与 depparse_scorer.py 评估标注任务则走 tagging_scorers.py。此时训练循环中各任务有标注数据按平方根权重交替采样一致性损失则是各任务辅助视图损失之和见 multitask_model.py 第 76–78 行。复现清单按 README 顺序安装 TensorFlow 1.10.1 与 Numpy 1.14.5Python 2 环境bash fetch_data.sh—— 下载 GloVe 词向量、1B 词无标注语料、CoNLL-2000 chunking 数据至data/raw_data/python preprocessing.py—— 构建词表/词向量、生成 chunk dev 集、写标签映射按需修改 base/configure.py任务、超参数、batch 大小等python cvt.py --modetrain --model_namechunking_model训练约 200k 步产物在data/models/chunking_model/python cvt.py --modeeval --model_namechunking_model评估核对 dev F1 ≥ 97.1、test F1 ≥ 96.6 的基准。引用与联系若在你的研究中使用本代码请按 README 要求引用原论文inproceedings{clark2018semi, title {Semi-Supervised Sequence Modeling with Cross-View Training}, author {Kevin Clark and Minh-Thang Luong and Christopher D. Manning and Quoc V. Le}, booktitle {EMNLP}, year {2018} }README 还给出了作者联系方式Kevin Clark clarkkev、Thang Luong lmthang。小结research/cvt_text 是 CVT 半监督序列建模在 TensorFlow 1.x 上的完整可复现实现fetch_data.shpreprocessing.py打通数据链路base/configure.py集中管理全部超参数并自动留档CNN-BiLSTM 编码器同时输出单向/双向四种表示TaggingModule用 5 个结构不同的预测视图对齐教师软概率实现一致性训练Trainer则负责有标注/无标注交替、dev 最优模型选择与断点续训。对于想在标注数据有限、无标注语料充足的场景下做序列标注或句法分析的工程实践者这份代码提供了从数据格式、超参数到一致性损失实现的端到端参考。【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考