File size: 11,144 Bytes
94e9257 8a08c70 94e9257 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 | """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)
|