news 2026/9/23 17:49:12

基于学生-导师协同训练的噪声数据集清洗:student_mentor_dataset_cleaning 完整使用指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于学生-导师协同训练的噪声数据集清洗:student_mentor_dataset_cleaning 完整使用指南
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

本篇技术指南以 Google Research 仓库中的 student_mentor_dataset_cleaning/README.md 为主体,结合其 main.py 与training/子模块源码,系统讲解"学生-导师(Student-Mentor)协同训练"这一噪声数据集清洗与训练框架的完整使用流程。读完本文,你将掌握该项目的环境搭建、两种运行模式(softmax / triplet)、全部命令行参数语义,以及训练循环、噪声注入、梯度快照、模型持久化等底层机制,能够独立复现官方示例命令并扩展到自己的人工噪声数据集上。

一、项目定位与核心思想

student_mentor_dataset_cleaning是一个用于从带噪数据集中学习的研究性训练框架。其基本思想是维护两个模型:

  • Student(学生):承担实际分类/度量学习任务的目标模型,在带噪、类不平衡的数据上训练;
  • Mentor(导师):一个二元分类器,以学生模型在每条样本上的梯度快照为输入,输出该样本是否"可信"的置信度(权重),从而引导学生过滤噪声样本。

二者的训练交替进行:先由导师给出样本权重,学生据此加权训练;随后根据学生在新数据上的梯度重新生成导师的训练数据,再更新导师。通过这种"以梯度为媒介"的协同机制,框架在不需要人工标注噪声标签的前提下,自动学习如何区分干净样本与错误标签样本。

从源码结构看,项目按职责划分为三层(training/trainers/ 下的trainer.pytrainer_triplet.py是两套训练循环实现,training/datasets/ 负责数据加载与污染,training/loss/triplet_loss.py 提供三元组损失,training/utils.py 提供梯度快照与自定义回调),入口统一收敛到 main.py。

二、环境准备与依赖安装

依赖清单见 requirements.txt,核心包包括:

依赖用途
tensorflow>=2.0模型构建与训练(代码使用tf.keras+ eager execution)
absl-py命令行参数解析(absl.flags/absl.app
tensorflow-datasetssoftmax 模式加载 MNIST
tensorflow-probability数据集重采样与统计相关功能
numpy / pandas / scipy / scikit-learn数值计算、CSV 读取、稀疏矩阵与线性回归(triplet 模式)
scann近邻检索(triplet 模式三元组挖掘的预留依赖)
PillowCSV 模式图像读取

run.sh 给出了一个"零手工干预"的端到端环境搭建脚本:先virtualenv -p python3 .在当前目录创建虚拟环境并激活,再pip install -r .../requirements.txt安装依赖,随后直接运行训练命令并在结束后清理临时目录。需要说明的是,该脚本内部引用的模块名与仓库实际目录名存在差异,官方 README 推荐的运行方式仍是下述python -m命令(要求从仓库根目录google-research/执行,以保证模块路径可被 Python 解析)。

三、快速开始:官方示例命令

README 给出的最小可用命令如下(从仓库根目录执行):

# From google-research/ python -m student_mentor_dataset_cleaning.main --save_dir=/tmp/models \ --student_epoch_count=1 --mentor_epoch_count=1

该命令的含义是:

  • --save_dir=/tmp/models:指定模型检查点保存目录;
  • --student_epoch_count=1:每个迭代轮次中学生只训练 1 个 epoch;
  • --mentor_epoch_count=1:每个迭代轮次中导师只训练 1 个 epoch。

由于max_iteration_count默认值为 20,完整运行会执行最多 20 轮"训练学生→训练导师"的交替迭代(但会在验证损失连续 20 轮不改善时提前终止,见下文训练循环一节)。不指定--mode时默认进入softmax模式,自动通过tensorflow-datasets加载 MNIST,无需准备任何本地数据文件。

四、命令行参数全面解析

所有参数均在 main.py 中通过absl.flags定义,并在verify_arguments()(main.py)中做合法性校验。完整参数表如下:

参数默认值取值范围/说明
--mini_batch_size32整数,须为正。学生与导师共享的 mini-batch 大小
--max_iteration_count20整数,须为正。学生-导师交替训练的最大轮次数
--student_epoch_count30整数,须为正。每轮中学生最多训练的 epoch 数
--mentor_epoch_count30整数,须为正。每轮中导师最多训练的 epoch 数
--modesoftmaxsoftmaxtriplet。softmax 模式固定使用 MNIST 并忽略csv_path;triplet 模式使用 CSV 指定的数据集
--save_dir''模型保存目录路径
--tensorboard_log_dir''TensorBoard 日志目录路径(为空则不写日志)
--train_dataset_dir''训练图像所在目录(triplet 模式使用)
--csv_path''训练数据 DataFrame/CSV 文件路径(triplet 模式使用)
--student_initial_model''学生模型初始化路径(当前入口中已定义但暂未接入训练流程)
--delg_embedding_layer_dim2048triplet 模式中损失层的 embedding 维度

五、softmax 模式:MNIST 上的默认实验

run_softmax()(main.py)演示了框架的完整装配方式:

学生模型是一个简单的全连接网络:Flatten(28×28) → Dense(128, relu) → Dense(10),使用 Adam(lr=0.001)优化器,损失为SparseCategoricalCrossentropy(from_logits=True),并挂载了 Top-1~Top-4 准确率与交叉熵等指标。

导师模型是一个二元分类网络:Flatten(101770) → Dense(50, relu) → Dense(1, sigmoid)。输入维度 101770 与学生模型全部可训练参数被展平后的梯度向量长度一致(梯度快照见下文),输出经 sigmoid 归一化为 (0,1) 区间的"样本可信度",使用BinaryCrossentropy训练,并记录 BinaryAccuracy、FalseNegatives、FalsePositives、TrueNegatives、TruePositives 等指标。

之后调用trainer.train(...)进入交替训练主循环。

六、triplet 模式:基于 CSV 数据集的度量学习

--mode=triplet时,程序进入 run_triplet()(实现位于 trainer_triplet.py,其模块注释明确标注still work-in-progress):

  • 学生模型换为ResNet152V2include_top=False、ImageNet 预训练权重、输入 321×321×3、pooling='avg'),损失为自定义 TripletLoss;
  • 导师模型结构不变,但输入维度变为 104000(对应 triplet 梯度快照的展平长度,见 utils.py);
  • 数据通过 CsvDataset 从 CSV 加载:图像按base_dir/x/y/z/id.jpg的三级目录结构存放(x/y/z 为 image id 的前三个字符),CSV 默认第 0 列为 image id、第 2 列为标签,加载时按标签排序并截取前 1996 条。

