- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
本指南以 data_free_distillation/README.md 为核心,系统拆解 google-research 仓库中"大规模生成式无数据蒸馏"(Large-Scale Generative Data-Free Distillation)论文实验实现的技术骨架:包括标签条件生成器(数据来源)、ResNet 教师/学生网络(resnet.py 与 generators.py)、模型统一接入层(models.py)以及单元测试与运行流程(run.sh)。读完本文,你将掌握该实验实现的目录结构、生成器与 ResNet 的完整网络参数与结构设计、测试验证方式以及本地复现运行的方法。
一、主题定位:什么是"生成式无数据蒸馏"
传统知识蒸馏(Knowledge Distillation)依赖一个关键前提:能够访问教师模型训练时使用的原始数据集。但在真实场景中,原始数据往往因隐私、版权或存储成本无法获取,此时只能借助"无数据蒸馏"(Data-Free Distillation)在不访问原始训练数据的前提下完成知识迁移。
"大规模生成式无数据蒸馏"采取生成式(Generative)路线:训练一个生成器网络,在给定随机噪声与类别标签的条件下合成与教师网络数据分布一致的"伪样本",再将这些伪样本用于训练学生网络。从本仓库源码结构看,这一思路被完整落地为三大部分:
- 生成器:generators.py 负责合成标签条件化的图像;
- 教师/学生网络:resnet.py 提供 ResNet-18/34/50/101 系列分类网络,models.py 的模块注释明确说明其用途为"为 student 和 teacher 提供模型访问";
- 验证与运行:resnet_test.py 与 run.sh 负责结构正确性验证与测试执行。
需要说明的前提是:本仓库作为该论文的实验实现,当前目录下主要包含生成器、网络定义与单元测试,用于复现与验证论文中的核心网络结构;仓库本身是只读的,你可以在本地安装依赖、运行测试进行验证。
二、仓库结构与运行环境
整个实验位于仓库根目录下的data_free_distillation/子目录,结构如下:
data_free_distillation/ ├── README.md # 论文实验实现说明与引用信息 ├── requirements.txt # 依赖清单:tensorflow、tf-slim ├── run.sh # 单元测试一键运行脚本 └── main/ ├── generators.py # 生成器模型(simple / conditioned / generator) ├── models.py # 教师与学生模型的统一接入层 ├── resnet.py # ResNet-18/34/50/101 网络定义 └── resnet_test.py # ResNet 结构与形状单元测试依赖说明(requirements.txt)
requirements.txt 仅声明两个依赖:
tensorflow:框架本体,代码中统一通过tensorflow.compat.v1导入,即按TF1 风格编写(如tf.variable_scope、tf.nn.leaky_relu),测试入口还显式调用tf.disable_v2_behavior()关闭 TF2 行为;tf-slim:提供slim.conv2d、slim.batch_norm、slim.arg_scope等高层封装,以及tf_slim.nets中的resnet_v1.utils、resnet_utils工具函数。
三、数据来源:标签条件生成器(generators.py)
生成器是整个无数据蒸馏流程中"数据供给"的核心。文件 generators.py 实现了三个层层封装的函数:simple_generator→conditioned_generator→generator。
3.1 simple_generator:基础生成网络
simple_generator是论文中使用的简单生成器模型(generators.py),其结构设计参考了 XNOR-Net 论文(代码注释中给出出处 arXiv:1603.05279)。关键签名与默认值如下:
def simple_generator(z, # 随机噪声向量 image_size, # 输出图像边长 num_interpolate = 2, # 上采样插值次数 channels = None, # 各插值层输出通道数 depthwise_separate = None, # 是否使用深度可分离卷积 output_bn = True, # 输出层后是否接 BatchNorm is_training = True, # BN 是否处于训练模式 reuse=None, scope=None):默认通道与结构:当不显式传入channels时,默认按channels = [128 // (i + 1) for i in range(num_interpolate)]生成,即num_interpolate=2时默认通道为[128, 64];depthwise_separate默认全为False(不使用深度可分离卷积)。两者的长度必须与num_interpolate一致(代码中有断言校验)。
网络前向流程(结合 generators.py):
- 初始尺寸:
init_size = image_size // (2**num_interpolate),即经过num_interpolate次×2 的上采样后恢复到目标分辨率; - 全连接升维:
slim.fully_connected将噪声z映射为init_size × init_size × channels[0]维向量(无激活、无偏置),再 reshape 成[-1, init_size, init_size, channels[0]]的特征图; - 首个 BN 层:
bn_0使用epsilon=1e-5。代码注释说明:这是为了与 DAFL 论文(Data-Efficient Model Compression)的 BN 超参保持一致以复现结果,随后接leaky_relu激活; - 插值上采样循环:每一轮先通过最近邻插值(
tf.image.resize+ResizeMethod.NEAREST_NEIGHBOR)将特征图尺寸翻倍,再执行 3×3 卷积(若depthwise_separate=True则拆成 3×3 depthwise + 1×1 pointwise 两步); - 输出层:3×3 卷积将通道数压缩为3(RGB 图像),激活函数为
tanh,不接 normalizer; - 可选输出 BN:
output_bn=True时在输出后再接一层center=False, scale=True的 BatchNorm。
全局 arg_scope 超参:生成器内的 BatchNorm 统一设置为decay=0.9, center=True, scale=True, epsilon=0.8;卷积与可分离卷积统一使用leaky_relu激活 + BatchNorm 归一化。
3.2 conditioned_generator:标签条件化
def conditioned_generator(z, one_hot_label, image_size, ...): with tf.variable_scope(scope, 'conditioned_generator', [z, one_hot_label], reuse=reuse): z = tf.concat([z, one_hot_label], axis=1) # 噪声与标签沿通道拼接 return simple_generator(z, image_size, ...)conditioned_generator的核心操作是把随机噪声z与 one-hot 标签沿axis=1拼接后送入simple_generator,从而让生成器能够按类别合成图像(generators.py)。
3.3 generator:蒸馏脚本的对外入口
def generator(z, label, image_size, n_classes, ...): one_hot_label = tf.one_hot(label, n_classes) return conditioned_generator(z, one_hot_label, image_size, ...)generator是生成器训练与蒸馏脚本中最常调用的函数(代码注释原文:"This is the function we would typically use in generator training and distillation script")。它接收整数标签label,内部先通过tf.one_hot(label, n_classes)转为 one-hot 向量,再委托给conditioned_generator,并以scope='generator'作为默认命名空间(generators.py)。
可以推断,无数据蒸馏的基本数据流为:随机噪声 + 类别标签 → 生成器 → 合成的类别图像 → 教师/学生网络,生成器由此充当原始数据集的替代品。
四、教师与学生网络:ResNet 系列(resnet.py)
resnet.py 提供了 ResNet-18/34/50/101 四种规格的完整实现,同时承担教师网络与学生网络的网络定义职责。文件开头 docstring 明确指出一个重要设计差异:
该模块中的网络面向 CIFAR-10 数据集训练,但为了复现论文 [2](DeepInversion: Dreaming to Distill,arXiv:1912.08795)的结果,其结构与原始 ResNet [1] 在 CIFAR-10 上的实现不同——原始版本下采样 3 次,而这里下采样4 次,与 ImageNet 上的网络结构类似。
4.1 两种残差单元:basic_block 与 bottleneck
- basic_block(resnet.py):标准基础块,包含两个 3×3 卷积(第二个不带激活);当
stride != 1或输入/输出通道数不一致时,shortcut 使用 1×1 卷积(stride与残差路径一致),否则 shortcut 直接取输入;输出为relu(shortcut + residual); - bottleneck(resnet.py):1×1 卷积降维 → 3×3
conv2d_same卷积 → 1×1 卷积升维(第三个不带激活),同样在维度/步长变化时插入 1×1 shortcut。
4.2 resnet() 核心函数:根块与全局池化
resnet()(resnet.py)是通用骨架,几个关键行为:
- 根块自适应选择:
_use_small_root_block(inputs)依据输入图像边长自动选择——尺寸 ≤ 64 时使用 3×3、stride=1 的"小根块"(适配 32×32 的 CIFAR);否则使用 7×7、stride=2 的"大根块"并接 3×3、stride=2 的 max pooling(适配 ImageNet)。_skip_first_max_pooling则针对 128×128 输入跳过首个 max pooling; - 残差块堆叠:通过
resnet_utils.stack_blocks_dense完成; - 全局平均池化:
global_pool=True时执行tf.reduce_mean(net, axis=[1, 2], keepdims=True),输出记为end_points['global_pool']; - 分类头:1×1 卷积输出
num_classes通道 logits,并附slim.softmax的predictionsend point; - BatchNorm 模式:通过
is_training控制slim.batch_norm的训练/推理状态。
4.3 低分辨率 ImageNet 的 stride 策略:skip_first_n_strides
_create_blocks(resnet.py)以depths = [64, 128, 256, 512]构建四个残差块,并支持skip_first_n_strides参数调整前 N 个块的下采样步长,代码注释给出了完整的适用性对照表(面向低分辨率 ImageNet 输入):
| skip_first_n_strides | 四个块的 stride 方案 | 适用输入尺寸 |
|---|---|---|
| 0 | [1, 2, 2, 2] | 56 或 64 |
| 1 | [1, 1, 2, 2] | 28 或 32 |
| 2 | [1, 1, 1, 2] | 14 或 16 |
| 3 | [1, 1, 1, 1] | 7 或 8 |
该参数有0 ≤ skip_first_n_strides ≤ 3的断言约束,为复现不同分辨率下的训练结果提供了灵活性。
4.4 四种规格的默认超参
四种网络均提供统一参数接口,默认值如下(以resnet_18为例,resnet.py):
def resnet_18(inputs, num_classes = None, is_training = True, global_pool = True, weight_decay = 5e-4, # 权重衰减 batch_norm_decay = 0.9, # BN 滑动平均衰减 skip_first_n_strides = 0, reuse=None, scope='resnet_18'):各规格的 block 结构与单元数配置如下:
| 模型 | 每块单元数 | 残差单元类型 | scope 默认值 |
|---|---|---|---|
| resnet_18 | [2, 2, 2, 2] | basic_block | resnet_18 |
| resnet_34 | [3, 4, 6, 3] | basic_block | resnet_34 |
| resnet_50 | [3, 4, 6, 3] | bottleneck | resnet_50 |
| resnet_101 | [3, 4, 23, 3] | bottleneck | resnet_101 |
其中 bottleneck 变体的输出深度为depth * 4(bottleneck 内部通道为depth),首层卷积统一为conv1_depth=64,所有变体均通过resnet_utils.resnet_arg_scope(weight_decay=..., batch_norm_decay=...)设置全局卷积与 BN 超参。四种规格内部通过model_fn(architecture)(resnet.py)完成名字到构造函数的映射,不支持的名字会触发断言错误。
五、模型统一接入层(models.py)
models.py 是教师/学生模型的外部统一入口,模块 docstring 为 "Proves access to models for student and teacher":
def model_fn(model_name): if model_name.startswith('resnet'): return resnet.model_fn(model_name) raise RuntimeError('Unsupported model: %s' % model_name)其设计意图清晰:以字符串模型名驱动网络构建。只要模型名以resnet开头即委托给 resnet.py 的model_fn(支持resnet_18/34/50/101),其余名称一律抛出RuntimeError。这意味着在蒸馏流程中,教师与学生网络可以共用同一套model_fn机制按需实例化,便于后续扩展其他架构。
六、单元测试与结构验证(resnet_test.py)
resnet_test.py 继承tf.test.TestCase,对 ResNet 的结构正确性做了系统性验证,是理解网络行为的"可运行说明书":
resnet_small:测试专用的浅薄网络(block1深度 2 /block2深度 4 /block3深度 8 /block4深度 16),便于快速构建与断言;testClassificationEndpoints:验证logits形状为[batch, 1, 1, num_classes]、predictions与global_poolend point 存在且形状正确;testEndpointNames/testEndpointNamesWithBottleneckBlock:分别验证 basic_block 与 bottleneck 两种结构下 end point 命名集合(如resnet/blockN/unit_M/basic_block/conv1、shortcut等)与预期完全一致;testClassificationShapes/testClassificationShapesWithBottleneckBlock:验证各 block 输出特征图的空间尺寸随下采样逐步减半(如block1: [2,32,32,...] → block4: [2,4,4,...]);testShapesWithInputSize128x128/testShapesWithInputSize256x256:验证大输入尺寸下根块与 pooling 的行为(128×128 跳过首个 max pooling);testSkipStrideShapes:用resnet_18+skip_first_n_strides=1验证 stride 跳过逻辑(此时block2保持 32×32 分辨率);testUnknownBatchSize:验证动态 batch 维度(占位符)下前向传播可用。
测试入口统一tf.disable_v2_behavior()后调用tf.test.main(),再次印证其运行环境为 TensorFlow 1.x 兼容模式。
七、本地运行与验证(run.sh)
run.sh 提供了开箱即用的验证流程:在$PWD下创建 Python 3 虚拟环境、安装依赖并运行 ResNet 单元测试:
#!/bin/bash set -e set -x virtualenv -p python3 . # 在当前目录创建虚拟环境 source ./bin/activate # 激活虚拟环境 pip install -r data_free_distillation/requirements.txt python -m data_free_distillation.main.resnet_test执行要点:
- 脚本要求环境中有
virtualenv工具与 Python 3 解释器; - 虚拟环境创建在当前工作目录(
$PWD)下,因此建议在仓库根目录执行bash data_free_distillation/run.sh; set -e保证任一步失败即中止,便于快速定位依赖或环境问题;- 最终通过
python -m data_free_distillation.main.resnet_test以模块方式运行测试,前提是仓库根目录在PYTHONPATH中(data_free_distillation.main的包路径依赖仓库根目录),这也与 models.py 中from data_free_distillation.main import resnet的导入方式一致。
八、论文引用
如果你在研究中使用了本仓库的代码,README 提供了官方推荐引用格式(BibTeX):
@article{Luo2020DataFreeDistill, author = {Luo, Liangchen and Sandler, Mark and Lin, Zi and Zhmoginov, Andrey and Howard, Andrew}, title = {Large-Scale Generative Data-Free Distillation}, journal = {arXiv preprint arXiv:2012.05578}, year = {2020} }该引用指向论文Large-Scale Generative Data-Free Distillation(arXiv:2012.05578),本仓库即其配套实验实现。
小结
围绕 data_free_distillation/README.md 所定义的实验主题,本仓库给出了一个结构清晰的生成式无数据蒸馏实现骨架:以 generators.py 的标签条件生成器作为"伪数据"来源,以 resnet.py + models.py 提供教师/学生 ResNet 网络,并以 resnet_test.py + run.sh 保证网络结构可验证、可复现。对研究无数据蒸馏、生成式知识迁移以及 TensorFlow 1.x + tf-slim 网络工程化的开发者而言,这是一个可直接阅读源码、运行测试并在此基础上扩展完整蒸馏训练流程的可靠起点。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
Google Research Subclass Distillation 实战解析:基于 MNIST Colab 的子类蒸馏复现与损失函数原理
Google Research Subclass Distillation 实战解析:基于 MNIST Colab 的子类蒸馏复现与损失函数原理 导读:本文围绕
人工智能深度学习NLP计算机视觉强化学习Wand-Enhancer完整教程:Wand专业版本地解锁工具,三步搞定
Wand Enhancer完整教程:Wand专业版本地解锁工具,三步搞定 关键对局正打得火热,屏幕上突然弹出"2小时已用完",刚调好的数值当场清零,想继续就得掏
桌面应用前端Amazon Bedrock 模型蒸馏(Model Distillation)实战指南:从 JSONL 训练数据与历史调用日志到蒸馏模型部署
Amazon Bedrock 模型蒸馏(Model Distillation)实战指南:从 JSONL 训练数据与历史调用日志到蒸馏模型部署 本文基于 amaz
示例工程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考