Coladan: A High-Performance, Uncertainty-guided Multimodal-Multiomic Whole Slide AI Model generates Genome-wide Spatial Gene Expression and Image-only Virtual Perturbation from Histopathology Images
git and environment
https://github.com/99wzj/Coladan/
inference demo
import h5py
import torch
from transformers import AutoConfig
import numpy as np
import pandas as pd
from demo_data import load_10X_demo_h5
from modeling_coladan import Coladan
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# load model
cfg = AutoConfig.from_pretrained("WangZj99/Coladan", trust_remote_code=True)
model = Coladan(cfg)
model = model.to(device)
from huggingface_hub import hf_hub_download
state_dict_path = hf_hub_download(
"WangZj99/Coladan",
filename="state_dict.pt",
)
model.load_state_dict(torch.load(state_dict_path))
#load data
demo_path = hf_hub_download(
"WangZj99/Coladan",
filename="10X_demo.h5",
)
loaded_data = load_10X_demo_h5(
demo_path,
load_expression=True,
)
patches = loaded_data['he_image']
coords = loaded_data['coordinates']
predict_genes = loaded_data['predict_gene'] # length 16000 for 16000 genes
# predict from HE , only support one slide per time
with torch.autocast(device_type='cuda', dtype=torch.bfloat16), torch.inference_mode():
patches = patches
coords = coords.to(device)
predict_genes = torch.tensor(loaded_data['predict_gene'].to_numpy(), dtype=torch.long)
predict_matrix = model.predict_gene_from_image(patches, coords, predict_genes)
#extra
#calculate perason and spearman between ori and predict
gene_order = predict_genes.cpu().numpy()
ori_matrix_df = loaded_data['expression_matrix'][gene_order]
ori_matrix = ori_matrix_df.to_numpy(dtype=np.float32) # shape: [N_cells, 16000]
predict_matrix = predict_matrix.to(dtype=torch.float32).cpu().numpy()
from coladan import calculate_pearson_and_spearman
mean_pearson,mean_spearman = calculate_pearson_and_spearman(ori = ori_matrix,predict = predict_matrix)
print(f"mean Pearson : {mean_pearson:.4f}")
print(f"mean Spearman : {mean_spearman:.4f}")
image-only perturb demo
from contextlib import nullcontext
import torch
from huggingface_hub import hf_hub_download
from transformers import AutoConfig
from modeling_coladan import Coladan
from demo_data import load_10X_demo_h5
HF_REPO_ID = "WangZj99/Coladan"
BASE_STATE_DICT_FILENAME = "state_dict.pt"
PERTURB_DECODER_FILENAME = "perturb_decoder.pt"
DEMO_FILENAME = "10X_demo.h5"
# Each gene can have its own perturbation direction.
# Supported values: "up", "down", or an explicit numeric value.
# Examples:
# {"AR": "up", "PTEN": "down"}
# {"AR": 200.0, "PTEN": -200.0}
PERTURBATIONS = {
"AR": "up",
"FOXA1": "up",
# "PTEN": "down",
}
# Default expression values injected for direction-based perturbations.
# Numeric values in PERTURBATIONS override these defaults.
UP_VALUE = 200.0
DOWN_VALUE = -200.0
def _to_long_tensor(x) -> torch.LongTensor:
"""Convert pandas/numpy/list/tensor gene IDs to a torch.LongTensor."""
if isinstance(x, torch.Tensor):
return x.long()
if hasattr(x, "to_numpy"):
x = x.to_numpy()
return torch.tensor(x, dtype=torch.long)
def main() -> None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Load the standard Coladan model first. This contains image encoder,
# slide encoder, projector and the standard inference gene decoder.
cfg = AutoConfig.from_pretrained(HF_REPO_ID, trust_remote_code=True)
model = Coladan(cfg).to(device)
state_dict_path = hf_hub_download(HF_REPO_ID, filename=BASE_STATE_DICT_FILENAME)
model.load_state_dict(torch.load(state_dict_path, map_location=device))
model.eval()
demo_path = hf_hub_download(
HF_REPO_ID,
filename=DEMO_FILENAME,
)
loaded_data = load_10X_demo_h5(
demo_path,
load_expression=False,
)
patches = loaded_data["he_image"]
coords = loaded_data["coordinates"]
if not isinstance(coords, torch.Tensor):
coords = torch.as_tensor(coords)
coords = coords.to(device)
# length 16000 for 16000 genes in the demo file
predict_genes = _to_long_tensor(loaded_data["predict_gene"])
# Image-only virtual perturbation from H&E; only one slide is supported at a time.
# predict_perturbation_from_image(...) automatically downloads perturb_decoder.pt
# and replaces model.gene_decoder with the perturbation-specific decoder.
autocast_context = (
torch.autocast(device_type="cuda", dtype=torch.bfloat16)
if device.type == "cuda"
else nullcontext()
)
with autocast_context, torch.inference_mode():
predict_perturb_matrix = model.predict_perturbation_from_image(
patches=patches,
coords=coords,
predict_genes=predict_genes,
perturbations=PERTURBATIONS,
up_value=UP_VALUE,
down_value=DOWN_VALUE,
perturb_decoder_repo_id=HF_REPO_ID,
perturb_decoder_filename=PERTURB_DECODER_FILENAME,
)
print(f"Perturb prediction matrix shape: {tuple(predict_perturb_matrix.shape)}")
# Optional: save the matrix locally.
# torch.save(predict_perturb_matrix, "predict_perturb_matrix.pt")
# np.save("predict_perturb_matrix.npy", predict_perturb_matrix.cpu().numpy())
if __name__ == "__main__":
main()
- Downloads last month
- -
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support