-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpretrain.py
More file actions
84 lines (72 loc) · 2.79 KB
/
Copy pathpretrain.py
File metadata and controls
84 lines (72 loc) · 2.79 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
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
import os
import random
import torch
import numpy as np
from src.train import train
from src.datas import transforms
from src.datas.dataloader import get_dataloader
from src.models import mae_vit
from src.utils.args import get_train_args
from src.utils.logging import get_log_writer
from src.utils.optim import get_optimizer_lr_scheduler
def main(args):
random.seed(42)
np.random.seed(42)
torch.manual_seed(42)
log_writer = get_log_writer(args)
if args.transform == "instance_normalize":
dataloader = get_dataloader(
ispretrain=True,
annotations_file=args.annotation_file,
input_dir=args.input_dir,
val_annotations_file=args.val_annotation_file,
val_input_dir=args.val_input_dir,
batch_size=args.batch_size,
transform=transforms.InstanceNorm(),
num_workers=args.num_workers,
pin_memory=args.pin_memory,
)
elif args.transform == "normalize":
# TODO: calculate the mean and variance for each channel.
norm_mean = torch.Tensor(torch.load("src/datas/xpt_spe_mean.pth"))
norm_std = torch.Tensor(torch.load("src/datas/xpt_spe_std.pth"))
dataloader = get_dataloader(
ispretrain=True,
annotations_file=args.annotation_file,
input_dir=args.input_dir,
val_annotations_file=args.val_annotation_file,
val_input_dir=args.val_input_dir,
batch_size=args.batch_size,
transform=transforms.Normalize(norm_mean, norm_std),
num_workers=args.num_workers,
pin_memory=args.pin_memory,
)
elif args.transform == "log":
dataloader = get_dataloader(
ispretrain=True,
annotations_file=args.annotation_file,
input_dir=args.input_dir,
val_annotations_file=args.val_annotation_file,
val_input_dir=args.val_input_dir,
batch_size=args.batch_size,
transform=transforms.LogTransform(),
num_workers=args.num_workers,
pin_memory=args.pin_memory,
)
else:
raise NotImplementedError
if args.model == "base":
model = mae_vit.mae_vit_base_patch16(mask_ratio=args.mask_ratio)
elif args.model == "large":
model = mae_vit.mae_vit_large_patch16(mask_ratio=args.mask_ratio)
elif args.model == "huge":
model = mae_vit.mae_vit_huge_patch14(mask_ratio=args.mask_ratio)
optimizer, scheduler = get_optimizer_lr_scheduler(model.parameters(), args)
train.trainer(
model, dataloader, optimizer, scheduler, args.epochs, log_writer, args
)
model = model.cpu()
torch.save(model.state_dict(), os.path.join(args.output_dir, "model.ckpt"))
if __name__ == "__main__":
args = get_train_args()
main(args)