Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 7 additions & 7 deletions climanet/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,14 +121,14 @@ def train_monthly_model(
optimizer.zero_grad(set_to_none=True)

# Calculate average epoch loss
avg_epoch_loss = epoch_loss.item() / (i + 1)
writer.add_scalar("Loss/train", avg_epoch_loss, epoch)
avg_train_loss = epoch_loss.item() / (i + 1)
writer.add_scalar("Loss/train", avg_train_loss, epoch)
avg_epoch_loss = avg_train_loss # Initially use training loss

# Validation loss (optional)
if validation_dataset is not None:
# Store train loss for gap calculation
avg_train_loss = avg_epoch_loss
_, avg_epoch_loss = predict_monthly_var(
_, avg_val_loss = predict_monthly_var(
model,
validation_dataset,
batch_size=batch_size,
Expand All @@ -140,17 +140,17 @@ def train_monthly_model(
run_dir=run_dir,
dataloader_num_workers=dataloader_num_workers,
)
writer.add_scalar("Loss/validation", avg_epoch_loss, epoch)
writer.add_scalar("Loss/validation", avg_val_loss, epoch)
avg_epoch_loss = avg_val_loss # Use validation loss if exists

if verbose and epoch % verbose_epoch_interval == 0:
gap = avg_epoch_loss - avg_train_loss
gap = avg_val_loss - avg_train_loss
print(f"Epoch {epoch}: gap between train and val loss: {gap:.6f}")

# Step scheduler
scheduler.step(avg_epoch_loss)

# Log to TensorBoard
writer.add_scalar("Loss/train", avg_epoch_loss, epoch)
writer.add_scalar("Loss/best", best_loss, epoch)

# Early stopping check
Expand Down
6 changes: 4 additions & 2 deletions climanet/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import numpy as np
import xarray as xr
import torch
import time
import psutil

from torch.utils.tensorboard import SummaryWriter
Expand Down Expand Up @@ -231,7 +232,8 @@ def set_seed(seed: int = 42):
def setup_logging(log_dir: str) -> SummaryWriter:
"""Set up TensorBoard logging directory and writer."""
Path(log_dir).mkdir(parents=True, exist_ok=True)
return SummaryWriter(log_dir)
timestamp_utc = time.strftime("%Y%m%dT%H%M%S", time.gmtime())
return SummaryWriter(log_dir, filename_suffix=f"_UTC{timestamp_utc}")


def compute_masked_loss(
Expand Down Expand Up @@ -487,7 +489,7 @@ def plot_nobs_vs_err(

for i, ax in enumerate(axes):
ax.set_title(f"Month = {err_baseline.time.dt.strftime('%Y-%m-%d').values[i]}")

# Get unique number of observations for this month, ignoring NaNs and zeros
n_obs_unique = np.unique(nobs.isel(time=i).values)
n_obs_unique = n_obs_unique[(~np.isnan(n_obs_unique)) & (n_obs_unique > 0)]
Expand Down
158 changes: 107 additions & 51 deletions notebooks/example_daily.ipynb

Large diffs are not rendered by default.

1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ dependencies = [
"matplotlib>=3.10.8",
"netcdf4>=1.7.4",
"scipy>=1.17.1",
"tbparse>=0.0.9",
"tensorboard",
"torch>=2.10.0",
"xarray>=2025.12.0",
Expand Down
21 changes: 19 additions & 2 deletions tests/test_utils.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,19 @@
def test_dummy():
assert True
from climanet.utils import setup_logging
from tbparse import SummaryReader


def test_setup_logging(tmp_path):
writer = setup_logging(tmp_path)
log_text = "This is a test log entry."
writer.add_scalar("Test Scalar", 42)
writer.add_text("Test Text", log_text)
writer.close()

# Test that there is one event file
# The file should have "UTC" keyword in timestamp suffix
assert len(list(tmp_path.glob("events*UTC*"))) == 1

# Load the events file with SummaryReader
reader = SummaryReader(tmp_path)
assert reader.text["value"].iloc[0] == log_text # check text
assert reader.scalars["value"].iloc[0] == 42 # check scalar