scRep / scRep_pretrain /train_utils.py
jlu-wsj's picture
Add files using upload-large-folder tool
94e9257 verified
Raw History Blame Contribute Delete
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()},
)