(self, state, construct_dpo=False)
| 294 | |
| 295 | |
| 296 | def sample(self, state, construct_dpo=False): |
| 297 | args, config = self.args, self.config |
| 298 | |
| 299 | model = state['model'] |
| 300 | model_cond = state['model_cond'] |
| 301 | model_score = state['model_score'] |
| 302 | model_score_cond = state['model_score_cond'] |
| 303 | dataset = state['dataset'] |
| 304 | dataloader = state['dataloader'] |
| 305 | epoch = state['epoch'] |
| 306 | dataset_type = state['dataset_type'] |
| 307 | |
| 308 | rank = torch.distributed.get_rank() |
| 309 | |
| 310 | multi_n = self.multihypo_n |
| 311 | with torch.no_grad(): |
| 312 | pred = {} |
| 313 | dpo_pair_list = [] |
| 314 | for i in range(multi_n+2): |
| 315 | pred[i] = {} |
| 316 | for i, (inps, labels, img_ids, bboxes, img_path) in enumerate(tqdm(dataloader, ncols=100)): |
| 317 | n = inps.size(0) |
| 318 | output = {} |
| 319 | for k, _ in labels.items(): |
| 320 | labels[k] = labels[k].to(self.device) |
| 321 | input = inps.float().to(self.device) |
| 322 | save_trans_inv = labels['trans_inv'].cpu().numpy().astype(np.float32) |
| 323 | scale = torch.tensor([self.image_size[0],self.image_size[1],dataset.bbox_3d_shape[2]]).float().to(self.device) / self.config.diffusion.scale |
| 324 | labels['trans_inv']= labels['trans_inv'].unsqueeze(1).repeat(1,multi_n+2,1,1).view(-1,2,3) |
| 325 | labels['intrinsic_param']= labels['intrinsic_param'].unsqueeze(1).repeat(1,multi_n+2,1,1).view(-1,3,3) |
| 326 | labels['joint_root']= labels['joint_root'].unsqueeze(1).repeat(1,multi_n+2,1).view(-1,3) |
| 327 | |
| 328 | state['input'] = input.clone() |
| 329 | state['scale'] = scale.clone() |
| 330 | output, score_input = self.gen_mesh(state, multi_n) # get denoised twist and joints |
| 331 | |
| 332 | if construct_dpo: |
| 333 | # construct with GT as winner |
| 334 | save_trans_inv = save_trans_inv.tolist() |
| 335 | save_pred_joints = output['pred_joints'].cpu().numpy().astype(np.float32).tolist() |
| 336 | save_pred_twist = output['pred_twist'].cpu().numpy().astype(np.float32).tolist() |
| 337 | save_score_joint = score_input['joint'].cpu().numpy().astype(np.float32).tolist() |
| 338 | save_score_twist = score_input['twist'].cpu().numpy().astype(np.float32).tolist() |
| 339 | save_joints_uvd_29 = labels['joints_uvd_29'].cpu().numpy().astype(np.float32).tolist() |
| 340 | save_w_twist = labels['target_twist'].cpu().numpy().astype(np.float32).tolist() |
| 341 | save_trans = labels['trans'].cpu().numpy().astype(np.float32).tolist() |
| 342 | # import ipdb; ipdb.set_trace() |
| 343 | for idx in range(n): |
| 344 | for idx_hypo in range(multi_n): |
| 345 | dpo_pair = { |
| 346 | 'img_path': img_path[idx], |
| 347 | # 'db_idx': labels['db_idx'][idx].item(), |
| 348 | 'img_idx': img_ids[idx].item(), # absolute img_ids |
| 349 | 'l_joints': save_pred_joints[idx][idx_hypo], |
| 350 | 'l_twist': save_pred_twist[idx][idx_hypo], |
| 351 | 'w_joints': save_joints_uvd_29[idx], |
| 352 | 'w_twist': save_w_twist[idx], |
| 353 | 'trans': save_trans[idx], |
no test coverage detected