TripletLoss 的关键参数(triplet_loss.py)包括:embedding_size(embedding 维度)、triplet_loss_margin(默认 0.1)、train_ratio(默认 0.1,用于切分近邻索引的训练子集)、num_partitions(默认 1000,乘积量化分区数)、num_neighbors(默认 100)、anchor_reuse_count_max(默认 20,每个 anchor 最多复用的三元组数)等。损失计算时按"easy positive + hard negative"策略挖掘三元组,并用半硬三元组损失relu(‖a−p‖−‖a−n‖+margin)(triplet_semihard_loss_fn,margin 默认 1.0)累计。源码中近邻检索(internal_get_nearest_neighbors)暂以pass占位,等待接入 scann 检索器,属实验性预留。

七、训练循环的底层原理

核心主循环位于 trainer.py 的 train(),每轮迭代执行以下步骤:

  1. 计算样本权重_get_weights_dataset()(trainer.py)用utils.get_gradients_dataset_from_labelled_data计算学生在训练集上的梯度快照,再逐条交给导师模型map(mentor),输出即每条样本的置信度权重;
  2. 重置并训练学生:从save_dir/student/init.hdf5重新加载初始学生(_reinitilize_student),将(x, y, 权重)三元组数据传入student.fit(...)加权训练student_epoch_count个 epoch;
  3. 重置并训练导师_create_mentor_dataset()(trainer.py)对导师训练集施加噪声污染(corrupt_dataset),把"未被修改的样本"标为 1、"标签被随机化的样本"标为 0,作为导师的监督信号;随后从mentor/init.hdf5重新加载初始导师并训练mentor_epoch_count个 epoch,训练时通过class_weight={0: 1-noise_rate, 1: noise_rate}平衡正负类;
  4. 早停与保优:若本轮导师验证损失优于历史最优,则保存best.hdf5并重置等待计数;否则waiting += 1,超过patience=20即终止整个迭代。

代码中硬编码的默认污染参数为noise_rate=0.1(10% 的标签被随机化)、target_distribution_parameter=0.01(指数分布重采样的陡峭度,控制类不平衡程度)。训练过程中的关键回调(trainer.py)包括:学生EarlyStopping(patience=60)、导师EarlyStopping(patience=100)、ReduceLROnPlateau(factor=0.5、min_lr=1e-7)以及按 epoch 保存权重的ModelCheckpointweights.{epoch:04d}.hdf5)。其中CustomEarlyStoppingCustomReduceLROnPlateau(utils.py)重写了on_train_begin为空操作,目的是让状态在多次fit调用间不被重置。

梯度快照机制(utils.py):对每条样本,在tf.GradientTape内前向计算损失,然后对学生的全部可训练权重求梯度并展平拼接,得到一个一维向量——这正是导师模型的输入特征。triplet 变体(utils.py)则先生成三元组,再对每个三元组计算半硬三元组损失的梯度,通过from_generator构造输出形状为[104000]的梯度数据集。

八、数据准备与噪声注入机制

