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)