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

Function main

train_3d.py:21–108  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

19from func_3d.dataset import get_dataloader
20
21def main():
22
23 args = cfg.parse_args()
24
25 GPUdevice = torch.device('cuda', args.gpu_device)
26
27 net = get_network(args, args.net, use_gpu=args.gpu, gpu_device=GPUdevice, distribution = args.distributed)
28 net.to(dtype=torch.bfloat16)
29 if args.pretrain:
30 print(args.pretrain)
31 weights = torch.load(args.pretrain)
32 net.load_state_dict(weights,strict=False)
33
34 sam_layers = (
35 []
36 # + list(net.image_encoder.parameters())
37 # + list(net.sam_prompt_encoder.parameters())
38 + list(net.sam_mask_decoder.parameters())
39 )
40 mem_layers = (
41 []
42 + list(net.obj_ptr_proj.parameters())
43 + list(net.memory_encoder.parameters())
44 + list(net.memory_attention.parameters())
45 + list(net.mask_downsample.parameters())
46 )
47 if len(sam_layers) == 0:
48 optimizer1 = None
49 else:
50 optimizer1 = optim.Adam(sam_layers, lr=1e-4, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False)
51 if len(mem_layers) == 0:
52 optimizer2 = None
53 else:
54 optimizer2 = optim.Adam(mem_layers, lr=1e-8, betas=(0.9, 0.999), eps=1e-08, weight_decay=0, amsgrad=False)
55 # scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5) #learning rate decay
56
57 torch.autocast(device_type="cuda", dtype=torch.bfloat16).__enter__()
58
59 if torch.cuda.get_device_properties(0).major >= 8:
60 # turn on tfloat32 for Ampere GPUs (https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices)
61 torch.backends.cuda.matmul.allow_tf32 = True
62 torch.backends.cudnn.allow_tf32 = True
63
64 args.path_helper = set_log_dir('logs', args.exp_name)
65 logger = create_logger(args.path_helper['log_path'])
66 logger.info(args)
67
68 nice_train_loader, nice_test_loader = get_dataloader(args)
69
70 '''checkpoint path and tensorboard'''
71 checkpoint_path = os.path.join(settings.CHECKPOINT_PATH, args.net, settings.TIME_NOW)
72 #use tensorboard
73 if not os.path.exists(settings.LOG_DIR):
74 os.mkdir(settings.LOG_DIR)
75 writer = SummaryWriter(log_dir=os.path.join(
76 settings.LOG_DIR, args.net, settings.TIME_NOW))
77
78 #create checkpoint folder to save model

Callers 1

train_3d.pyFile · 0.70

Calls 5

get_networkFunction · 0.90
set_log_dirFunction · 0.90
create_loggerFunction · 0.90
get_dataloaderFunction · 0.90
deviceMethod · 0.45

Tested by

no test coverage detected