(gpu, args)
| 63 | |
| 64 | |
| 65 | def main_worker(gpu, args): |
| 66 | args.output_dir = os.path.join(args.output_folder, args.exp_name) |
| 67 | |
| 68 | # local rank & global rank |
| 69 | args.gpu = gpu |
| 70 | args.rank = args.rank * args.ngpus_per_node + gpu |
| 71 | torch.cuda.set_device(args.gpu) |
| 72 | |
| 73 | # logger |
| 74 | setup_logger(args.output_dir, |
| 75 | distributed_rank=args.gpu, |
| 76 | filename="train.log", |
| 77 | mode="a") |
| 78 | # dist init |
| 79 | dist.init_process_group(backend=args.dist_backend, |
| 80 | init_method=args.dist_url, |
| 81 | world_size=args.world_size, |
| 82 | rank=args.rank) |
| 83 | # wandb |
| 84 | if args.rank == 0: |
| 85 | wandb.init(job_type="training", |
| 86 | mode="offline", |
| 87 | config=args, |
| 88 | project=args.exp_name, |
| 89 | name=args.exp_name, |
| 90 | tags=[args.dataset]) |
| 91 | dist.barrier() |
| 92 | # build model |
| 93 | model, param_list = build_segmenter(args) |
| 94 | # logger.info(model) |
| 95 | logger.info(args) |
| 96 | |
| 97 | # build optimizer & lr scheduler |
| 98 | optimizer = torch.optim.AdamW(param_list, |
| 99 | lr=args.lr, |
| 100 | weight_decay=args.weight_decay, |
| 101 | amsgrad=args.amsgrad |
| 102 | ) |
| 103 | |
| 104 | scaler = amp.GradScaler() |
| 105 | |
| 106 | # build dataset |
| 107 | args.batch_size = int(args.batch_size / args.ngpus_per_node) |
| 108 | args.batch_size_val = int(args.batch_size_val / args.ngpus_per_node) |
| 109 | args.workers = int( |
| 110 | (args.workers + args.ngpus_per_node - 1) / args.ngpus_per_node) |
| 111 | train_data = RefDataset(lmdb_dir=args.train_lmdb, |
| 112 | mask_dir=args.mask_root, |
| 113 | dataset=args.dataset, |
| 114 | split=args.train_split, |
| 115 | mode='train', |
| 116 | input_size=args.input_size, |
| 117 | word_length=args.word_len |
| 118 | ) |
| 119 | val_data = RefDataset(lmdb_dir=args.val_lmdb, |
| 120 | mask_dir=args.mask_root, |
| 121 | dataset=args.dataset, |
| 122 | split=args.val_split, |
nothing calls this directly
no test coverage detected