aboutsummaryrefslogtreecommitdiff
path: root/main.py
diff options
context:
space:
mode:
Diffstat (limited to 'main.py')
-rw-r--r--main.py10
1 files changed, 8 insertions, 2 deletions
diff --git a/main.py b/main.py
index 1c05580..34640c6 100644
--- a/main.py
+++ b/main.py
@@ -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")