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()
|