DeepGraphLearning / DeepGraphLearning/PerturbDiff

Unable to reproduce Table 4 Replogle results using released finetuned_replogle.ckpt

Abierto
#7 1 comentario 0 reacciones 0 asignados Ver en GitHub

Nadie ha tomado este issue todavía.

Lenguaje dominante
Python
Estrellas
63
Forks
10
Métricas de merge de PR
Sin PR fusionados en 30 d

Descripción

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.ckpt from katarinayuan/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!

Guía de contribución

No hay ninguna guía de contribución indexada para este repositorio

Primeros pasos

  1. Lee el issue completo y luego la guía de contribución del proyecto.
  2. Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
  3. Haz un fork del repositorio y trabaja en una rama.
  4. Abre un pull request que haga referencia al número del issue.

Línea de trabajo

Comienza con src/apps/run/rawdata_diffusion_sampling.py y la load_plmodel_checkpoint() del repositorio, comprobando cómo interactúan las anulaciones de rutas en tiempo de ejecución con los hyper_parameters del checkpoint. Compara el checkpoint publicado, el preprocesamiento de Replogle, el vocabulario de genes y las opciones de muestreo con la configuración de Table 4. El trabajo estará terminado cuando se identifique la discrepancia en la configuración o se reproduzca el Overall R² de 0.988 comunicado.

Escrito por el modelo de indexación a partir del texto del issue.

Evaluación

Stack tecnológico
python, pytorch
Área
machine-learning
Tipo de issue
Error
Dificultad
4/5
Tiempo estimado
3-5 días
Estado de actividad
Activo
Claridad
Necesita aclaración
Aptitud para principiantes
38/100

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.