(self, scene)
| 442 | return scaled |
| 443 | |
| 444 | def predict(self, scene): |
| 445 | if self.model is None: |
| 446 | raise RuntimeError("R3 model is not loaded") |
| 447 | |
| 448 | _activate_r3_imports() |
| 449 | from R3.utils.pose_enc import pose_encoding_to_extri_intri |
| 450 | |
| 451 | images_raw = scene["images_raw"] |
| 452 | n, _, h, w = images_raw.shape |
| 453 | images, proc_hw, resized = self._prepare_images(images_raw) |
| 454 | h_proc, w_proc = proc_hw |
| 455 | |
| 456 | amp_dtype = torch.bfloat16 if self.amp_dtype == "bf16" else torch.float16 |
| 457 | use_cuda_amp = bool(self.use_amp) and torch.device(self.device).type == "cuda" |
| 458 | amp_context = ( |
| 459 | torch.autocast(device_type="cuda", dtype=amp_dtype) |
| 460 | if use_cuda_amp else nullcontext() |
| 461 | ) |
| 462 | |
| 463 | rel_pose_kwargs = self._rel_pose_kwargs() |
| 464 | with torch.no_grad(): |
| 465 | with amp_context: |
| 466 | if hasattr(self.model, "clear_online_state"): |
| 467 | self.model.clear_online_state() |
| 468 | predictions = self.model( |
| 469 | images, |
| 470 | mode=self.attention_mode or "causal", |
| 471 | use_ray_pose=False, |
| 472 | pose_max_recent=int(self.pose_max_recent), |
| 473 | bootstrap_full_attention_frames=int(self.bootstrap_full_attention_frames), |
| 474 | online_finalize_pose_reconstruction=bool( |
| 475 | self.online_finalize_pose_reconstruction |
| 476 | ), |
| 477 | rel_pose_reconstruction_method=self.rel_pose_reconstruction_method, |
| 478 | rel_pose_reconstruction_kwargs=rel_pose_kwargs, |
| 479 | ) |
| 480 | |
| 481 | output_frame_ids = predictions.get("output_frame_ids", list(range(n))) |
| 482 | output_frame_ids = [int(frame_id) for frame_id in output_frame_ids] |
| 483 | if output_frame_ids != list(range(n)): |
| 484 | raise ValueError( |
| 485 | "R3 returned a subset or reordered frames " |
| 486 | f"(output_frame_ids={output_frame_ids}); disable output eviction for benchmark evaluation." |
| 487 | ) |
| 488 | |
| 489 | result = {} |
| 490 | |
| 491 | depth = predictions.get("depth") |
| 492 | if isinstance(depth, torch.Tensor): |
| 493 | pred_depth = depth[0, :, :, :, 0].detach().cpu().float().numpy() |
| 494 | if resized: |
| 495 | pred_depth = self._resize_nhw(pred_depth, (h, w)) |
| 496 | result["pred_depth"] = pred_depth.astype(np.float32) |
| 497 | |
| 498 | depth_conf = predictions.get("depth_conf") |
| 499 | if isinstance(depth_conf, torch.Tensor): |
| 500 | pred_conf = depth_conf[0].detach().cpu().float().numpy() |
| 501 | if resized: |
no test coverage detected