MCPcopy Create free account
hub / github.com/ImprintLab/Medical-SAM2 / main

Function main

train_2d.py:26–120  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

24
25
26def 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

Callers 1

train_2d.pyFile · 0.70

Calls 5

REFUGEClass · 0.85
get_networkFunction · 0.50
set_log_dirFunction · 0.50
create_loggerFunction · 0.50
deviceMethod · 0.45

Tested by

no test coverage detected