scRep / scRep_inference.py
jlu-wsj's picture
Add files using upload-large-folder tool
8a08c70 verified
Raw
History Blame Contribute Delete
11.1 kB
"""Standalone inference helpers for the public scRep release.
This is the small, public equivalent of scDINO's experiment-side inference
adapter. It deliberately supports a checkpoint containing only weights by
inferring the architecture and taking vocab assets from ``assets/``.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Sequence
import anndata as ad
import numpy as np
import torch
from scRep_pretrain.model import CellContrastivePredictModel, ContrastiveSplitCellCollator, FoundationConfig
from scRep_pretrain.vocab import VocabSpec, gene_vocab_to_map
ROOT = Path(__file__).resolve().parent
LABEL_CANDIDATES = ["cell_type", "author_cell_type", "cell_type_ontology_term_id", "time_label"]
BATCH_CANDIDATES = ["extended_batch_id", "batch_id", "batch", "sample_id", "sample", "donor_id"]
@dataclass
class EvalExample:
gene_ids: list[int]
expr: list[float]
label: str | None = None
source_file: str = ""
cell_idx: int = -1
batch_id: str | None = None
def safe_name(value: Any) -> str:
result = "".join(c if c.isalnum() or c in "._-" else "-" for c in str(value)).strip("-")
return result or "item"
def infer_model_cache_name(model_dir: str | Path) -> str:
path = Path(model_dir).expanduser().resolve()
return safe_name(f"{path.parent.name}_{path.name}" if path.name.startswith("checkpoint-") else path.name)
def choose_label_key(explicit_key: str, available: Sequence[str]) -> str | None:
if explicit_key:
return explicit_key if explicit_key in set(available) else None
return next((key for key in LABEL_CANDIDATES if key in set(available)), None)
def choose_batch_key(available: Sequence[str]) -> str | None:
return next((key for key in BATCH_CANDIDATES if key in set(available)), None)
def list_data_files(data_dir: str | Path, data_files: Sequence[str | Path] = ()) -> list[Path]:
root = Path(data_dir).expanduser().resolve()
if data_files:
paths = [Path(p).expanduser() for p in data_files]
return [(p if p.is_absolute() else root / p).resolve() for p in paths]
paths = sorted(root.glob("*.h5ad")) if root.is_dir() else [root]
if not paths or not all(p.is_file() for p in paths):
raise FileNotFoundError(f"No .h5ad files found in: {root}")
return paths
def _extract_row(matrix: Any, row: int) -> tuple[np.ndarray, np.ndarray]:
value = matrix[row]
if hasattr(value, "indices"):
return np.asarray(value.indices), np.asarray(value.data, dtype=np.float32)
dense = np.asarray(value).reshape(-1)
idx = np.flatnonzero(dense)
return idx, dense[idx].astype(np.float32)
def _process_expr(values: np.ndarray, use_raw: bool, cp_target_sum: float) -> np.ndarray:
values = np.asarray(values, dtype=np.float32)
if not use_raw or not values.size:
return values
total = float(values.sum())
return np.log1p(values * (float(cp_target_sum) / total)) if total > 0 else np.empty(0, dtype=np.float32)
def load_examples(data_dir: str | Path, gene_name_to_id: dict[str, int], *, max_cells_per_file: int = 0,
max_total_cells: int = 0, seed: int = 42, backed: str = "r", use_raw: bool = False,
cp_target_sum: float = 1e4, label_key: str = "") -> tuple[list[EvalExample], str | None]:
rng, result, chosen_label = np.random.default_rng(seed), [], None
for path in list_data_files(data_dir):
adata = ad.read_h5ad(path, backed=backed)
try:
matrix, names = (adata.raw.X, adata.raw.var_names) if use_raw else (adata.X, adata.var_names)
if matrix is None:
raise ValueError(f"Missing expression matrix: {path}")
if use_raw and adata.raw is None:
raise ValueError(f"--use_raw requested but raw is missing: {path}")
if chosen_label is None:
chosen_label = choose_label_key(label_key, list(adata.obs.columns))
local_to_global = np.fromiter((gene_name_to_id.get(str(g), 0) for g in names), dtype=np.int64, count=len(names))
indices = np.arange(adata.n_obs); rng.shuffle(indices)
if max_cells_per_file > 0:
indices = indices[:max_cells_per_file]
for row in indices:
idx, values = _extract_row(matrix, int(row)); values = _process_expr(values, use_raw, cp_target_sum)
gids = local_to_global[idx]; valid = gids > 0
if not np.any(valid) or not values.size:
continue
label = str(adata.obs.iloc[int(row)][chosen_label]) if chosen_label else None
result.append(EvalExample(gids[valid].astype(int).tolist(), values[valid].astype(float).tolist(), label,
path.name, int(row), None))
if max_total_cells > 0 and len(result) >= max_total_cells:
return result, chosen_label
finally:
if getattr(adata, "file", None) is not None:
adata.file.close()
return result, chosen_label
def _config_from_state(state: dict[str, torch.Tensor], vocab: VocabSpec) -> FoundationConfig:
emb = state["student_backbone.tokenizer.gene_emb.weight"]
layers = len({key.split(".")[2] for key in state if key.startswith("student_backbone.encoder.") and key.endswith("norm1.scale")})
return FoundationConfig(d_model=int(emb.shape[1]), n_heads=12, encoder_layers=layers,
gene_vocab_size=int(emb.shape[0]), bin_vocab_size=int(state["student_backbone.tokenizer.bin_emb.weight"].shape[0]),
gene_pad_id=vocab.gene_pad_id, bin_pad_id=vocab.bin_pad_id, gene_cls_id=vocab.gene_cls_id,
dino_out_dim=int(state["student_head.prototypes.weight"].shape[0]),
dino_hidden_dim=int(state["student_head.mlp.0.weight"].shape[0]),
dino_bottleneck_dim=int(state["student_head.mlp.4.weight"].shape[0]))
def load_scRep_bundle(model_dir_or_args: str | Path | Any, args: Any | None = None):
if args is None:
args, model_dir = model_dir_or_args, getattr(model_dir_or_args, "model_dir")
else:
model_dir = model_dir_or_args
weights_dir = Path(model_dir).expanduser().resolve()
from safetensors.torch import load_file
state = load_file(str(weights_dir / "model.safetensors"), device="cpu")
asset_arg = Path(str(getattr(args, "asset_dir", "") or weights_dir / "assets")).expanduser()
candidates = [weights_dir, weights_dir / "assets", weights_dir / "tokenizer", asset_arg]
def asset(name: str) -> Path:
found = next((p / name for p in candidates if (p / name).is_file()), None)
if found is None: raise FileNotFoundError(f"Missing {name}; pass --asset_dir")
return found
gene_vocab = json.loads(asset("gene_vocab.json").read_text())
meta_path, manifest_path = next((p / "meta_vocab.json" for p in candidates if (p / "meta_vocab.json").is_file()), None), next((p / "manifest.json" for p in candidates if (p / "manifest.json").is_file()), None)
vocab = VocabSpec()
config_path = next((candidate for p in candidates for candidate in (p / "config.json", p / "model_config.json") if candidate.is_file()), None)
cfg = FoundationConfig(**json.loads(config_path.read_text())) if config_path else _config_from_state(state, vocab)
model = CellContrastivePredictModel(cfg); model.load_state_dict(state, strict=True)
device = str(getattr(args, "device", "cpu")); model.to(device).eval()
return model, gene_vocab, json.loads(meta_path.read_text()) if meta_path else None, json.loads(manifest_path.read_text()) if manifest_path else None, {"model_family": "model", "weights_dir": str(weights_dir), "config_path": str(config_path or "")}
def _select_input_genes(gene_ids: np.ndarray, expr: np.ndarray, max_input_genes: int):
if max_input_genes <= 0 or len(gene_ids) <= max_input_genes: return gene_ids, expr
order = np.argsort(-expr, kind="stable")[:max_input_genes]
return gene_ids[order], expr[order]
def _make_eval_collator(vocab: VocabSpec, gene_vocab_size: int, n_bins: int, max_input_genes: int, model_family: str = "model", expr_embedding_mode: str = "bin"):
limit = max_input_genes if max_input_genes > 0 else gene_vocab_size
return ContrastiveSplitCellCollator(vocab=vocab, gene_vocab_size=gene_vocab_size, n_bins=n_bins, max_input_genes=max_input_genes,
shuffle_genes=False, teacher_top_genes=limit, teacher_num_views=1, teacher_global_ratio=1.0,
student_global_ratio_min=1.0, student_global_ratio_max=1.0, student_global_min_genes=limit,
student_local_min_genes=limit, student_local_max_genes=limit, student_num_local_views=0,
student_global_dropout_prob=0.0, student_local_dropout_prob=0.0, student_global_bin_mask_prob=0.0,
expr_embedding_mode=expr_embedding_mode)
def build_eval_batch(examples: Sequence[EvalExample], *, vocab: VocabSpec, collator: Any, max_input_genes: int, use_bin: bool):
rows = []; width = 1
for ex in examples:
gids, expr = _select_input_genes(np.asarray(ex.gene_ids, dtype=np.int64), np.asarray(ex.expr, dtype=np.float32), max_input_genes)
order = np.argsort(-expr, kind="stable") if use_bin else np.argsort(gids, kind="stable")
gids, expr = gids[order], expr[order]
bins = collator._encode_expression_ids(expr) if use_bin else np.full(len(gids), vocab.bin_pad_id, dtype=np.int64)
rows.append((gids, bins)); width = max(width, len(gids) + 1)
gene = torch.full((len(rows), width), vocab.gene_pad_id, dtype=torch.long); bins = torch.full_like(gene, vocab.bin_pad_id); mask = torch.zeros_like(gene, dtype=torch.bool)
gene[:, 0] = vocab.gene_cls_id; mask[:, 0] = True
for i, (gids, vals) in enumerate(rows):
gene[i, 1:len(gids)+1] = torch.from_numpy(gids); bins[i, 1:len(gids)+1] = torch.from_numpy(vals); mask[i, 1:len(gids)+1] = True
return {"gene_ids": gene, "bin_ids": bins, "attention_mask": mask}
def encode_embeddings(model: Any, examples: Sequence[EvalExample], args: Any, *, use_bin: bool = True, model_family: str = "model") -> np.ndarray:
vocab = VocabSpec(); collator = _make_eval_collator(vocab, model.cfg.gene_vocab_size, int(args.n_bins), int(args.max_input_genes), model_family, model.cfg.expr_embedding_mode)
backbone = None if getattr(args, "backbone", "auto") == "auto" else args.backbone; pieces = []
with torch.inference_mode():
for start in range(0, len(examples), int(args.batch_size)):
batch = build_eval_batch(examples[start:start+int(args.batch_size)], vocab=vocab, collator=collator, max_input_genes=int(args.max_input_genes), use_bin=use_bin)
states = model._encode_tokens(batch["gene_ids"].to(args.device), batch["bin_ids"].to(args.device), batch["attention_mask"].to(args.device), use_bin_embeddings=use_bin, backbone=backbone)
pieces.append(torch.nn.functional.normalize(states[:, 0], dim=-1).cpu())
return torch.cat(pieces).numpy() if pieces else np.empty((0, model.cfg.d_model), dtype=np.float32)