You need to agree to share your contact information to access this model

This repository is publicly accessible, but you have to accept the conditions to access its files and content.

Log in or Sign Up to review the conditions and access this model content.

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