| 144 | |
| 145 | |
| 146 | class 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) |