From 8ec693f04c35aadd12a11a9f372b36814b690ed1 Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Tue, 31 Mar 2026 20:52:51 +0000 Subject: [PATCH] fix(logging): improve JSON serialization and add wandb directory to gitignore - Fix float32 serialization issue in unified logger metrics flushing - Add jax.numpy import for proper type handling - Add wandb/ directory to .gitignore to exclude temporary tracking files - Tested wandb integration: metrics, artifacts, and local backup working correctly --- .gitignore | 3 +++ src/experiment_logger/unified_logger.py | 12 +++++++++++- 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/.gitignore b/.gitignore index 8204ce9..6cc1e08 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,9 @@ artifacts/* runs/* +# Experiment tracking +wandb/ + # Python-generated files __pycache__/ *.py[oc] diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index 87e3006..1501bf5 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -13,6 +13,7 @@ from pathlib import Path from typing import Any, Dict, Optional import flax +import jax.numpy as jnp import numpy as np from experiment_logger.wandb_utils import finish_wandb, init_wandb @@ -153,7 +154,16 @@ class UnifiedLogger: metrics_file = self.metrics_dir / "metrics.jsonl" with open(metrics_file, "a") as f: for metric in self.metrics_buffer: - f.write(json.dumps(metric) + "\n") + # Convert numpy/jax types to native Python types for JSON serialization + serializable_metric = {} + for k, v in metric.items(): + if hasattr(v, "item"): # numpy/jax scalar + serializable_metric[k] = v.item() + elif isinstance(v, (np.ndarray, jnp.ndarray)): + serializable_metric[k] = v.tolist() + else: + serializable_metric[k] = v + f.write(json.dumps(serializable_metric) + "\n") self.metrics_buffer.clear() except Exception as e: logger.error(f"Error flushing metrics: {e}")