diff options
| author | gdamms <damguillotin@gmail.com> | 2026-02-06 11:07:37 +0100 |
|---|---|---|
| committer | gdamms <damguillotin@gmail.com> | 2026-02-06 11:07:37 +0100 |
| commit | f1fee867bbbffdd2686cc2f92313496f77d08b2f (patch) | |
| tree | e4a2cd033c8d9fdc372c6fbe033436db68183a2d /main.py | |
| parent | 8ee90daeedef75135ab50dd2bd5c5349929b53dd (diff) | |
| download | diffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.tar.gz diffusion-mnist-f1fee867bbbffdd2686cc2f92313496f77d08b2f.zip | |
add patience for early stopping or infinit epochs
Diffstat (limited to 'main.py')
| -rw-r--r-- | main.py | 10 |
1 files changed, 8 insertions, 2 deletions
@@ -32,20 +32,26 @@ def main(): # Train diffusion model train_parser = subparsers.add_parser("train", help="Train diffusion model") - train_parser.add_argument("--epochs", type=int, default=10, help="Number of epochs") + train_parser.add_argument("--epochs", type=int, default=None, help="Number of epochs") train_parser.add_argument("--lr", type=float, default=2e-4, help="Learning rate") train_parser.add_argument("--batch-size", type=int, default=64, help="Batch size") train_parser.add_argument("--attention", action="store_true", help="Use self-attention") train_parser.add_argument("--checkpoint", type=str, default=None, help="Resume from checkpoint") train_parser.add_argument("--name", type=str, default=None, help="Run name") + train_parser.add_argument("--val-split", type=float, default=0.1, help="Validation split fraction") + train_parser.add_argument("--test-split", type=float, default=0.1, help="Test split fraction") + train_parser.add_argument("--patience", type=int, default=5, help="Early stopping patience") # Train autoencoder ae_parser = subparsers.add_parser("train-ae", help="Train autoencoder") - ae_parser.add_argument("--epochs", type=int, default=10, help="Number of epochs") + ae_parser.add_argument("--epochs", type=int, default=None, help="Number of epochs") ae_parser.add_argument("--lr", type=float, default=1e-3, help="Learning rate") ae_parser.add_argument("--batch-size", type=int, default=64, help="Batch size") ae_parser.add_argument("--latent-channels", type=int, default=1, help="Latent channels") ae_parser.add_argument("--checkpoint", type=str, default=None, help="Resume from checkpoint") + ae_parser.add_argument("--val-split", type=float, default=0.1, help="Validation split fraction") + ae_parser.add_argument("--test-split", type=float, default=0.1, help="Test split fraction") + ae_parser.add_argument("--patience", type=int, default=5, help="Early stopping patience") # Sample from model sample_parser = subparsers.add_parser("sample", help="Generate samples") |
