DeepGraphLearning / DeepGraphLearning/PerturbDiff
Unable to reproduce Table 4 Replogle results using released finetuned_replogle.ckpt
Nessuno ha ancora preso questa issue.
- Lingua principale
- Python
- Stelle
- 63
- Fork
- 10
- Metriche di merge delle PR
- Nessuna PR unita negli ultimi 30g
Descrizione
Thanks for releasing the code and checkpoints!
I'm trying to reproduce the Replogle row of Table 4 (PerturbDiff Finetuned) using the released preprocessed data and checkpoint, running inference only (no retraining). I'm getting results very different from the paper and would appreciate some guidance.
Setup
- Checkpoint:
finetuned_replogle.ckptfromkatarinayuan/PerturbDiff_release_ckpt - Data:
katarinayuan/PerturbDiff_data, Replogle processed data - Gene vocabulary:
merged_pbmc_tahoe_rep_cellxgene_genes_mapped.pkl(12,626 genes, 12626 model mode). I confirmed the Replogle shared-gene ratio is 5,760/12,626 = 45.6%, which matches Table 3. - Evaluation: Cell-Eval v0.6.6
- Git commit:
f4e27c155be5325418c4cb3182453d4022754e91(origin/main, 2026-04-07) - Seed / devices: seed 42 (repo default via
optimization.seed, not overridden); single GPU (trainer.devices=[0]) - Hardware: NVIDIA GB10 (aarch64), CUDA 13.0, PyTorch cu130, Python 3.10
** Command: **
python ./src/apps/run/rawdata_diffusion_sampling.py \
run_name=replogle_finetuned_full \
model_checkpoint_path=<path_to>/finetuned_replogle.ckpt \
trainer.use_distributed_sampler=false \
trainer.devices=[0] \
data.normalize_counts=10 \
path=trixie_path \
cov_encoding=trixie_onehot \
cov_encoding.batch_encoding=onehot \
cov_encoding.celltype_encoding=llm \
cov_encoding.replogle_gene_encoding=genept \
model.p_drop_control=0 \
data.keep_control_cell=false \
sampling.use_ddim=true \
sampling.num_sampled_batches=null \
data=replogle_finetune \
data.sample_replogle_only=true \
data.selected_gene_file=<path_to>/merged_pbmc_tahoe_rep_cellxgene_genes_mapped.pkl \
data.pad_length=12626 \
model.hidden_num=[12626,512] \
model.input_dim=12626 \
data.embed_key=X \
optimization.micro_batch_size=128 \
data.use_cell_set=32 \
optimization.optimizer.lr=0.002
** Note on checkpoint loading: **
the checkpoint's baked-in hyper_parameters reference the original training cluster's absolute paths (e.g. /projects/AI4D/core-132/...), which don't resolve on a different machine. I had to route checkpoint loading through the repo's own load_plmodel_checkpoint() (which already supports runtime path overrides) instead of a plain PlModel.load_from_checkpoint(...). Flagging in case it's relevant to reproducibility for others as well.
Result
Running sampling directly from the released checkpoint, I get Overall R² = -11.24, which is far from the reported 0.988.
Looking at the raw predictions, some of the predicted expression values are abnormally large (max ≈ 4000+), concentrated in a few perturbations (e.g. hepg2 DNAJA1), whereas the ground-truth values look normal (max ≈ 6.8).
Question
Running the released checkpoint as-is gives R² = -11.24 instead of the 0.988 reported in Table 4, so something in my setup clearly differs from yours. Could you help me figure out what's going wrong? In particular, has the released (refactored) code + checkpoint been verified to actually reproduce the Table 4 Replogle numbers on your side?
Any pointers on where to look first would be greatly appreciated. Thanks again for the great work!
Guida per i contributori
Nessuna guida per i contributori indicizzata per questo repository
Come iniziare
- Leggi tutta la issue e poi la guida ai contributi del progetto.
- Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
- Fai un fork del repository e lavora su un branch.
- Apri una pull request che faccia riferimento al numero della issue.
Direzione di ricerca
Inizia da src/apps/run/rawdata_diffusion_sampling.py e dalla load_plmodel_checkpoint() del repository, verificando come le sovrascritture dei percorsi a runtime interagiscono con gli hyper_parameters del checkpoint. Confronta il checkpoint rilasciato, il preprocessing di Replogle, il vocabolario dei geni e le opzioni di sampling con la configurazione riportata in Table 4. Il lavoro è concluso quando viene identificata la discrepanza nella configurazione oppure viene riprodotto l’Overall R² di 0.988 riportato.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Valutazione
- Stack tecnologico
- python, pytorch
- Ambito
- machine-learning
- Tipo di issue
- Bug
- Difficoltà
- 4/5
- Tempo stimato
- 3-5 giorni
- Stato di attività
- Attiva
- Chiarezza
- Da chiarire
- Idoneità per principianti
- 38/100