()
| 19 | from func_3d.dataset import get_dataloader |
| 20 | |
| 21 | def 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 |
no test coverage detected