1
Fork 0

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
This commit is contained in:
Tibo De Peuter 2026-03-31 20:52:51 +00:00
parent dc3531071f
commit 8ec693f04c
2 changed files with 14 additions and 1 deletions

3
.gitignore vendored
View file

@ -2,6 +2,9 @@
artifacts/* artifacts/*
runs/* runs/*
# Experiment tracking
wandb/
# Python-generated files # Python-generated files
__pycache__/ __pycache__/
*.py[oc] *.py[oc]

View file

@ -13,6 +13,7 @@ from pathlib import Path
from typing import Any, Dict, Optional from typing import Any, Dict, Optional
import flax import flax
import jax.numpy as jnp
import numpy as np import numpy as np
from experiment_logger.wandb_utils import finish_wandb, init_wandb from experiment_logger.wandb_utils import finish_wandb, init_wandb
@ -153,7 +154,16 @@ class UnifiedLogger:
metrics_file = self.metrics_dir / "metrics.jsonl" metrics_file = self.metrics_dir / "metrics.jsonl"
with open(metrics_file, "a") as f: with open(metrics_file, "a") as f:
for metric in self.metrics_buffer: 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() self.metrics_buffer.clear()
except Exception as e: except Exception as e:
logger.error(f"Error flushing metrics: {e}") logger.error(f"Error flushing metrics: {e}")