MCPcopy Create free account
hub / github.com/Ropedia/SpatialBench / predict

Method predict

benchmark/evaluation/model_adapters/r3_adapter.py:444–518  ·  view source on GitHub ↗
(self, scene)

Source from the content-addressed store, hash-verified

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:

Callers 3

_measure_oneFunction · 0.45
mainFunction · 0.45
mainFunction · 0.45

Calls 10

_prepare_imagesMethod · 0.95
_rel_pose_kwargsMethod · 0.95
_resize_nhwMethod · 0.95
_scale_intrinsicsMethod · 0.95
_activate_r3_importsFunction · 0.85
clear_online_stateMethod · 0.80
_invert_se3Method · 0.80
deviceMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected