MuPaD-HE2mIF

Predict a 16-channel Orion multiplexed immunofluorescence (mIF) panel from a single H&E tile, by finetuning MuPaD-512 as a conditional mean predictor.

Code: github.com/jinxixiang/MUPAD5-HE2mIF/

Why a mean predictor

The evaluation metric is a per-channel Pearson correlation against ground truth, and that is maximized by the conditional mean E[mIF | H&E] — not by a sharp sample from p(mIF | H&E). The model is trained with flow matching but read out at t = 1, where the input is pure noise and the loss-optimal output is the conditional mean. Inference is a single forward pass, not a multi-step sampler.

z_he  = LatentNorm_he( HE_VAE(he) )              # pretrained H&E latent, frozen
z_mif = LatentNorm_mif( mIF_VAE(mif) )           # stage-1 mIF latent, frozen
z_t   = (1-t)·z_mif + t·ε
x̂     = SiT( [z_t ; z_he], t, MUSK(he) )         # x-prediction

inference:  x̂ = SiT( [ε ; z_he], t=1, MUSK(he) )  → decode

H&E enters as input, not context: z_he is concatenated channel-wise at the patch embedder, and the new condition slots are initialized by copying the pretrained patch-embed weights, so z_he lives in exactly the latent space the trunk was pretrained on. MUSK cross-attention is a secondary semantic channel.

Contents

stage1_vae/
  vae/config.json
  vae/diffusion_pytorch_model.safetensors   16-channel mIF VAE  (335 MB)
  latent_stats.json                         mIF latent mean/var, fitted at 256px
  latent_stats_512.json                     same, fitted at 512px  <- used at inference
  mif_input_scale.json                      raw-intensity scaling used by the VAE
  config.yaml                               stage-1 training config
stage2_denoising/
  diffusion_pytorch_model.safetensors       SiT trunk, 1.39B params, fp32 (5.56 GB)
  config.yaml                               stage-2 training config

Latent statistics are resolution-specific. latent_stats.json was fitted at the VAE's own 256px; applying it at 512px leaves normalized latents at std 0.68–0.79 instead of 1.0, which unbalances the interpolant. Use latent_stats_512.json for 512px inference — the code selects it automatically.

Requires

This repository holds only the HE2mIF-specific weights. You also need xiangjx/MuPaD-512 for the SiT trunk architecture, the H&E VAE, and the MUSK encoder.

Usage

git clone https://github.com/jinxixiang/MUPAD && cd MUPAD/5-HE2mIF
pip install -r ../requirements.txt

huggingface-cli download xiangjx/MuPaD-512     --local-dir pretrained/MuPaD-512
huggingface-cli download xiangjx/MuPaD-HE2mIF  --local-dir exps
export MUPAD_PRETRAINED_ROOT=$PWD/pretrained/MuPaD-512
export ORION_DATA_ROOT=/path/to/ORIONCRC_dataset_tile_20x
python generate.py \
    --config configs/stage2_denoising.yaml \
    --denoise-ckpt exps/stage2_denoising/diffusion_pytorch_model.safetensors \
    --vae-ckpt exps/stage1_vae \
    --val-csv $ORION_DATA_ROOT/val_dataframe.csv \
    --out-dir predictions/run \
    --guidance-scale 1.25

python eval/evaluate.py --pred-dir predictions/run \
    --csv $ORION_DATA_ROOT/val_dataframe.csv --out results.csv

Guidance uses zeroed conditioning as the null branch (there is no conditioning dropout in training). --guidance-scale 1.25 measurably beats the unguided default.

Markers

16 of Orion's 17 channels. PD-1 is excluded: its positivity column is 1.0000 on every training tile — a broken threshold, not biology — so it carries no usable supervision.

Hoechst, CD31, CD45, CD68, CD4, FOXP3, CD8a, CD45RO, CD20, PD-L1, CD3e, CD163, E-cadherin, Ki67, Pan-CK, SMA

Data

Trained on ORION-CRC tiles at 20× (512 × 512 px), Zenodo record 15340874.

License

MIT, following the REPA/SiT lineage the trunk derives from.

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

Model tree for xiangjx/MuPaD-HE2mIF

Finetuned
(1)
this model

Collection including xiangjx/MuPaD-HE2mIF