MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / sample

Method sample

ADHMR/lib/runners/scorenet.py:296–507  ·  view source on GitHub ↗
(self, state, construct_dpo=False)

Source from the content-addressed store, hash-verified

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],

Callers 1

validateMethod · 0.95

Calls 10

gen_meshMethod · 0.95
process_predFunction · 0.90
printFunction · 0.85
removeMethod · 0.80
trans_backFunction · 0.50
process_outputFunction · 0.50
itemsMethod · 0.45
toMethod · 0.45
loadMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected