
SANE: Scalable and Versatile Weight Space Learning
本仓库包含论文 "Towards Scalable and Versatile Weight Space Learning" 的代码,该论文发表于 ICML 2024。这项工作引入了 SANE*,一种用于学习神经网络任务无关表示的新方法,该方法可扩展至更大的模型并适用于各种任务。论文可在此处找到:ICML proceedings | arxiv.
*Sequential Autoencoder for Neural Embeddings
摘要
学习训练良好的神经网络模型的表示,有望提供对这些模型内部工作机制的理解。然而,先前的研究在处理较大网络时面临局限性,或者仅针对判别式或生成式任务。本文介绍了用于权重空间学习的 SANE 方法。SANE 通过学习与任务无关的神经网络表示,克服了先前的局限性,这些表示可扩展到具有不同架构的较大模型,并展现出超越单一任务的能力。我们的方法将超表示(hyper-representations)的概念扩展到对神经网络权重的子集进行顺序处理,从而允许将较大的神经网络作为一组标记嵌入到学习到的表示空间中。SANE 能够从逐层嵌入中揭示全局模型信息,并顺序生成未见过的神经网络模型,这是先前的超表示学习方法无法实现的。如下图所示,广泛的实证评估表明,SANE(浅蓝色)在多个权重表示学习基准上达到或超越了最先进性能,特别是在新任务的初始化和较大的 ResNet 架构方面。

关键方法
- 顺序分解:将神经网络权重分解为更小、更易管理的 token 序列。
- 自监督预训练:采用自监督方法,在模型子序列上针对多种任务和架构对 SANE 模型进行预训练。
- 模型分析:通过其嵌入序列分析模型。
- 模型采样:通过对学习到的表示空间进行采样,生成新的神经网络模型。
结果
- 模型属性预测:SANE 嵌入在跨不同数据集和架构的模型属性(如测试准确率、epoch 和泛化差距)上展现出高预测性能。
- 生成能力:与从头训练相比,SANE 能够以显著更少的计算开销从头生成高性能神经网络模型或对其进行微调。
- 可扩展性:该方法可扩展至 ResNet-18 等大型模型,在长 token 序列中保留有意义的信息。
代码结构
- data/:用于数据预处理和加载的脚本。
- experiments/:在 CIFAR100-ResNet18 模型库上预训练 SANE、预测属性和采样模型的示例实验
- src/:包含 SANE 包,用于预处理模型检查点数据集、预训练 SANE 以及执行判别式和生成式下游任务。
运行实验
我们提供了运行示例实验的代码,并展示了如何使用我们的代码。
下载模型动物园数据集
我们在 modelzoos.cc 上提供了多个模型动物园。这些动物园中的任何一个都可以在我们的流水线中使用,只需进行少量调整。
要开始一个小规模实验,请导航至 ./data/ 并运行
bash download_cifar10_cnn_sample.sh
这将下载并解压一个包含在 CIFAR-10 上训练的 CNN 模型的小型模型动物园示例。 在大型模型动物园上进行训练需要预处理以提高训练效率。我们提供了用于预处理训练样本的代码。要编译这些数据集,请运行
python3 preprocess_dataset_cnn_cifar10_sample.py
位于 ./data/ 中。在同一目录中,我们还提供了用于下载和预处理其他动物园的脚本。
预处理后的数据集没有特定的依赖要求,除了常规的 numpy 和 pytorch。
请注意,这并非论文中使用的确切模型,因此会产生不同的结果。完整的动物园可以从 modelzoos.cc 下载,并以与动物园示例相同的方式使用。
预训练 SANE
在示例动物园上预训练 SANE 的代码包含在 experiments/pretrain_sane_cifar100_resnet18.py 中。该代码依赖 ray.tune 来管理资源,但目前仅运行单个配置。
要更改任何配置,请将值替换为 tune.grid_search([value_1, ..., value_n])。要运行实验,请运行
python3 pretrain_sane_cifar100_resnet18.py
位于 experiments/
使用 SANE 嵌入预测属性
SANE 嵌入保留了模型的顺序分解。与全局模型嵌入相比,这使得能够对模型进行更细粒度的分析。下图展示了 SANE 嵌入(右侧)与 WeightWatcher 库中使用的特征(左侧)之间的比较,后者基于权重矩阵的特征分解。两者在 ResNet 模型中均显示出相似的层属性趋势,但 SANE 似乎能捕捉到中间层的额外信号。

我们在 experiments/property_prediction_cifar100_resnet18.py 中提供了使用 SANE 嵌入预测此类模型属性的代码。它假设已下载数据集并预训练了如上所述的 SANE。在 property_prediction_cifar100_resnet18.py 中,设置预训练 SANE 模型和 epoch 的路径。然后在
python3 property_prediction_cifar100_resnet18.py
中运行 experiments。这将计算来自 SANE 嵌入和权重统计基线的属性预测结果,并将它们保存在一个 json 中。
生成模型
生成模型可以为新任务和新架构提供初始化,从而相比随机初始化具有优势,见下图。

在 experiments 中,还有用于生成和评估模型的代码。
cnn-cifar10_exploration.ipynb 是一个快速入门笔记本,用于探索数据集、SANE 模型、训练循环,以及模型的编码和解码。
我们进一步提供了针对 cnns 样本数据集和更大的 resnet 数据集的实验代码。对于后者,sample_finetune_cifar100_resnet18.py 包含一个模型采样的示例。如上所述,设置预训练 SANE 模型的路径和 epoch。然后,在 experiments 中运行
python3 sample_finetune_cifar100_resnet18.py
生成模型需要一个经过预处理的 CIFAR100 数据集,可以通过运行以下命令生成
python3 prepare_cifar100_dataset.py
小型 cnn 样本实验代码的结构与此对应。
联系方式
如果您对本项目有任何疑问,或希望获取数据和/或预训练模型的访问权限,请随时与我们联系。请联系 konstantin.schuerholt@unisg.ch。
引用
如果您在研究中使用此代码,请引用我们的论文:
@inproceedings{schuerholt2024sane,
title={Towards Scalable and Versatile Weight Space Learning},
author={Konstantin Sch{"u}rholt and Michael W. Mahoney and Damian Borth},
booktitle={Proceedings of the 41st International Conference on Machine Learning (ICML)},
year={2024},
organization={PMLR}
}
许可证
本项目采用 MIT 许可证。