(config, labels, output, smpl)
| 181 | return img_path_list, img_dir, fps |
| 182 | |
| 183 | def get_estimates(config, labels, output, smpl): |
| 184 | with torch.enable_grad(): |
| 185 | init_depth = labels['joint_root'][:,2].clone().view(-1,config.sampling.multihypo_n+2).detach() |
| 186 | n = init_depth.shape[0] |
| 187 | depth_model = DeltaDepth(n).cuda(init_depth.device) |
| 188 | depth_model.train() |
| 189 | optim_list = [{"params":depth_model.delta_d,"lr":config.inference.optim_lr}] |
| 190 | optimizer = torch.optim.Adam(optim_list) |
| 191 | scheduler = torch.optim.lr_scheduler.StepLR(optimizer,step_size=config.inference.step_size, gamma=config.inference.gamma) |
| 192 | for i in range(config.inference.optim_step): |
| 193 | optimizer.zero_grad() |
| 194 | new_label = {} |
| 195 | for k in labels.keys(): |
| 196 | try: |
| 197 | new_label[k] = labels[k].detach().clone() |
| 198 | except Exception: |
| 199 | pass |
| 200 | new_output = {} |
| 201 | for k in output.keys(): |
| 202 | new_output[k] = output[k].detach().clone() |
| 203 | new_label['joint_root'][:,2] = depth_model(init_depth).view(-1) |
| 204 | output_final = process_output(new_output,new_label,smpl,process=True) |
| 205 | output_final['loss'].backward(retain_graph=True) |
| 206 | optimizer.step() |
| 207 | scheduler.step() |
| 208 | |
| 209 | return output_final |
no test coverage detected