DeepGraphLearning / DeepGraphLearning/PerturbDiff

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

Open
#7 1 comment 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
63
Forks
10
PR merge metrics
No merged PRs in 30d

Description

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!

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with src/apps/run/rawdata_diffusion_sampling.py and the repo's load_plmodel_checkpoint(), checking how runtime path overrides interact with the checkpoint hyper_parameters. Compare the released checkpoint, Replogle preprocessing, gene vocabulary, and sampling options with the reported Table 4 setup. Done means identifying the setup mismatch or reproducing the reported Overall R² of 0.988.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Active
Clarity
Needs clarification
Newbie friendliness
38/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.