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