aboutsummaryrefslogtreecommitdiff
path: root/main.py
blob: 34640c63d8beb4b8f09c0d532a5ec63f95c5a395 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""
MNIST Diffusion Model

A diffusion-based generative model for MNIST digits.

Usage:
    Train diffusion model:
        python main.py train --epochs 10

    Train autoencoder:
        python main.py train-ae --epochs 10

    Generate samples:
        python main.py sample --checkpoint checkpoints/diffusion_latest.pt

    Visualize diffusion process:
        python main.py visualize --all
"""

import argparse
import torch


def main():
    parser = argparse.ArgumentParser(
        description="MNIST Diffusion Model",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=__doc__
    )

    subparsers = parser.add_subparsers(dest="command", help="Available commands")

    # Train diffusion model
    train_parser = subparsers.add_parser("train", help="Train diffusion model")
    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=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")
    sample_parser.add_argument("--checkpoint", type=str, default="checkpoints/diffusion_latest.pt",
                               help="Path to model checkpoint")
    sample_parser.add_argument("--n-samples", type=int, default=10, help="Samples per class")
    sample_parser.add_argument("--attention", action="store_true", help="Use attention in model")

    # Visualize diffusion
    viz_parser = subparsers.add_parser("visualize", help="Visualize diffusion process")
    viz_parser.add_argument("--checkpoint", type=str, default="checkpoints/diffusion_latest.pt",
                            help="Path to model checkpoint")
    viz_parser.add_argument("--attention", action="store_true", help="Use attention in model")
    viz_parser.add_argument("--forward", action="store_true", help="Visualize forward diffusion")
    viz_parser.add_argument("--backward", action="store_true", help="Visualize backward diffusion")
    viz_parser.add_argument("--all", action="store_true", help="Run all visualizations")

    args = parser.parse_args()

    if args.command is None:
        parser.print_help()
        return

    # Set multiprocessing start method
    torch.multiprocessing.set_start_method("spawn", force=True)

    if args.command == "train":
        from src.train_diffusion import train_diffusion
        train_diffusion(
            epochs=args.epochs,
            learning_rate=args.lr,
            batch_size=args.batch_size,
            use_attention=args.attention,
            checkpoint_path=args.checkpoint,
            run_name=args.name,
        )

    elif args.command == "train-ae":
        from src.train_autoencoder import train_autoencoder
        train_autoencoder(
            epochs=args.epochs,
            learning_rate=args.lr,
            batch_size=args.batch_size,
            latent_channels=args.latent_channels,
            checkpoint_path=args.checkpoint,
        )

    elif args.command == "sample":
        import os
        from src.config import DEVICE
        from src.sample import generate_grid
        from src.utils import load_checkpoint
        from models import UNetMNIST

        model = UNetMNIST(use_attention=args.attention).to(DEVICE)
        if os.path.exists(args.checkpoint):
            model = load_checkpoint(model, os.path.basename(args.checkpoint))
        else:
            print(f"Warning: Checkpoint {args.checkpoint} not found.")

        generate_grid(model, n_per_class=args.n_samples)

    elif args.command == "visualize":
        import os
        from src.config import DEVICE
        from src.sample import (
            visualize_forward_diffusion,
            visualize_backward_diffusion,
            generate_grid,
        )
        from src.utils import load_checkpoint
        from models import UNetMNIST

        model = UNetMNIST(use_attention=args.attention).to(DEVICE)
        if os.path.exists(args.checkpoint):
            model = load_checkpoint(model, os.path.basename(args.checkpoint))

        if args.forward or args.all:
            visualize_forward_diffusion()

        if args.backward or args.all:
            visualize_backward_diffusion(model)

        if args.all:
            generate_grid(model)


if __name__ == "__main__":
    main()