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

Class ShapeAssemblyNet_B_vnn

src/shape_assembly/models/train/network_vnn_B.py:146–444  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

144
145
146class ShapeAssemblyNet_B_vnn(nn.Module):
147
148 def __init__(self, cfg, data_features):
149 super().__init__()
150 self.cfg = cfg
151 self.encoder = self.init_encoder()
152
153 self.encoder_dgcnn = DGCNN_New(feat_dim=cfg.model.pc_feat_dim)
154
155 self.pose_predictor_rot = self.init_pose_predictor_rot()
156 self.pose_predictor_trans = self.init_pose_predictor_trans()
157 if self.cfg.model.recon_loss:
158 self.decoder = self.init_decoder()
159 self.data_features = data_features
160
161 self.iter_counts = 0
162 self.close_eps = 0.1
163 self.L2 = nn.MSELoss()
164 self.R = torch.tensor([[0.26726124, -0.57735027, 0.77151675],
165 [0.53452248, -0.57735027, -0.6172134],
166 [0.80178373, 0.57735027, 0.15430335]], dtype=torch.float64).unsqueeze(0)
167 self.chamLoss = dist_chamfer_3D.chamfer_3DDist()
168
169 self.mlp_color = nn.Sequential(
170 nn.Linear(512*2*3, 1024))
171
172
173 def init_encoder(self):
174 if self.cfg.model.encoderB == 'dgcnn':
175 encoder = DGCNN(feat_dim=self.cfg.model.pc_feat_dim)
176 elif self.cfg.model.encoderB == 'vn_dgcnn':
177 encoder = VN_DGCNN_New(feat_dim=self.cfg.model.pc_feat_dim)
178 elif self.cfg.model.encoderB == 'pointnet':
179 encoder = PointNet(feat_dim=self.cfg.model.pc_feat_dim)
180 return encoder
181
182 def init_pose_predictor_rot(self):
183 if self.cfg.model.encoderB == 'vn_dgcnn':
184 pc_feat_dim = self.cfg.model.pc_feat_dim * 2 * 3
185 if self.cfg.model.pose_predictor_rot == 'original':
186 pose_predictor_rot = Regressor_CR(pc_feat_dim= pc_feat_dim, out_dim=6)
187 elif self.cfg.model.pose_predictor_rot == 'vn':
188 pose_predictor_rot = VN_equ_Regressor(pc_feat_dim= pc_feat_dim/3, out_dim=6)
189
190 return pose_predictor_rot
191
192 def init_pose_predictor_trans(self):
193 if self.cfg.model.encoderB == 'vn_dgcnn':
194 pc_feat_dim = self.cfg.model.pc_feat_dim * 2 * 3
195 if self.cfg.model.pose_predictor_trans == 'original':
196 pose_predictor_trans = Regressor_CR(pc_feat_dim=pc_feat_dim, out_dim=3)
197 elif self.cfg.model.pose_predictor_trans == 'vn':
198 pose_predictor_trans = VN_inv_Regressor(pc_feat_dim=pc_feat_dim/3, out_dim=3)
199 return pose_predictor_trans
200
201 def init_decoder(self):
202 pc_feat_dim = self.cfg.model.pc_feat_dim
203 decoder = MLPDecoder(feat_dim=pc_feat_dim, num_points=self.cfg.data.num_pc_points)

Callers 2

trainFunction · 0.90
trainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected