Request access to these checkpoints

Access requests are reviewed manually by the authors. Expect a few days for processing.

By requesting access, you agree to the following terms:

  1. The checkpoints are provided for non-commercial research and evaluation purposes only.
  2. You will not redistribute them, in whole or in part, to any third party.
  3. You will cite the associated work in any publication that uses them.

Log in or Sign Up to review the conditions and access this model content.

SAE–CBM Unification — SAE and backbone checkpoints

Sparse autoencoders trained on frozen backbone features, plus the fine-tuned backbones they were trained on, for the unification of sparse dictionaries and concept bottleneck models.

Contenu

231 fichiers .pt, 38.3 Go.

Axe Valeurs
Datasets cub, awa2, paco, skincon, imagenet
Backbones resnet50 (layer4), vit_b_16 (penultimate), dinov2_vitb14_reg (penultimate)
Seeds 0, 42, 2026
Dictionnaire TopK, K ∈ {1536, 4096}

Arborescence :

results/metrics_curve_cov2/<dataset>/seed_<s>/
    backbone_<dataset>_<backbone>.pt                # backbone fine-tuné
    metrics_curve_<DS>_<backbone>_s<s>_<ts>.csv     # courbe de métriques
    sae_artifacts_<DS>_topk_K<K>_<norm>_<layer>_<ts>/
        sae_<layer>_K<K>_topk_<norm>_<ts>_<lambda>.pt
results/metrics_curve_<dataset>/seed_<s>/           # série antérieure (cov1)
code/src/cbm_pipeline/                              # code d'inférence

Format des checkpoints SAE

torch.load(path, map_location="cpu") renvoie un dict :

Clé Contenu
sae_state_dict dictionary._weights, dictionary.multiplier, encoder.final_block.0.{weight,bias}
feat_to_concept_state_dict tête features → concepts
concept_clf_state_dict classifieur concept (fc.weight, fc.bias)
W_dec, W_cbm, Z décodeur, matrice CBM, codes
embeddings activations mises en cache : H_log, H_prob, splits train/test
métadonnées layer, method (topk), normal, lambda, split

Les .pt embarquent des activations en cache (embeddings) en plus des poids — c'est ce qui explique leur taille. Un export poids-seuls serait nettement plus léger.

Chargement

import torch
ckpt = torch.load("results/.../sae_penultimate_K1536_topk_normal_<ts>_<lambda>.pt", map_location="cpu")
sae_weights = ckpt["sae_state_dict"]
W_dec = ckpt["W_dec"]

Le code d'inférence est sous code/src/cbm_pipeline/.

cub/ — features CUB et reconstructions (self-contained)

De quoi vérifier les SAE CUB sans rien recalculer : les features d'entrée, la reconstruction, et les codes.

cub/features/<backbone>/{train,val,test}.npy         # [N, D] fp32, entrée du SAE
cub/features/<backbone>/{train,val,test}_labels.npy  # [N] int64, classe
cub/features/<backbone>/meta.json
cub/reconstruction/<sae_tag>/{train,val,test}_xhat.npy   # [N, D] fp32
cub/reconstruction/<sae_tag>/{train,val,test}_l0.npy     # [N] int32, L0 par échantillon
cub/reconstruction/<sae_tag>/metrics.json
cub/summary.json

Splits CUB-200-2011 : train 4796 / val 1198 / test 5794 (val_ratio=0.2, 200 classes).

Fidélité de reconstruction (FVU, seed 0)

Backbone Couche D K top-k train val test
resnet50 layer4 2048 4096 20 0.179 0.178 0.253
vit_b_16 penultimate 768 1536 8 0.267 0.271 0.371
dinov2_vitb14_reg penultimate 768 1536 8 0.178 0.185 0.232

FVU = ||x - x̂||² / ||x - mean(x)||² (0 = parfait, 1 = prédicteur moyen). Le L0 mesuré vaut exactement le top_k visé sur tous les échantillons.

vit_b_16 et dinov2_vitb14_reg produisent tous deux des dossiers sae_artifacts_CUB_topk_K1536_normal_penultimate_<ts> en dimension 768. Seul le timestamp les distingue, apparié via metrics_curve_CUB_<backbone>_s0_<ts>.csv : 124812 = ViT-B/16, 124813 = DINOv2.

Reproduction : code/scripts/cub_cbm_features_recon.py (+ run_cub.slurm).

Features d'entrée ImageNet

Les features backbone ImageNet fp32 sont publiées séparément — voir le repo imagenet-backbone-features-fp32 (même gating).

Citation

TODO — référence de l'article.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support