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:
- The checkpoints are provided for non-commercial research and evaluation purposes only.
- You will not redistribute them, in whole or in part, to any third party.
- 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.
Supervised Concept SAE — crossmodal checkpoints
Crossmodal sparse autoencoders binding image features to a text-side concept space, trained across 46 experimental runs (datasets × text encoders × architecture ablations).
Contenu
449 checkpoints .pt (27.6 Go) + 1058 fichiers JSON de métriques (5.3 Mo) + le code.
weights/<run>/<tag>.pt # poids, cachés par cfg.tag()
artifacts/<run>/*.json # métriques, évaluations, vocabulaires
code/src/xmodal_sae/ # package d'inférence
code/scripts/ # scripts d'entraînement et d'évaluation
Runs principaux
| Run | Checkpoints | Contenu |
|---|---|---|
ds_v1 |
104 | balayage multi-datasets |
imnet200_sbert |
92 | ImageNet-200, encodeur texte SBERT |
datasets |
64 | grille datasets |
nwpu, rsicd, rs5m_sub |
35 / 18 / 3 | télédétection |
inat, inat_vlm, awa2, clevr, ben_vlm |
7–21 | domaines spécialisés |
textenc_{clip,bge,mpnet,sbert} |
1–2 | ablation d'encodeur texte |
abl_*, grid_v* |
1–2 | ablations et grilles |
*_scalemae, *_dinov3sat, *_rsclip |
3 | backbones de télédétection |
Convention de nommage
Le nom de fichier est la configuration (cfg.tag()), par exemple :
A_nmse1_auxk1_k32_x8_mean_dc0_da0_s0_ds_cub_tmexpand5.pt
│ │ │ │ │ │ │ │ │ │ └ mode d'expansion texte
│ │ │ │ │ │ │ │ │ └ dataset
│ │ │ │ │ │ │ │ └ seed
│ │ │ │ │ │ └───┴ flags de décorrélation
│ │ │ │ │ └ pooling
│ │ │ │ └ facteur d'expansion du dictionnaire
│ │ │ └ sparsité TopK
│ │ └ perte auxiliaire k
│ └ perte NMSE
└ architecture
archc_bce_<dataset>_k32_x8_lx1.0_dc0_da0_cb1.0_e150_s0.pt # variante Arch-C / BCE
Format
import torch
ckpt = torch.load("weights/ds_v1/<tag>.pt", map_location="cpu")
sae = ckpt["sae"] # state_dict du SAE
etxt = ckpt["etxt"] # state_dict de l'encodeur texte (None si gelé)
cfg = ckpt.get("cfg") # configuration sérialisée, selon le script
Les scripts d'entraînement écrivent selon les cas {"sae", "etxt"}, {"model", "cfg"}
ou un state_dict nu — vérifier les clés à l'ouverture.
cub/ — embeddings, SBERT VLG-CBM et reconstructions
cub/embeddings/ # cache CUB complet du run ds_v1 (A_img, A_txt, A_txt5,
# A_txtnp, cat_embeds, Y_vlgbce, image_ids, scaler.pkl, JSON)
cub/vlgbce_sbert/{base,rich,tro}.npy # concepts VLG-CBM encodés SBERT
cub/vlgbce_sbert/{base,rich,tro}_terms.json
cub/reconstruction/<tag>/img_hat.npy # decode_img(z_image)
cub/reconstruction/<tag>/txt_hat.npy # decode_txt(z_image) — chemin cross-modal
cub/reconstruction/<tag>/txt_hat_cyclic.npy # decode_txt(TopK(E_t(A_txt))) — chemin cyclique
cub/reconstruction/<tag>/codes_topk_{idx,val}.npy
cub/reconstruction/<tag>/metrics.json
cub/captions/ # captions CUB générées (4 variantes x 3 versions)
code/concept_captions/ # pipeline de génération des captions
Concepts VLG-CBM encodés par SBERT
| Variante | Concepts | Dim | Source |
|---|---|---|---|
base |
471 | 384 | vlgbce_concepts.json |
rich |
657 | 384 | vlgbce_concepts_rich.json |
tro |
393 | 384 | vlgbce_concepts_tro.json |
Encodeur sentence-transformers/all-MiniLM-L6-v2, embeddings L2-normalisés — le
même encodeur que le côté texte des SAE, donc directement comparables à A_txt.
Lire les métriques de reconstruction
L'entrée du SAE est A_img standardisé, pas A_img brut : le pipeline ajuste un
StandardScaler sur le split train (rng = default_rng(0), VAL_FRAC = 0.20) et
c'est A_img_std qui alimente l'encodeur. Le scaler.pkl du cache permet de
rejouer exactement la transformation ; reconstruire depuis l'array brut donne des
FVU aberrants.
Deux chemins texte, à ne pas confondre :
- cyclique —
decode_txt(TopK(E_t(A_txt))). C'est le chemin caption réellement appris. FVU val de 0.18 à 1.3 selon la config, cosinus jusqu'à 0.97. - cross-modal —
decode_txt(z_image). Entraîné uniquement silam_x > 0. Sur les tagslx0.0il n'a jamais reçu de gradient : le FVU y monte à ~2700, contre ~26 pour le même modèle enlx1.0. Un FVU ≫ 1 y signale un chemin non entraîné, pas une erreur de calcul.
Exemples (val, seed 0) :
| Checkpoint | image FVU | caption cyclique FVU / cos | caption cross FVU |
|---|---|---|---|
C_nmse1_auxk1_k32_x8_mean_..._tmexpand5 |
0.476 | 0.179 / 0.972 | 4.5 |
archc_bce_cub_k32_x8_lx0.0_... |
0.246 | 0.191 / 0.970 | 2731.3 |
archc_bce_cub_k32_x8_lx1.0_... |
0.990 | 0.547 / 0.909 | 26.4 |
geom_cub_sub1_pair0_prof0_sym0_x64_s0 |
0.261 | 0.310 / 0.951 | 1538.4 |
Les checkpoints *_wh_* (whitening ZCA côté image) sont exclus des reconstructions :
la matrice de whitening n'est pas sauvegardée, l'entrée n'est pas reproductible.
Reproduction : code/scripts/cub_xmodal_recon.py (+ run_cub.slurm).
Reproduction
Le code attend un répertoire de sortie XMODAL_OUT contenant weights/, artifacts/,
figures/, cache/. Pointer XMODAL_OUT sur la racine de ce repo permet de réutiliser
les poids sans réentraîner (ils sont résolus par cfg.tag()).
Citation
TODO — référence de l'article.