MCPcopy Create free account
hub / github.com/TEA-Lab/TwoByTwo / train

Function train

src/script/our_train_B.py:43–268  ·  view source on GitHub ↗
(rank, world_size, conf)

Source from the content-addressed store, hash-verified

41 dist.destroy_process_group()
42
43def 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

Callers

nothing calls this directly

Calls 7

load_dataMethod · 0.95
OurDatasetClass · 0.90
setupFunction · 0.70
cleanupFunction · 0.70
training_stepMethod · 0.45
forward_passMethod · 0.45

Tested by

no test coverage detected