Build decoders for each pose.
(self, mean_poses_dict)
| 210 | num_stages=3) |
| 211 | |
| 212 | def load_param_decoder(self, mean_poses_dict): |
| 213 | """Build decoders for each pose.""" |
| 214 | start = 0 |
| 215 | mean_lst = [] |
| 216 | self.pose_param_decoders = {} |
| 217 | for pose_param in self.pose_param_conf: |
| 218 | pose_name = pose_param['name'] |
| 219 | num_angles = pose_param['num_angles'] |
| 220 | if pose_param['use_mean']: |
| 221 | pose_decoder = ContinuousRotReprDecoder( |
| 222 | num_angles, |
| 223 | dtype=torch.float32, |
| 224 | mean=mean_poses_dict.get(pose_name, None)) |
| 225 | else: |
| 226 | pose_decoder = ContinuousRotReprDecoder(num_angles, |
| 227 | dtype=torch.float32, |
| 228 | mean=None) |
| 229 | self.pose_param_decoders['{}_decoder'.format( |
| 230 | pose_name)] = pose_decoder |
| 231 | pose_dim = pose_decoder.get_dim_size() |
| 232 | pose_mean = pose_decoder.get_mean() |
| 233 | if pose_param['rotate_axis_x']: |
| 234 | pose_mean[3] = -1 |
| 235 | idxs = list(range(start, start + pose_dim)) |
| 236 | idxs = torch.tensor(idxs, dtype=torch.long) |
| 237 | self.register_buffer('{}_idxs'.format(pose_name), idxs) |
| 238 | start += pose_dim |
| 239 | mean_lst.append(pose_mean.view(-1)) |
| 240 | return start, mean_lst |
| 241 | |
| 242 | def get_camera_param(self, camera_cfg): |
| 243 | """Build camera param.""" |
no test coverage detected