softmax 模式的数据由 mnist.create_dataset() 生成:将 MNIST 训练集按0.4 / 0.1 / 0.5划分为学生训练集、学生验证集、导师训练集三份,图像归一化到 [0,1],并只对学生训练集施加污染。

污染流水线封装在 datasets.corrupt_dataset() 中,包含两个环节:

  • 类不平衡重采样imbalanced_sample_exponential,datasets/init.py):按指数分布对训练样本做过采样,target_distribution_parameter=0时近似均匀采样;
  • 标签噪声add_noise,datasets/init.py):以noise_rate概率将样本标签随机替换为另一个类别,并将该样本的权重置 0——被噪声污染的样本因此天然带上了"低权重"标签,供导师学习。

此外dataset_split()(datasets/init.py)按随机阈值将导师数据集按 0.6/0.4 切分为训练与验证两部分。

九、模型产物与 TensorBoard 日志

训练产出统一写入--save_dir,目录结构为:

save_dir/ ├── student/ │ ├── init.hdf5 # 初始学生(每轮重置回此状态) │ ├── best.hdf5 # 历史最优学生 │ └── iteration_0000/ # 每轮的按 epoch 权重 │ └── weights.0000.hdf5 └── mentor/ ├── init.hdf5 # 初始导师 ├── best.hdf5 # 历史最优导师 └── iteration_0000/ └── weights.0000.hdf5

若指定--tensorboard_log_dir,则会在其下按student/iteration_XXXXmentor/iteration_XXXX生成 TensorBoard 日志(histogram_freq=1),并通过LearningRateLogger(utils.py)额外记录每个 epoch 结束时的学习率标量,便于在 TensorBoard 中观察学习率衰减与训练曲线。注意训练开始时trainer.train()会对save_dirlog_dir执行shutil.rmtree清理,重复运行同一目录前请确认无需保留旧产物。

十、注意事项与扩展建议

  1. 运行目录python -m方式要求从仓库根目录google-research/执行,否则模块student_mentor_dataset_cleaning无法被导入;
  2. 快速验证:官方示例将两个 epoch 参数设为 1,配合默认max_iteration_count=20即可在几分钟内跑通完整交替训练并观察早停行为;如需彻底快速冒烟,也可进一步调小--max_iteration_count
  3. 模式选择:softmax 模式开箱即用(自动下载 MNIST);triplet 模式需要自行准备train_dataset_dir+csv_path数据,且当前实现标注为 work-in-progress,近邻检索逻辑为占位实现;
  4. 依赖版本:代码基于 TF2 eager 模式(入口处调用tf.compat.v1.enable_eager_execution()),建议在 TensorFlow 2.x 环境中运行;requirements.txttensorflow>=2.0z的写法按 pip 语义会被视为>=2.0处理;
  5. 可扩展方向--student_initial_model参数已预留但尚未接入训练流程,若需从预训练权重启动学生,可从 main.py 与trainer.train()的调用处入手扩展;噪声率与重采样参数当前为代码内硬编码(noise_rate=0.1target_distribution_parameter=0.01),可在训练器实现中按需调整以适配不同噪声强度的数据集。
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 17:48:48

识别图片文字的软件性能优化实战与最佳实践指南

识别图片文字的软件性能优化实战与最佳实践指南 上周陪一个做外包的后端兄弟面大厂,面试官甩了张带噪点的物流单图片,问:“你的OCR接口P99延迟突然飙到800ms,怎么排查?”他愣了五秒,支支吾吾说“可能是图片太大”。面试官摇头走了。这场景太典型了,很多开发盯着业务逻辑写,一碰到【识别图片文字的软件】…

作者头像 李华
网站建设 2026/9/23 17:48:33

无线网怎么修改密码源码解析:3步搞定底层逻辑

无线网怎么修改密码源码解析:3步搞定底层逻辑 官方文档往往冗长且晦涩,让你抓不住重点。其实,无线网怎么修改密码的核心在于理解WPA2加密协议的密钥派生机制。通过源码解析,你能看清密码变更背后的数据流转。…

作者头像 李华
网站建设 2026/9/23 17:48:25

景气指数编制全流程:从指标筛选到合成计算与验证维护

1. 景气指数到底在测什么景气指数这个词,乍一听挺唬人,其实说白了就是给经济或行业的"体温"量个体温。它不直接告诉你GDP涨了多少,而是通过一组先行、同步、滞后指标的组合,判断当前经济处于扩张还是收缩区间&#xff0…

作者头像 李华
网站建设 2026/9/23 17:48:15

3个坑让你少加班,blackcock保姆级教程

3个坑让你少加班,blackcock保姆级教程 代码从网上复制下来,本地一跑直接报错?别急着删库,这往往是环境或版本不对。很多新手卡在“为什么我这边行,他那边不行”的循环里,其实90%的问题出在依赖解析和配置细节上。今天这篇blackcock保姆级教程,专门拆解那些让你深夜抓狂的隐性Bug,帮你把调…

作者头像 李华