(rank, world_size, conf)
| 41 | dist.destroy_process_group() |
| 42 | |
| 43 | def train(rank, world_size, conf): |
| 44 | os.makedirs(conf.exp.log_dir, exist_ok=True) |
| 45 | os.makedirs(conf.exp.vis_dir, exist_ok=True) |
| 46 | setup(rank, world_size, conf) |
| 47 | |
| 48 | if dist.get_rank() == 0: |
| 49 | wandb.init(project='shape-matching', notes='weak baseline', config=conf) |
| 50 | wandb.define_metric("test/epoch/*", step_metric="test/epoch/epoch") |
| 51 | wandb.define_metric("train/network_B/*", step_metric="train/network_B/step") |
| 52 | |
| 53 | data_features = ['src_pc', 'src_rot', 'src_trans', 'tgt_pc', 'tgt_rot', 'tgt_trans', 'partA_symmetry_type', 'partB_symmetry_type','predicted_partB_rotation', 'predicted_partB_position', 'predicted_partA_rotation', 'predicted_partA_position'] |
| 54 | |
| 55 | network_B = ShapeAssemblyNet_B_vnn(cfg=conf, data_features=data_features) |
| 56 | network_B.cuda(rank) |
| 57 | network_B = DDP(network_B, device_ids=[rank]) |
| 58 | |
| 59 | # Initialize train dataloader |
| 60 | train_data = OurDataset( |
| 61 | data_root_dir=conf.data.root_dir, |
| 62 | data_csv_file=conf.data.train_csv_file, |
| 63 | data_features=data_features, |
| 64 | num_points=conf.data.num_pc_points |
| 65 | ) |
| 66 | train_data.load_data() |
| 67 | |
| 68 | print('Len of Train Data: ', len(train_data)) |
| 69 | |
| 70 | train_sampler = DistributedSampler(train_data, num_replicas=world_size, rank=rank) |
| 71 | train_dataloader = DataLoader( |
| 72 | dataset=train_data, |
| 73 | batch_size=conf.exp.batch_size, |
| 74 | num_workers=conf.exp.num_workers, |
| 75 | pin_memory=True, |
| 76 | shuffle=False, |
| 77 | drop_last=False, |
| 78 | sampler=train_sampler |
| 79 | ) |
| 80 | |
| 81 | # Initialize val dataloader |
| 82 | val_data = OurDataset( |
| 83 | data_root_dir=conf.data.root_dir, |
| 84 | data_csv_file=conf.data.val_csv_file, |
| 85 | data_features=data_features, |
| 86 | num_points=conf.data.num_pc_points |
| 87 | ) |
| 88 | val_data.load_data() |
| 89 | |
| 90 | print('Len of Val Data: ', len(val_data)) |
| 91 | |
| 92 | val_sampler = DistributedSampler(val_data, num_replicas=world_size, rank=rank) |
| 93 | val_dataloader = DataLoader( |
| 94 | dataset=val_data, |
| 95 | batch_size=conf.exp.batch_size, |
| 96 | num_workers=conf.exp.num_workers, |
| 97 | pin_memory=True, |
| 98 | shuffle=False, |
| 99 | drop_last=False, |
| 100 | sampler=val_sampler |
nothing calls this directly
no test coverage detected