From 795ae48520740679709c0027407cc7c1d473ce25 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Wed, 8 Apr 2026 18:58:33 +0200 Subject: [PATCH] feat(log): add TensorBoard support and enhance WandB API key checking --- src/experiment_logger/unified_logger.py | 23 +++++++++++++++++++++++ src/experiment_logger/wandb_utils.py | 21 +++++++++++++++++++++ 2 files changed, 44 insertions(+) diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index 15c8a19..06ca3a9 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -124,6 +124,16 @@ class UnifiedLogger: # Save config to disk self._save_config() + # Setup TensorBoard + self.writer = None + try: + from torch.utils.tensorboard import SummaryWriter + + self.writer = SummaryWriter(self.run_dir) + self.info("TensorBoard SummaryWriter initialized.") + except ImportError: + self.warning("tensorboard not installed. Skipping SummaryWriter.") + # Initialize WandB if requested if self.use_wandb: self._init_wandb(project_name, entity, save_code) @@ -206,6 +216,16 @@ class UnifiedLogger: except Exception as e: self.warning(f"WandB logging failed: {e}") + # Log to TensorBoard + if self.writer is not None: + for k, v in metrics.items(): + if isinstance(v, (int, float, np.floating, np.integer)): + self.writer.add_scalar(k, v, step) + elif hasattr(v, "item"): + self.writer.add_scalar(k, v.item(), step) + elif isinstance(v, (np.ndarray, jnp.ndarray)) and v.size == 1: + self.writer.add_scalar(k, v.item(), step) + # Buffer for disk storage self.metrics_buffer.append(metrics_with_metadata) @@ -331,6 +351,9 @@ class UnifiedLogger: # Flush remaining metrics self._flush_metrics() + if self.writer is not None: + self.writer.close() + self.info(f"Run complete. Results saved to: {self.run_dir.absolute()}") # Finish WandB run diff --git a/src/experiment_logger/wandb_utils.py b/src/experiment_logger/wandb_utils.py index 5fd8837..2c162fb 100644 --- a/src/experiment_logger/wandb_utils.py +++ b/src/experiment_logger/wandb_utils.py @@ -36,6 +36,27 @@ def init_wandb( """ try: import wandb + import os + import sys + + # Robust HPC checking: check for API key + has_key = os.environ.get("WANDB_API_KEY") is not None + if not has_key: + try: + # Check if logged in locally via settings/netrc + has_key = wandb.setup().settings.api_key is not None + except Exception: + pass + + is_interactive = sys.stdout.isatty() + + if not has_key and not is_interactive and os.environ.get("WANDB_MODE") != "offline": + logger.warning( + "WANDB_API_KEY not found and environment is non-interactive. Switching to offline mode." + ) + sync_path = f"runs/{name}" if name else "runs" + logger.warning(f"WandB is offline. Use 'wandb sync {sync_path}' to upload logs later.") + os.environ["WANDB_MODE"] = "offline" run = wandb.init( project=project,