1
Fork 0

feat(log): add TensorBoard support and enhance WandB API key checking

This commit is contained in:
Tibo De Peuter 2026-04-08 18:58:33 +02:00
parent d4ac45b34e
commit 795ae48520
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
2 changed files with 44 additions and 0 deletions

View file

@ -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

View file

@ -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,