Download pipeline.py from Sentinal4D/PhenoSeq: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/Sentinal4D/PhenoSeq/resolve/main/pipeline.py
- Command line
-
hf download hf://Sentinal4D/PhenoSeq/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/Sentinal4D/PhenoSeq/resolve/main/pipeline.py
11.8 kB
| """ | |
| PhenoSeq inference pipeline. | |
| Self-contained entry point for generating scGPT RNA-seq embeddings from | |
| ViT-L microscopy imaging features using a pretrained diffusion model. | |
| Quick start | |
| ----------- | |
| import numpy as np | |
| from pipeline import PhenoSeqPipeline | |
| >>> | |
| # Load from a local directory (or a HuggingFace repo id) | |
| pipe = PhenoSeqPipeline.from_pretrained(".") | |
| >>> | |
| # img_features: (n_cells, n_imaging_cells, 5120) β raw ViT-L embeddings | |
| # img_norm is loaded automatically from img_norm.npz when present | |
| rna = pipe(img_features) # β np.ndarray (n_cells, 512) | |
| Input format | |
| ------------ | |
| img_features : np.ndarray or torch.Tensor | |
| Shape (B, N, 5120) where | |
| B = number of target RNA cells to generate | |
| N = number of imaging cells per well (default 16) | |
| Features should be raw (unnormalized) ViT-L embeddings when img_norm | |
| is available, or pre-normalized when img_norm is None. | |
| Output format | |
| ------------- | |
| np.ndarray of shape (B, 512) β scGPT embedding space predictions, | |
| un-normalized back to the original scGPT scale. | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import os | |
| from pathlib import Path | |
| from typing import Optional, Union | |
| import numpy as np | |
| import torch | |
| import yaml | |
| logger = logging.getLogger(__name__) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Pipeline | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class PhenoSeqPipeline: | |
| """ | |
| Wraps the trained PhenoSeq diffusion model for single-cell RNA prediction. | |
| Parameters | |
| ---------- | |
| diffusion : GaussianDiffusion | |
| Loaded diffusion model (weights applied, eval mode). | |
| img_norm : dict with 'mean' and 'std' arrays of shape (5120,), optional | |
| Imaging normalisation statistics computed from the training split. | |
| If None, input features are assumed to be pre-normalized. | |
| device : str or torch.device | |
| Device for inference ('cuda', 'cpu', etc.). | |
| ddim_steps : int | |
| Number of DDIM denoising steps (50 is a good default; 0 = full DDPM). | |
| """ | |
| def __init__( | |
| self, | |
| diffusion, | |
| img_norm: Optional[dict] = None, | |
| device: Union[str, torch.device] = "cpu", | |
| ddim_steps: int = 50, | |
| ): | |
| self.diffusion = diffusion | |
| self.img_norm = img_norm | |
| self.device = torch.device(device) | |
| self.ddim_steps = ddim_steps | |
| self.diffusion.to(self.device).eval() | |
| # ββ Factory βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def from_pretrained( | |
| cls, | |
| model_dir: Union[str, Path] = ".", | |
| checkpoint_name: str = "best_model.pt", | |
| img_norm_name: str = "img_norm.npz", | |
| device: Optional[Union[str, torch.device]] = None, | |
| ddim_steps: int = 50, | |
| use_ema: bool = True, | |
| ) -> "PhenoSeqPipeline": | |
| """ | |
| Load a pipeline from a local directory or HuggingFace Hub repo. | |
| Parameters | |
| ---------- | |
| model_dir : str or Path | |
| Local path or HuggingFace Hub repo id containing the checkpoint. | |
| checkpoint_name : str | |
| Filename of the PyTorch checkpoint inside model_dir. | |
| img_norm_name : str | |
| Filename of the imaging normalisation stats (.npz with 'mean', 'std'). | |
| If the file is not found, img_norm is set to None and a warning is logged. | |
| device : str, torch.device, or None | |
| Target device; auto-selects CUDA when available if None. | |
| ddim_steps : int | |
| DDIM sampling steps. | |
| use_ema : bool | |
| Prefer EMA weights when available in the checkpoint (recommended). | |
| """ | |
| model_dir = Path(model_dir) | |
| # ββ Resolve from HuggingFace Hub if path doesn't exist locally ββββ | |
| if not model_dir.exists(): | |
| try: | |
| from huggingface_hub import snapshot_download | |
| model_dir = Path(snapshot_download(str(model_dir))) | |
| logger.info(f"Downloaded from HuggingFace Hub β {model_dir}") | |
| except Exception as exc: | |
| raise FileNotFoundError( | |
| f"Directory '{model_dir}' not found locally and Hub download failed: {exc}" | |
| ) from exc | |
| if device is None: | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| device = torch.device(device) | |
| # ββ Load checkpoint βββββββββββββββββββββββββββββββββββββββββββββββ | |
| ckpt_path = model_dir / checkpoint_name | |
| if not ckpt_path.exists(): | |
| raise FileNotFoundError(f"Checkpoint not found: {ckpt_path}") | |
| logger.info(f"Loading checkpoint from {ckpt_path}") | |
| ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False) | |
| # ββ Build model from embedded config ββββββββββββββββββββββββββββββ | |
| cfg = ckpt.get("config", {}) | |
| mc = ckpt.get("model_cfg", cfg.get("model", {})) | |
| dc = cfg.get("diffusion", {}) | |
| from models.denoiser import Img2RNADenoiser | |
| from models.diffusion import GaussianDiffusion | |
| # RNA normalisation is stored in the checkpoint | |
| rna_norm_raw = ckpt.get("rna_norm") | |
| rna_norm = ( | |
| {"mean": np.array(rna_norm_raw["mean"]), "std": np.array(rna_norm_raw["std"])} | |
| if rna_norm_raw is not None else None | |
| ) | |
| denoiser = Img2RNADenoiser( | |
| img_dim = mc.get("img_dim", 5120), | |
| rna_dim = mc.get("rna_dim", 512), | |
| model_dim = mc.get("model_dim", 1024), | |
| num_heads = mc.get("num_heads", 8), | |
| num_layers= mc.get("num_layers", 6), | |
| time_dim = mc.get("time_dim", 256), | |
| ff_mult = mc.get("ff_mult", 4), | |
| dropout = 0.0, | |
| ) | |
| diffusion = GaussianDiffusion( | |
| denoiser = denoiser, | |
| num_steps = dc.get("num_steps", 1000), | |
| schedule = dc.get("schedule", "cosine"), | |
| beta_start = dc.get("beta_start", 1e-4), | |
| beta_end = dc.get("beta_end", 0.02), | |
| rna_norm = rna_norm, | |
| ) | |
| # ββ Load weights ββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| ema_state = ckpt.get("ema_state_dict") | |
| if use_ema and ema_state is not None: | |
| logger.info("Loading EMA weights") | |
| diffusion.denoiser.load_state_dict( | |
| {k.removeprefix("denoiser."): v for k, v in ema_state.items()}, | |
| strict=False, | |
| ) | |
| else: | |
| if use_ema and ema_state is None: | |
| logger.warning("EMA weights not found; using standard model weights") | |
| diffusion.load_state_dict(ckpt["model_state_dict"], strict=True) | |
| n_params = sum(p.numel() for p in denoiser.parameters()) | |
| logger.info(f"Model loaded ({n_params:,} parameters) on {device}") | |
| # ββ Imaging normalisation stats βββββββββββββββββββββββββββββββββββ | |
| img_norm = None | |
| norm_path = model_dir / img_norm_name | |
| if norm_path.exists(): | |
| data = np.load(norm_path) | |
| img_norm = {"mean": data["mean"], "std": data["std"]} | |
| logger.info(f"Loaded imaging normalisation stats from {norm_path}") | |
| else: | |
| logger.warning( | |
| f"{img_norm_name} not found in {model_dir}. " | |
| "Pass pre-normalized imaging features, or provide img_norm manually. " | |
| "See save_img_norm.py to generate this file from your training data." | |
| ) | |
| return cls(diffusion, img_norm=img_norm, device=device, ddim_steps=ddim_steps) # type: ignore[return-value] | |
| # ββ Inference βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def __call__( | |
| self, | |
| img_features: Union[np.ndarray, torch.Tensor], | |
| batch_size: int = 256, | |
| ddim_steps: Optional[int] = None, | |
| ) -> np.ndarray: | |
| """ | |
| Generate scGPT RNA-seq embeddings from imaging features. | |
| Parameters | |
| ---------- | |
| img_features : array-like, shape (B, N, 5120) or (N, 5120) | |
| ViT-L imaging embeddings. When img_norm is available these should | |
| be raw (unnormalized); otherwise provide pre-normalized features. | |
| If 2D (N, 5120), a batch dimension is added automatically. | |
| batch_size : int | |
| Number of cells processed per forward pass. | |
| ddim_steps : int, optional | |
| Override the pipeline's default DDIM steps for this call. | |
| Pass 0 to use full DDPM sampling (slower, ~1000 steps). | |
| Returns | |
| ------- | |
| np.ndarray of shape (B, 512) | |
| Predicted scGPT-space RNA embeddings in the original (denormalized) scale. | |
| """ | |
| steps = ddim_steps if ddim_steps is not None else self.ddim_steps | |
| # Coerce to numpy then torch | |
| if isinstance(img_features, torch.Tensor): | |
| img_np = img_features.cpu().float().numpy() | |
| else: | |
| img_np = np.asarray(img_features, dtype=np.float32) | |
| # Add batch dim if single sample | |
| if img_np.ndim == 2: | |
| img_np = img_np[np.newaxis] | |
| if img_np.ndim != 3 or img_np.shape[-1] != 5120: | |
| raise ValueError( | |
| f"Expected img_features shape (B, N, 5120), got {img_np.shape}" | |
| ) | |
| # Normalize imaging features if stats are available | |
| if self.img_norm is not None: | |
| mean = self.img_norm["mean"].astype(np.float32) # (5120,) | |
| std = self.img_norm["std"].astype(np.float32) # (5120,) | |
| img_np = (img_np - mean) / std | |
| # Batch inference | |
| all_preds: list[np.ndarray] = [] | |
| for start in range(0, len(img_np), batch_size): | |
| chunk = torch.from_numpy(img_np[start : start + batch_size]).to(self.device) | |
| if steps > 0: | |
| preds = self.diffusion.sample_ddim( | |
| img_features = chunk, | |
| num_inference_steps = steps, | |
| eta = 0.0, | |
| ) | |
| else: | |
| preds = self.diffusion.sample(chunk) | |
| all_preds.append(preds.cpu().float().numpy()) | |
| return np.concatenate(all_preds, axis=0) # (B, 512) | |
| # ββ Convenience βββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def rna_dim(self) -> int: | |
| return self.diffusion.denoiser.rna_dim | |
| def __repr__(self) -> str: | |
| d = self.diffusion.denoiser | |
| return ( | |
| f"PhenoSeqPipeline(" | |
| f"rna_dim={d.rna_dim}, " | |
| f"model_dim={d.model_dim}, num_layers={len(d.layers)}, " | |
| f"ddim_steps={self.ddim_steps}, device={self.device})" | |
| ) | |