aboutsummaryrefslogtreecommitdiff
path: root/src/train_diffusion.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/train_diffusion.py')
-rw-r--r--src/train_diffusion.py35
1 files changed, 21 insertions, 14 deletions
diff --git a/src/train_diffusion.py b/src/train_diffusion.py
index cb70650..2ec4d60 100644
--- a/src/train_diffusion.py
+++ b/src/train_diffusion.py
@@ -18,7 +18,8 @@ import os
import torch
import torch.nn as nn
import numpy as np
-import matplotlib.pyplot as plt
+import plotly.graph_objects as go
+from plotly.subplots import make_subplots
from rich.progress import track
import mlflow
@@ -156,20 +157,26 @@ def evaluate_and_log(model: nn.Module, epoch: int, predict_x0: bool = True):
mlflow.log_metric("KL Divergence", kl_score, step=epoch)
mlflow.log_metric("JSD", jsd_score, step=epoch)
- # Log sample images
- fig, axes = plt.subplots(4, 8, figsize=(16, 8))
- for i, ax in enumerate(axes.flat):
- if i < len(fakes):
- ax.imshow(fakes[i].transpose(1, 2, 0).squeeze(), cmap='gray')
- ax.axis('off')
- fig.suptitle(f"Generated Samples - Epoch {epoch}")
- plt.tight_layout()
+ # Log sample images using plotly
+ fig = make_subplots(rows=4, cols=8, horizontal_spacing=0.01, vertical_spacing=0.02)
+ for i in range(min(32, len(fakes))):
+ row = i // 8 + 1
+ col = i % 8 + 1
+ img = fakes[i].transpose(1, 2, 0).squeeze()[::-1]
+ fig.add_trace(
+ go.Heatmap(z=img, colorscale='gray', showscale=False),
+ row=row, col=col
+ )
+ fig.update_layout(
+ title_text=f"Generated Samples - Epoch {epoch}",
+ width=800,
+ height=400,
+ showlegend=False
+ )
+ fig.update_xaxes(showticklabels=False, showgrid=False, zeroline=False)
+ fig.update_yaxes(showticklabels=False, showgrid=False, zeroline=False)
- mlflow.log_figure(fig, f"samples_epoch_{epoch:03d}.png")
-
- # Save to plots folder
- fig.savefig(os.path.join(PLOTS_DIR, f"samples_epoch_{epoch:03d}.png"))
- plt.close(fig)
+ mlflow.log_figure(fig, f"samples/epoch_{epoch:03d}.png")
if __name__ == "__main__":