- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
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.py与trainer_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-datasets | softmax 模式加载 MNIST |
| tensorflow-probability | 数据集重采样与统计相关功能 |
| numpy / pandas / scipy / scikit-learn | 数值计算、CSV 读取、稀疏矩阵与线性回归(triplet 模式) |
| scann | 近邻检索(triplet 模式三元组挖掘的预留依赖) |
| Pillow | CSV 模式图像读取 |
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_size | 32 | 整数,须为正。学生与导师共享的 mini-batch 大小 |
--max_iteration_count | 20 | 整数,须为正。学生-导师交替训练的最大轮次数 |
--student_epoch_count | 30 | 整数,须为正。每轮中学生最多训练的 epoch 数 |
--mentor_epoch_count | 30 | 整数,须为正。每轮中导师最多训练的 epoch 数 |
--mode | softmax | softmax或triplet。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_dim | 2048 | triplet 模式中损失层的 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):
- 学生模型换为
ResNet152V2(include_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(),每轮迭代执行以下步骤:
- 计算样本权重:
_get_weights_dataset()(trainer.py)用utils.get_gradients_dataset_from_labelled_data计算学生在训练集上的梯度快照,再逐条交给导师模型map(mentor),输出即每条样本的置信度权重; - 重置并训练学生:从
save_dir/student/init.hdf5重新加载初始学生(_reinitilize_student),将(x, y, 权重)三元组数据传入student.fit(...)加权训练student_epoch_count个 epoch; - 重置并训练导师:
_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}平衡正负类; - 早停与保优:若本轮导师验证损失优于历史最优,则保存
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 保存权重的ModelCheckpoint(weights.{epoch:04d}.hdf5)。其中CustomEarlyStopping与CustomReduceLROnPlateau(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_XXXX、mentor/iteration_XXXX生成 TensorBoard 日志(histogram_freq=1),并通过LearningRateLogger(utils.py)额外记录每个 epoch 结束时的学习率标量,便于在 TensorBoard 中观察学习率衰减与训练曲线。注意训练开始时trainer.train()会对save_dir与log_dir执行shutil.rmtree清理,重复运行同一目录前请确认无需保留旧产物。
十、注意事项与扩展建议
- 运行目录:
python -m方式要求从仓库根目录google-research/执行,否则模块student_mentor_dataset_cleaning无法被导入; - 快速验证:官方示例将两个 epoch 参数设为 1,配合默认
max_iteration_count=20即可在几分钟内跑通完整交替训练并观察早停行为;如需彻底快速冒烟,也可进一步调小--max_iteration_count; - 模式选择:softmax 模式开箱即用(自动下载 MNIST);triplet 模式需要自行准备
train_dataset_dir+csv_path数据,且当前实现标注为 work-in-progress,近邻检索逻辑为占位实现; - 依赖版本:代码基于 TF2 eager 模式(入口处调用
tf.compat.v1.enable_eager_execution()),建议在 TensorFlow 2.x 环境中运行;requirements.txt中tensorflow>=2.0z的写法按 pip 语义会被视为>=2.0处理; - 可扩展方向:
--student_initial_model参数已预留但尚未接入训练流程,若需从预训练权重启动学生,可从 main.py 与trainer.train()的调用处入手扩展;噪声率与重采样参数当前为代码内硬编码(noise_rate=0.1、target_distribution_parameter=0.01),可在训练器实现中按需调整以适配不同噪声强度的数据集。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
PokemonRedExperiments训练数据清洗:delete_empty_imgs.txt使用指南
PokemonRedExperiments训练数据清洗:delete_empty_imgs.txt使用指南 在使用强化学习(Reinforcement Lear
强化学习深度学习AI应用Gramophone性能优化:如何构建高效的Baseline Profile
Gramophone性能优化:如何构建高效的Baseline Profile Gramophone是一款严格遵循Android标准,采用media3和Mater
如何快速实现YOLOv7训练数据清洗:噪声样本自动检测与过滤方法
如何快速实现YOLOv7训练数据清洗:噪声样本自动检测与过滤方法 YOLOv7作为当前最先进的实时目标检测算法,其模型性能高度依赖训练数据的质量。噪声样本的存在
人工智能深度学习计算机视觉
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考