()
| 24 | |
| 25 | |
| 26 | def main(): |
| 27 | # use bfloat16 for the entire work |
| 28 | torch.autocast(device_type="cuda", dtype=torch.bfloat16).__enter__() |
| 29 | |
| 30 | if torch.cuda.get_device_properties(0).major >= 8: |
| 31 | # turn on tfloat32 for Ampere GPUs (https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices) |
| 32 | torch.backends.cuda.matmul.allow_tf32 = True |
| 33 | torch.backends.cudnn.allow_tf32 = True |
| 34 | |
| 35 | |
| 36 | args = cfg.parse_args() |
| 37 | GPUdevice = torch.device('cuda', args.gpu_device) |
| 38 | |
| 39 | net = get_network(args, args.net, use_gpu=args.gpu, gpu_device=GPUdevice, distribution = args.distributed) |
| 40 | |
| 41 | # optimisation |
| 42 | optimizer = optim.Adam(net.parameters(), lr=args.lr, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False) |
| 43 | # scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5) |
| 44 | |
| 45 | '''load pretrained model''' |
| 46 | |
| 47 | args.path_helper = set_log_dir('logs', args.exp_name) |
| 48 | logger = create_logger(args.path_helper['log_path']) |
| 49 | logger.info(args) |
| 50 | |
| 51 | |
| 52 | '''segmentation data''' |
| 53 | transform_train = transforms.Compose([ |
| 54 | transforms.Resize((args.image_size,args.image_size)), |
| 55 | transforms.ToTensor(), |
| 56 | ]) |
| 57 | |
| 58 | transform_test = transforms.Compose([ |
| 59 | transforms.Resize((args.image_size, args.image_size)), |
| 60 | transforms.ToTensor(), |
| 61 | ]) |
| 62 | |
| 63 | |
| 64 | # example of REFUGE dataset |
| 65 | if args.dataset == 'REFUGE': |
| 66 | '''REFUGE data''' |
| 67 | refuge_train_dataset = REFUGE(args, args.data_path, transform = transform_train, mode = 'Training') |
| 68 | refuge_test_dataset = REFUGE(args, args.data_path, transform = transform_test, mode = 'Test') |
| 69 | |
| 70 | nice_train_loader = DataLoader(refuge_train_dataset, batch_size=args.b, shuffle=True, num_workers=2, pin_memory=True) |
| 71 | nice_test_loader = DataLoader(refuge_test_dataset, batch_size=args.b, shuffle=False, num_workers=2, pin_memory=True) |
| 72 | '''end''' |
| 73 | |
| 74 | |
| 75 | '''checkpoint path and tensorboard''' |
| 76 | checkpoint_path = os.path.join(settings.CHECKPOINT_PATH, args.net, settings.TIME_NOW) |
| 77 | #use tensorboard |
| 78 | if not os.path.exists(settings.LOG_DIR): |
| 79 | os.mkdir(settings.LOG_DIR) |
| 80 | writer = SummaryWriter(log_dir=os.path.join( |
| 81 | settings.LOG_DIR, args.net, settings.TIME_NOW)) |
| 82 | |
| 83 | #create checkpoint folder to save model |
no test coverage detected