DEep MOdel GENeralization dataset(DEMOGEN):用 756 个真实训练模型研究深度网络泛化差距 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载DEMOGENDEep MOdel GENeralization dataset是 Google Research 公开的一个模型泛化数据集它提供 756 个在 CIFAR-10 / CIFAR-100 上真实训练好的深度模型及其完整的训练/测试性能记录用于支撑 ICLR 2019 论文《Predicting the Generalization Gap in Deep Networks with Margin Distributions》的研究复现。读完本文你将掌握如何在当前仓库中下载并加载这批预训练模型、通过ModelConfig精确索引任意模型变体、复现论文中的 margin边际分布与 total variation总变差两类泛化指标计算以及如何批量评估模型精度来验证泛化差距实验。数据集概述什么是 DEMOGENDEMOGEN 的核心思想是把泛化研究从理论推导向实证数据集转变与其从零训练大量模型来观察它们的泛化行为不如直接提供一批已经训练完毕、覆盖多种超参数组合的模型让研究者聚焦于分析什么因素与泛化差距相关。该数据集包含756 个训练好的深度模型每个模型附带其在训练集与测试集上的完整性能记录准确率与交叉熵两个经典基准数据集CIFAR-1010 类与CIFAR-100100 类两类网络架构变体结构类似Network-in-NetworkNIN的 CNN以及ResNet-32每种模型使用不同的正则化技术与超参数设置从而在泛化行为上产生宽谱分布。正如 demogen/README.md 所述NIN 模型在 CIFAR-10 上训练后测试准确率从 60% 一直延伸到 90.5%泛化差距训练与测试准确率之差从 1% 到 35% 不等——这为研究什么预测了泛化差距提供了足够大的观测空间。变体空间覆盖实践中最常见的正则化手段DEMOGEN 的 756 个模型并不是随机生成而是围绕研究者日常最常用的调参手段系统化构建的。根据 model_config.py 中的ALL_MODEL_SPEC常量可以精确还原每个实验族的参数网格实验族变体维度取值NIN_CIFAR10宽度倍率wide_multiplier1.0 / 1.5 / 2.0批归一化batchnormTrue / FalseDropout 概率dropout_prob0.0 / 0.2 / 0.5数据增强augmentationTrue / FalseL2 权重衰减decay_fac0.0 / 0.001 / 0.005副本copy1 / 2RESNET_CIFAR10宽度倍率1.0 / 2.0 / 4.0归一化方式batch / group数据增强True / FalseL2 权重衰减0.0 / 0.02 / 0.002初始学习率0.01 / 0.001副本1 / 2 / 3RESNET_CIFAR100宽度倍率1.0 / 2.0 / 4.0归一化方式batch / group数据增强True / FalseL2 权重衰减0.0 / 0.02 / 0.002初始学习率0.1 / 0.01 / 0.001副本1 / 2 / 3这些变化手段包括不同强度的权重衰减与 dropout、是否使用批归一化ResNet 额外提供组归一化、是否使用数据增强、隐层宽度或隐藏单元数量、以及针对 ResNet 的不同初始学习率。宽泛的超参数扫描保证了数据集中模型之间存在可量化的泛化行为差异。模型目录命名即超参数索引DEMOGEN 的一个巧妙设计是每个模型的存储目录名直接编码了它的全部超参数因此无需任何数据库即可按配置精确寻址。在 model_config.py 的get_model_dir_name中可以看到两种命名规则NIN 模型目录形如NIN_CIFAR10/nin_wide_1.5x_bn_dropout_0.2_aug_decay_0.001_1各片段依次为模型类型、宽度、是否 BN、dropout 概率、是否增强、衰减系数与副本序号ResNet 模型目录形如RESNET_CIFAR100/resnet_wide_2.0x_groupnorm_aug_decay_0.002_2且仅当学习率不等于默认值 0.01 时才追加lr_xxx片段。快速上手加载与评估一个预训练模型仓库提供了开箱即用的示例脚本 example.py通过python -m demogen.example即可运行运行前需设置好数据集根目录。该脚本演示了完整的加载—推理—评估链路model_config mc.ModelConfig( model_typenin, datasetcifar10, root_dirroot_dir) load_and_run(model_config, root_dir) eval_result evaluate_model(model_config, root_dir) print(Test Accuracy: {}.format(eval_result)) print(Stored Test Accuracy: {}.format(model_config.test_stats())) print(Stored Train Accuracy: {}.format(model_config.training_stats()))其中load_and_run展示了最基本的用法由配置生成 checkpoint 路径 → 构建模型函数 → 建立 Session → 灌入输入张量 → 恢复参数 → 前向推理evaluate_model则以 batch_size500 循环 20 个 batch在 10000 张测试图上统计准确率并把实测结果与数据集自带的eval.json/train.json中的存档精度model_config.test_stats()/training_stats()进行对照。ModelConfig数据集的门面接口ModelConfig是访问整个数据集的核心入口其构造参数与前述变体空间一一对应见 model_config.py。需要留意几个实现细节model_type仅支持nin与resnetdataset仅支持cifar10/cifar100传入组合必须存在于ALL_MODEL_SPEC中否则断言失败data_format由模型类型自动决定NIN 使用HWC通道在最后ResNet 使用CHW通道在前每个模型目录下都存有train.json与eval.json通过training_stats()/test_stats()可直接读取存档的训练/测试 Accuracy 与 CrossEntropy无需自己跑一遍推理checkpoint 文件统一命名为model.ckpt-150000见 model_config.py 的CKPT_NAME常量。从底层实现看get_model_fn()会将配置转换为tf.contrib.training.HParams再交给 models/get_model.py 分发NIN 的宽度按192 × wide_multiplier计算ResNet-32 的过滤器数按16 × wide_multiplier计算、并按(32-2)/6推导出每组 5 个 block。参数恢复则通过 load_parameters 在模型 scope 内收集变量并用tf.train.Saver恢复实现。依赖与运行环境根据 requirements.txt本仓库基于TensorFlow 1.x≥1.11 且 2.0、tensor2tensor ≥ 1.11.0与 numpy数据输入层通过 tensor2tensor 的 problems 接口加载 CIFAR见 data_util.py。仓库还提供了 run.sh展示基于 Python 2 virtualenv 的完整安装运行流程virtualenv -p python2 . source ./bin/activate pip install -r demogen/requirements.txt python -m demogen.example注意代码大量使用tensorflow.compat.v1且依赖已迁移至 contrib 模块因此建议在 TensorFlow 1.x 或提供 compat 兼容层的环境中运行。计算边际分布Margin线性逼近到决策边界的距离论文的核心指标是边际分布margin distribution对每个样本测量从该样本到决策边界在指定隐层激活空间内的距离。仓库在 margin_utils.py 中实现了这一指标的线性近似其文档字符串给出了标准用法input_fn data_util.get_input( datamodel_config.dataset, data_formatmodel_config.data_format) margins compute_margin(input_fn, root_dir, model_config) input_margins margins[inputs] h1_margins margins[h1] h2_margins margins[h2] h3_margins margins[h3]compute_margin会遍历整个训练集默认dataset_size50000、batchsize50同时计算inputs、h1、h2、h3四个层的边际并存入字典返回。margin 的实现原理margin()函数margin_utils.py的计算逻辑值得细读它分四步完成构造对抗类对每个样本取 logits 的 top-2若最高分类别就是真实标签则第二高类别作为竞争类indices_c否则用最高类别。梯度方向由one_hot(labels) - one_hot(indices_c)决定——即把样本推向竞争类的方向计算分子真实类 logit 与竞争类 logit 之差values_true - values_c一次求出各层梯度用tf.gradients(logits, layer_activations, grad_ys)同时得到 logits 对inputs/h1/h2/h3四个激活的梯度归一化得距离用梯度范数支持 L2、L1 与无穷范数由dist_norm控制对分子归一化numerator / ||g||即该层激活空间下到决策边界的线性逼近距离。实现还内置了数值稳定性保护epsilon1e-6被加入梯度范数避免除零所有用于归一化的梯度都经过tf.stop_gradient截断确保只对 logits 求导而不影响梯度本身的计算。支持dist_norm为 0无穷范数、1L1、2L2其他取值会抛出ValueError。计算总变差Total Variation度量隐层响应的离散程度另一个辅助指标是总变差它刻画某个隐层在整个训练集上的响应activation分布形态由 total_variation_util.py 实现。文档字符串给出的典型用法h1_total_variation compute_total_variation(input_fn, root_dir, h1, model_config)其计算过程与 margin 类似同样支持inputs/h1/h2/h3作为layer参数但在数据收集上一次只计算一个层——注释明确说明这是出于内存限制的考虑。算法步骤如下见 total_variation_util.py遍历训练集把指定层在全部样本上的激活拼接为一个大张量将激活展平为[样本数, 特征数]沿特征维度计算标准差response_std对标准差的平方求和再开方得到未归一化的总变差最后除以样本总数完成归一化。由于总变差本质上是层响应对输入变化的敏感程度与 margin 结合使用可以分别从决策边界距离与特征响应方差两个角度刻画模型这正是论文中 margin distribution 系列分析所需的两类底层统计量。数据集下载与引用用于本代码库的完整模型数据集约15.57GB可以从 Google Cloud Storage 公开地址下载https://storage.googleapis.com/margin_dist_public_files/demogen_models.tar.gz下载地址与大小以 demogen/README.md 为准。解压后将其根目录作为ModelConfig的root_dir传入即可。该数据集源自论文《Predicting the Generalization Gap in Deep Networks with Margin Distributions》ICLR 2019作者 Yiding Jiang、Dilip Krishnan、Hossein Mobahi、Samy Bengio。若你的研究使用了这一数据集建议按论文作者给出的 BibTeX 条目引用inproceedings{ jiang2018predicting, title{Predicting the Generalization Gap in Deep Networks with Margin Distributions}, author{Yiding Jiang and Dilip Krishnan and Hossein Mobahi and Samy Bengio}, booktitle{International Conference on Learning Representations}, year{2019}, url{https://openreview.net/forum?idHJlQfnCqKX}, }同时本代码库并非 Google 官方产品README 中明确声明 This is not an official Google product使用过程中如遇问题可通过 GitHub issue 反馈。结语DEMOGEN 的价值在于把泛化差距预测这一课题从零散的独立实验沉淀为一个可复现、可检索、可批量化的研究基础设施756 个真实训练的模型覆盖了主流正则化手段的完整参数网格目录命名即超参数索引ModelConfig一行代码即可定位任意模型margin 与 total variation 工具则直接复现了论文核心指标。对于任何想要在泛化理论上做实证验证的研究者这套代码库提供了从数据到指标的完整闭环。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐如何突破大模型训练瓶颈annotated_deep_learning_paper_implementations 可扩展性研究指南如何突破大模型训练瓶颈annotated_deep_learning_paper_implementations 可扩展性研究指南 annotated_dee人工智能深度学习大模型NLP计算机视觉强化学习LoRA为什么只有语言模型被量化解读 GOT-OCR2_0-4bit 的混合精度设计智慧为什么只有语言模型被量化解读 GOT OCR2_0 4bit 的混合精度设计智慧 当我们打开 mlx community/GOT OCR2_0 4bit 这个Gradle Kotlin DSL Samples最佳实践总结避免常见陷阱的20个技巧Gradle Kotlin DSL Samples最佳实践总结避免常见陷阱的20个技巧 Gradle Kotlin DSL Samples是Gradle官方提上一篇终极指南如何用foobox美化配置打造专业级foobar2000音乐播放器下一篇终极Magicast问题解决方案从入门到精通的常见问题解决指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考