| """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) |
|
|