-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
38 lines (31 loc) · 1.4 KB
/
Copy pathmain.py
File metadata and controls
38 lines (31 loc) · 1.4 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
import pprint
import torch
import wandb
import os
from args_helper import args
from logger import setup_logger
from regularize import trainer as regularize_trainer
from utils import set_seed, get_dataset, get_model, get_optimizer, get_scheduler, update_iter_and_epochs, get_criterion
def main():
logger = setup_logger(name=args.logger_name, args=args)
logger.info("Call with args: \n{}".format(pprint.pformat(vars(args))))
project_name = "{}_{}".format(args.which_dataset.lower(), args.arch.lower())
wandb.init(project=project_name, name=args.logger_name, config=vars(args))
set_seed(args.seed, logger)
if args.cuda:
device = torch.device("cuda:{}".format(args.gpu))
else:
device = torch.device("cpu")
torch.set_num_threads(4)
logger.info("Using device {}".format(device))
dataset = get_dataset(args=args, logger=logger)
update_iter_and_epochs(dataset=dataset, args=args, logger=logger)
model = get_model(args=args, logger=logger, dataset=dataset).to(device)
criterion = get_criterion(criterion_type=args.criterion)
optimizer = get_optimizer(args=args, model=model)
scheduler = get_scheduler(optimizer=optimizer, logger=logger, args=args)
regularize_trainer(
dataset=dataset, device=device, model=model,
args=args, optimizer=optimizer, scheduler=scheduler, criterion=criterion, wandb=wandb)
if __name__ == "__main__":
main()