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

Function train

src/script/our_train_A.py:47–273  ·  view source on GitHub ↗
(rank, world_size, conf)

Source from the content-addressed store, hash-verified

45 dist.destroy_process_group()
46
47def train(rank, world_size, conf):
48 os.makedirs(conf.exp.log_dir, exist_ok=True)
49 os.makedirs(conf.exp.vis_dir, exist_ok=True)
50 setup(rank, world_size, conf)
51
52 if dist.get_rank() == 0:
53 wandb.init(project='shape-matching', notes='weak baseline', config=conf)
54 wandb.define_metric("test/epoch/*", step_metric="test/epoch/epoch")
55 wandb.define_metric("train/network_A/*", step_metric="train/network_A/step")
56
57 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']
58
59 network_A = ShapeAssemblyNet_A_vnn(cfg=conf, data_features=data_features)
60 network_A.cuda(rank)
61 network_A = DDP(network_A, device_ids=[rank])
62
63 # Initialize train dataloader
64 train_data = OurDataset(
65 data_root_dir=conf.data.root_dir,
66 data_csv_file=conf.data.train_csv_file,
67 data_features=data_features,
68 num_points=conf.data.num_pc_points
69 )
70 train_data.load_data()
71
72 print('Len of Train Data: ', len(train_data))
73
74 train_sampler = DistributedSampler(train_data, num_replicas=world_size, rank=rank)
75 train_dataloader = DataLoader(
76 dataset=train_data,
77 batch_size=conf.exp.batch_size,
78 num_workers=conf.exp.num_workers,
79 pin_memory=True,
80 shuffle=False,
81 drop_last=False,
82 sampler=train_sampler
83 )
84
85 # Initialize val dataloader
86 val_data = OurDataset(
87 data_root_dir=conf.data.root_dir,
88 data_csv_file=conf.data.val_csv_file,
89 data_features=data_features,
90 num_points=conf.data.num_pc_points
91 )
92 val_data.load_data()
93
94 # Output the distribution of the validation data
95 print('Len of Val Data: ', len(val_data))
96
97 val_sampler = DistributedSampler(val_data, num_replicas=world_size, rank=rank)
98 val_dataloader = DataLoader(
99 dataset=val_data,
100 batch_size=conf.exp.batch_size,
101 num_workers=conf.exp.num_workers,
102 pin_memory=True,
103 # shuffle=True,
104 shuffle=False,

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