feat(log): add TensorBoard support and enhance WandB API key checking
This commit is contained in:
parent
d4ac45b34e
commit
795ae48520
2 changed files with 44 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Reference in a new issue