import os import json import time import logging from pathlib import Path from typing import Any, Dict, List, Optional, Tuple import torch from transformers import Trainer from .manifest import build_manifest_from_dir from .vocab import build_gene_vocab_from_manifest, build_meta_vocab_from_manifest logger = logging.getLogger("scRep_train") # ============================================================================= # Distributed helpers # ============================================================================= def dist_is_initialized() -> bool: return torch.distributed.is_available() and torch.distributed.is_initialized() def dist_rank_world() -> Tuple[int, int]: if dist_is_initialized(): return torch.distributed.get_rank(), torch.distributed.get_world_size() rank = int(os.environ.get("RANK", "0")) world = int(os.environ.get("WORLD_SIZE", "1")) return rank, world def dist_barrier() -> None: if dist_is_initialized(): torch.distributed.barrier() def is_world_process_zero() -> bool: rank, _ = dist_rank_world() return rank == 0 def wait_for_file(path: str, timeout: float = 300.0, poll: float = 0.5) -> None: """Wait until a file exists and is readable.""" p = Path(path) start = time.time() last_err = None while time.time() - start < timeout: if p.exists(): try: with open(p, "r") as f: _ = f.read(1) return except Exception as e: last_err = e time.sleep(poll) raise TimeoutError(f"Timed out waiting for file: {path}. Last error: {last_err}") # ============================================================================= # JSON helpers # ============================================================================= def load_json(path: str) -> Any: with open(path, "r") as f: return json.load(f) def save_json(obj: Any, path: str) -> None: os.makedirs(os.path.dirname(path), exist_ok=True) tmp_path = path + ".tmp" with open(tmp_path, "w") as f: json.dump(obj, f, ensure_ascii=False, indent=2) f.flush() os.fsync(f.fileno()) os.replace(tmp_path, path) # ============================================================================= # FSDP / Trainer helpers # ============================================================================= def parse_fsdp_list(fsdp_str: Optional[str]) -> Optional[List[str]]: s = (fsdp_str or "").strip() if not s: return None s = s.replace(",", " ") parts = [p.strip() for p in s.split() if p.strip()] return parts if parts else None def load_fsdp_config(path: Optional[str]) -> Optional[Dict[str, Any]]: if not path: return None if not os.path.exists(path): raise FileNotFoundError(f"fsdp_config not found: {path}") return load_json(path) def normalize_report_to(report_to: str): return [] if report_to == "none" else report_to class TrainerBatchWrapperCollator: """Wrap the collated batch as {"batch": ...} for Trainer.""" def __init__(self, base_collator): self.base = base_collator def __call__(self, examples): return {"batch": self.base(examples)} # ============================================================================= # Asset helpers # ============================================================================= def gene_vocab_to_map(gene_vocab: Dict[str, Any]) -> Dict[str, int]: gene_offset = int(gene_vocab.get("gene_offset", 2)) gene_names = list(gene_vocab["gene_names"]) return {str(name): i + gene_offset for i, name in enumerate(gene_names)} def _normalize_meta_keys(meta_keys: str) -> List[str]: return [k.strip() for k in meta_keys.split(",") if k.strip()] def _asset_metadata(args) -> Dict[str, Any]: return { "data_dir": os.path.abspath(args.data_dir), "manifest_json": os.path.abspath(args.manifest_json) if args.manifest_json else "", "gene_vocab_json": os.path.abspath(args.gene_vocab_json) if getattr(args, "gene_vocab_json", "") else "", "use_raw": bool(args.use_raw), "backed": args.backed, "meta_keys": _normalize_meta_keys(args.meta_keys), } def prepare_scratch_assets(args, output_dir: str) -> Tuple[Any, Dict[str, Any], Dict[str, Any], Dict[str, int], Dict[str, Dict[str, int]], Dict[str, int]]: is_main = is_world_process_zero() manifest_path = os.path.join(output_dir, "manifest.json") gene_vocab_path = os.path.join(output_dir, "gene_vocab.json") meta_vocab_path = os.path.join(output_dir, "meta_vocab.json") metadata_path = os.path.join(output_dir, "asset_metadata.json") expected_metadata = _asset_metadata(args) cache_complete = all( os.path.exists(path) for path in (manifest_path, gene_vocab_path, meta_vocab_path, metadata_path) ) metadata_matches = False cache_reason = "missing cached asset files" if cache_complete: cached_metadata = load_json(metadata_path) metadata_matches = cached_metadata == expected_metadata if not metadata_matches: cache_reason = "cached asset metadata does not match current dataset/config" if cache_complete and metadata_matches: logger.info("loading cached scratch assets from %s", output_dir) manifest = load_json(manifest_path) gene_vocab = load_json(gene_vocab_path) meta_vocab = load_json(meta_vocab_path) gene_name_to_id = gene_vocab_to_map(gene_vocab) meta_maps = meta_vocab["meta_maps"] meta_vocab_sizes = meta_vocab["meta_vocab_sizes"] logger.info( "loaded cached scratch assets: manifest=%s datasets, genes=%s, meta_keys=%s", len(manifest), gene_vocab.get("gene_vocab_size", 0), list(meta_vocab_sizes.keys()), ) return manifest, gene_vocab, meta_vocab, gene_name_to_id, meta_maps, meta_vocab_sizes if is_main: logger.info("scratch assets cache miss in %s: %s", output_dir, cache_reason) os.makedirs(output_dir, exist_ok=True) if args.manifest_json and os.path.exists(args.manifest_json): logger.info("building manifest from explicit manifest_json=%s", args.manifest_json) manifest = load_json(args.manifest_json) save_json(manifest, manifest_path) else: logger.info("building manifest from data_dir=%s", args.data_dir) manifest_start = time.time() manifest = build_manifest_from_dir( args.data_dir, backed=args.backed, use_raw=bool(args.use_raw), ) if len(manifest) == 0: raise RuntimeError( "manifest is empty when building from `data_dir`. " "This usually means your h5ad files were filtered out.\n\n" "Common causes:\n" "1) `--use_raw` is enabled but your .h5ad has no `adata.raw`.\n" "2) Wrong `data_dir` / file pattern (no .h5ad found).\n" "3) Required obs/meta columns mismatch (meta_keys).\n\n" f"args.use_raw={args.use_raw}, data_dir={args.data_dir}" ) save_json(manifest, manifest_path) logger.info("manifest saved to %s in %.1fs", manifest_path, time.time() - manifest_start) meta_keys = _normalize_meta_keys(args.meta_keys) if getattr(args, "gene_vocab_json", ""): logger.info("loading fixed gene vocab from %s", args.gene_vocab_json) gene_vocab = load_json(args.gene_vocab_json) gene_names = list(gene_vocab["gene_names"]) gene_name_to_id = gene_vocab_to_map(gene_vocab) logger.info("fixed gene vocab loaded with %s genes", len(gene_names)) else: logger.info("building gene vocab from %s manifest entries", len(manifest)) gene_start = time.time() gene_name_to_id, gene_names = build_gene_vocab_from_manifest( manifest, use_raw=args.use_raw, backed=args.backed, gene_offset=2, sort_paths=True, ) logger.info("gene vocab built with %s genes in %.1fs", len(gene_names), time.time() - gene_start) logger.info("building meta vocab for keys=%s", meta_keys) meta_start = time.time() meta_maps, meta_vocab_sizes = build_meta_vocab_from_manifest( manifest, keys=meta_keys, backed=args.backed, sort_paths=True, ) logger.info("meta vocab built in %.1fs", time.time() - meta_start) if not getattr(args, "gene_vocab_json", ""): gene_vocab = { "gene_offset": 2, "gene_names": gene_names, "gene_vocab_size": max(gene_name_to_id.values()) + 1, } meta_vocab = { "meta_vocab_sizes": meta_vocab_sizes, "meta_maps": meta_maps, } save_json(gene_vocab, gene_vocab_path) save_json(meta_vocab, meta_vocab_path) save_json(expected_metadata, metadata_path) logger.info("scratch assets saved to %s", output_dir) else: logger.info("waiting for scratch assets to be prepared in %s", output_dir) wait_for_file(metadata_path) manifest = load_json(manifest_path) gene_vocab = load_json(gene_vocab_path) meta_vocab = load_json(meta_vocab_path) gene_name_to_id = gene_vocab_to_map(gene_vocab) meta_maps = meta_vocab["meta_maps"] meta_vocab_sizes = meta_vocab["meta_vocab_sizes"] logger.info("loaded scratch assets prepared by rank0 from %s", output_dir) return manifest, gene_vocab, meta_vocab, gene_name_to_id, meta_maps, meta_vocab_sizes def save_bundle( trainer: Trainer, model, save_dir: str, *, gene_vocab: Dict[str, Any], meta_vocab: Dict[str, Any], manifest: Any, ) -> None: """ Save a final bundle that is directly loadable by ``CellContrastivePredictModel.from_pretrained``. Under FSDP we explicitly ask Accelerate for a gathered state dict, then save it through the unwrapped model. All ranks participate in ``get_state_dict``; only world-process-zero writes files. """ os.makedirs(save_dir, exist_ok=True) if hasattr(trainer, "accelerator") and trainer.accelerator is not None: unwrapped_model = trainer.accelerator.unwrap_model(model) state_dict = trainer.accelerator.get_state_dict(trainer.model) if trainer.is_world_process_zero(): unwrapped_model.save_pretrained( save_dir, save_safetensors=bool(trainer.args.save_safetensors), state_dict=state_dict, gene_vocab=gene_vocab, meta_vocab=meta_vocab, manifest=manifest, extra_files={"training_args.json": trainer.args.to_dict()}, ) else: if not trainer.is_world_process_zero(): return model.save_pretrained( save_dir, save_safetensors=bool(trainer.args.save_safetensors), gene_vocab=gene_vocab, meta_vocab=meta_vocab, manifest=manifest, extra_files={"training_args.json": trainer.args.to_dict()}, )