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 等)。
  • 即時評估回呼:如 OnlineProbeOnlineKNN,可在訓練過程中監控表示品質。
  • GPU 側批量增強(透過 kornia 實現):在批量上向量化增強操作,對不同模型大小與精度均實現可測量的吞吐量提升。
  • SLURM 級編排Manager 處理搶佔/重新排隊、原子級檢查點與可查詢的執行註冊表。
  • 實驗性 JAX/Flax-NNX 後端:與 torch 設計保持一致,並透過數值一致性回歸測試驗證。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案