diff options
| author | gdamms <damguillotin@gmail.com> | 2026-02-05 19:55:58 +0100 |
|---|---|---|
| committer | gdamms <damguillotin@gmail.com> | 2026-02-05 19:55:58 +0100 |
| commit | f1be02af8c33d4136ad7c90dd43158a8ff6e5af1 (patch) | |
| tree | 3d7facac2732046ce1447156afe7dacc3b64d585 /src/train_diffusion.py | |
| parent | a5d5f30fbd9c6c7c78834072401932c84bddaf14 (diff) | |
| download | diffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.tar.gz diffusion-mnist-f1be02af8c33d4136ad7c90dd43158a8ff6e5af1.zip | |
better plots, semi working autoencoder
Diffstat (limited to 'src/train_diffusion.py')
| -rw-r--r-- | src/train_diffusion.py | 35 |
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__": |
