Download scRep_pretrain/train_utils.py from jlu-wsj/scRep: direct link, hf CLI and curl.
- Browser
- Download file 11.5 kB
-
https://huggingface.co/jlu-wsj/scRep/resolve/main/scRep_pretrain/train_utils.py
- Command line
-
hf download hf://jlu-wsj/scRep/scRep_pretrain/train_utils.py
-
curl -L -o train_utils.py https://huggingface.co/jlu-wsj/scRep/resolve/main/scRep_pretrain/train_utils.py
11.5 kB
| 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()}, | |
| ) | |