from __future__ import annotations from typing import Any, Dict import os from pathlib import Path try: import wandb except Exception: # pragma: no cover wandb = None from .distributed import is_main_process def _read_text(path: Path) -> str | None: try: if path.exists(): text = path.read_text().strip() return text or None except Exception: pass return None def _write_text(path: Path, text: str) -> None: try: path.parent.mkdir(parents=True, exist_ok=True) path.write_text(text.strip() + "\n") except Exception: # W&B should never crash training just because the local ID file could not be written. pass def _env_int(name: str, default: int) -> int: try: return int(os.environ.get(name, default)) except Exception: return default def init_wandb(cfg: Dict[str, Any], run_name: str | None = None): """Initialize W&B on rank 0 only and never crash training if W&B is slow/down. Resume behavior: - If wandb.id or WANDB_RUN_ID is provided, use it. - Otherwise, if /wandb_run_id.txt exists, reuse that ID. - If a run ID is used, default resume mode is "allow" unless overridden by wandb.resume or WANDB_RESUME. Reliability behavior: - WANDB_INIT_TIMEOUT controls wandb.init timeout seconds (default 300). - WANDB_SERVICE_WAIT controls W&B service startup wait seconds (default 300). - If wandb.init raises any exception, training continues with W&B disabled. """ wb_cfg = cfg.get("wandb", {}) if not wb_cfg.get("enabled", False) or not is_main_process() or wandb is None: return None mode = os.environ.get("WANDB_MODE", wb_cfg.get("mode", "online")) entity = wb_cfg.get("entity") or os.environ.get("WANDB_ENTITY") or None project = os.environ.get("WANDB_PROJECT", wb_cfg.get("project", "SACFlow-FM")) out_dir = Path(cfg.get("output_dir", ".")) run_id_file = out_dir / "wandb_run_id.txt" run_id = ( wb_cfg.get("id") or os.environ.get("WANDB_RUN_ID") or _read_text(run_id_file) ) resume_mode = wb_cfg.get("resume") or os.environ.get("WANDB_RESUME") if run_id and resume_mode is None: resume_mode = "allow" init_timeout = _env_int("WANDB_INIT_TIMEOUT", int(wb_cfg.get("init_timeout", 300))) service_wait = _env_int("WANDB_SERVICE_WAIT", int(wb_cfg.get("service_wait", 300))) # Keep W&B files away from home/cache quotas if an output directory exists. os.environ.setdefault("WANDB_DIR", str(out_dir / "wandb")) os.environ.setdefault("WANDB_CACHE_DIR", str(out_dir / "wandb_cache")) os.environ.setdefault("WANDB_CONFIG_DIR", str(out_dir / "wandb_config")) try: settings = wandb.Settings(init_timeout=init_timeout, _service_wait=service_wait) run = wandb.init( project=project, entity=entity, mode=mode, name=run_name, id=run_id, resume=resume_mode, tags=wb_cfg.get("tags", []), config=cfg, settings=settings, ) except Exception as e: # pragma: no cover - depends on external W&B service/network print( f"[W&B WARNING] wandb.init failed ({type(e).__name__}: {e}). " "Continuing training with W&B disabled for this run. " "To avoid this, use WANDB_MODE=offline or increase WANDB_INIT_TIMEOUT.", flush=True, ) return None # Persist the generated or reused run ID so future --resume auto attaches to # the same W&B run from this output directory. if getattr(run, "id", None): _write_text(run_id_file, run.id) return run def wandb_log(run, data: Dict[str, Any], step: int | None = None): if run is not None: try: run.log(data, step=step) except Exception as e: print(f"[W&B WARNING] wandb.log failed: {e}", flush=True) def wandb_finish(run): if run is not None: try: run.finish() except Exception as e: print(f"[W&B WARNING] wandb.finish failed: {e}", flush=True)