-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathtrain.py
More file actions
33 lines (24 loc) · 1.35 KB
/
Copy pathtrain.py
File metadata and controls
33 lines (24 loc) · 1.35 KB
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
import argparse
import torch
from noise2noise import Model
def main():
parser = argparse.ArgumentParser(description="Train a Noise2Noise denoising model")
parser.add_argument("--train-data", type=str, required=True, help="Path to training data pickle file")
parser.add_argument("--epochs", type=int, default=10, help="Number of training epochs (default: 10)")
parser.add_argument("--batch-size", type=int, default=8, help="Batch size (default: 8)")
parser.add_argument("--num-workers", type=int, default=2, help="Number of data loading workers (default: 2)")
parser.add_argument("--lr", type=float, default=0.001, help="Learning rate (default: 0.001)")
parser.add_argument("--output", type=str, default="model_pytorch.pth", help="Output path for trained model")
parser.add_argument("--pretrained", type=str, default=None, help="Path to pretrained model to resume from")
args = parser.parse_args()
noisy_imgs_1, noisy_imgs_2 = torch.load(args.train_data, weights_only=True)
model = Model(lr=args.lr)
if args.pretrained:
model.load_pretrained_model(args.pretrained)
model.train(
noisy_imgs_1, noisy_imgs_2, num_epochs=args.epochs, batch_size=args.batch_size, num_workers=args.num_workers
)
model.save(args.output)
print(f"Model saved to {args.output}")
if __name__ == "__main__":
main()