galilai-group/stable-pretraining
Reliable, minimal and scalable library for pretraining foundation and world models
解决的问题
训练用于“基础模型”(如 CLIP、DINO 或 SimCLR 等学习通用视觉特征的模型)的大规模神经网络极为繁琐:研究人员必须同时处理数据加载、增强、日志记录、评估和集群作业管理,而微小的错误可能导致长时间运行无声失效。本项目是一个封装 PyTorch Lightning 的框架,旨在简化整个工作流程,使其更灵活、更稳定——让研究人员专注于模型和损失函数,而框架则负责处理底层 plumbing 工作。
如何工作
该框架围绕四个组件构建,它们通过 Python 字典相互传递数据:DataModule(数据)、Module(模型 + 前向传播)、Callbacks(监控/评估钩子)以及 Trainer/Manager(编排)。核心思想是:你的前向函数只需返回一个张量字典(如 {'loss': ..., 'embedding': ...}),其余所有操作——日志记录、评估探针、检查点等——都可直接从该字典读取,无需修改训练循环。框架内置了多种自监督学习方法的“前向”配方(SimCLR、DINO、MAE、BYOL、CLIP 等),还提供了 OnlineProbe(在冻结特征上训练小型线性分类器以实时追踪准确率)和 OnlineKNN(零训练最近邻评估器)等回调。它还将数据增强移至 GPU(通过 kornia 实现),消除了 CPU 瓶颈,并提供了一个 Manager,支持 SLURM 集群功能,如作业重新排队/恢复、原子级检查点和可查询的运行注册表。一个实验性的 JAX/Flax-NNX 后端也遵循相同的设计。
适用人群
从事基础模型或自监督学习研究的研究人员和工程师——尤其是那些在 GPU 集群上训练大型视觉模型,并希望实现实时评估、稳健检查点和减少样板代码的人。它也适用于任何希望以最少代码快速原型化或基准测试 SSL 方法(SimCLR、DINOv2、MAE 等)的人。
主要亮点
- 字典驱动设计:前向函数返回状态字典,因此任何中间张量均可自动记录,回调无需修改训练循环即可附加。
- 30+ 内置配方:涵盖 SSL、监督学习和多模态预训练(SimCLR、DINO/DINOv2、MAE、BYOL、VICReg、Barlow Twins、LeJEPA、CLIP 等)。
- 实时评估回调:如
OnlineProbe和OnlineKNN,可在训练过程中监控表示质量。 - GPU 侧批量增强(通过 kornia 实现):在批量上向量化增强操作,对不同模型大小和精度均实现可测量的吞吐量提升。
- SLURM 级编排:
Manager处理抢占/重新排队、原子级检查点和可查询的运行注册表。 - 实验性 JAX/Flax-NNX 后端:与 torch 设计保持一致,并通过数值一致性回归测试验证。
相关
- 项目
- 项目
- 项目
- 项目
- 项目