This document describes the workflow for training and inference using the LogitProd framework for both whole slide image (WSI)-level MIL tasks and patch-level classification, including ABMIL-based model training/inference and downstream logit fusion.
The pipeline consists of three main steps:
- Step 1: Feature extraction using Trident
- Step 2: Model training and inference (WSI-level MIL, survival analysis, gene mutation, and patch-level classification)
- Step 3: LogitProd logit-fusion aggregation and analysis
# Create a new conda environment
conda create -n LogitProd python=3.10 -y
# Activate the environment
conda activate LogitProd
# Install PyTorch (adjust CUDA version as needed)
conda install pytorch torchvision pytorch-cuda=12.1 -c pytorch -c nvidia
# Install other dependencies
pip install -r requirements.txtBefore training/inference (Step 2) and logit-fusion aggregation (Step 3), you need to extract patch-level features from whole slide images (WSI) using Trident.
Please refer to the Trident GitHub repository for installation and feature extraction instructions.
The output should be patch-level features in the following structure:
<trident_processed>/
└── 20x_256px/
└── features_{model_name}/ # e.g., features_uni_v2/
└── {slide_id}.h5
After Step 1, ensure you have:
-
Patch-level features: Extracted features from Trident (as shown above)
-
Data splits: CSV files containing train/val/test splits
<splits_dir>/ └── splits_{split_idx}_k.csv # e.g., splits_0_k.csv, splits_1_k.csv, ...
In Step 2, you run task-specific training and inference scripts for four types of tasks.
All scripts are located under scripts:
WSI_classification/Gene_mutation/Survival_analysis/Patch_classification/
Each folder contains its own training and inference scripts with consistent GitHub-release style CLI arguments (paths parameterized, no hardcoded user paths).
- Scripts location:
scripts/WSI_classification/ - Typical usage:
cd scripts/WSI_classification
# Training
python train_abmil_WSI_classification.py --help
# Inference
python infer_abmil_WSI_classification.py --help- Scripts location:
scripts/Gene_mutation/ - Typical usage:
cd scripts/Gene_mutation
# Training
python train_abmil_Gene_mutation.py --help
# Inference
python infer_abmil_Gene_mutation.py --help- Scripts location:
scripts/Survival_analysis/ - Typical usage:
cd scripts/Survival_analysis
# Training
python train_abmil_Survival_analysis.py --help
# Inference
python infer_abmil_Survival_analysis.py --help- Scripts location:
scripts/Patch_classification/ - Typical usage:
cd scripts/Patch_classification
# Training + inference are implemented in a single script
python train_infer_Patch_classification.py --helpIn Step 3, you run the LogitProd scripts to aggregate logits/features from Step 2 via centralized logit fusion across tasks/models.
All LogitProd-related scripts live alongside the task scripts in:
scripts/WSI_classification/LogitProd_WSI_classification.pyscripts/Gene_mutation/LogitProd_Gene_mutation.pyscripts/Survival_analysis/LogitProd_Survival_analysis.pyscripts/Patch_classification/LogitProd_Patch_classification.py
cd scripts/WSI_classification
python LogitProd_WSI_classification.py --helpYou can choose the appropriate LogitProd script for:
- WSI-level classification: multi-model fusion of slide-level logits
- Gene mutation prediction: fusion of mutation logits from multiple expert models
- Survival analysis: fusion of survival-related logits/outputs from multiple expert models
- Patch-level classification: fusion of patch-level logits
Each script exposes task-specific CLI arguments (paths to logits / features from Step 2,
output directory for fusion results, etc.). Use --help to inspect the exact options.
