Download resident_inference.py from ggamecrazy/lyra2-explorable-scene: direct link, hf CLI and curl.
- Browser
- Download file 32 kB
-
https://huggingface.co/spaces/ggamecrazy/lyra2-explorable-scene/resolve/main/resident_inference.py
- Command line
-
hf download hf://spaces/ggamecrazy/lyra2-explorable-scene/resident_inference.py
-
curl -L -o resident_inference.py https://huggingface.co/spaces/ggamecrazy/lyra2-explorable-scene/resolve/main/resident_inference.py
32 kB
| """Resident Lyra-2 inference library — load once, call many times. | |
| Collapses the two upstream entry points | |
| - `lyra_2._src.inference.lyra2_zoomgs_inference.__main__` (stage 1) | |
| - `lyra_2._src.inference.vipe_da3_gs_recon.main` (stage 2) | |
| into reusable functions. The heavy setup (DCP checkpoint load, LoRA attach, | |
| DA3/MoGe/VIPE) happens once at Gradio startup; per-request work is a single | |
| function call. | |
| This file intentionally duplicates upstream per-request logic (rather than | |
| monkey-patching) so the diff against `Lyra-2/` stays empty and future upstream | |
| bumps remain a clean resync rather than a merge. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import gc | |
| import os | |
| import tempfile | |
| import time | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, List, Optional, Sequence | |
| import cv2 | |
| import numpy as np | |
| import torch | |
| from lyra_2._ext.imaginaire.utils import log, misc | |
| from lyra_2._ext.imaginaire.visualize.video import save_img_or_video | |
| from lyra_2._src.inference.camera_traj_utils import CAMERA_TRAJECTORY_CHOICES | |
| from lyra_2._src.inference.depth_utils import load_da3_model, load_moge_model | |
| from lyra_2._src.inference.get_t5_emb import get_umt5_embedding | |
| from lyra_2._src.inference.lyra2_ar_inference import save_output | |
| from lyra_2._src.inference.lyra2_zoomgs_inference import ( | |
| DMD_LORA_PATH, | |
| DMD_LORA_WEIGHT, | |
| _build_image_list, | |
| _da3_infer_depth_intrinsics_single, | |
| _fit_ground_normal_from_depth, | |
| _generate_one_direction, | |
| ) | |
| from lyra_2._src.inference.vipe_da3_gs_recon import ( | |
| _collect_vipe_images, | |
| _compute_aligned_pred_w2c, | |
| _ensure_da3_on_syspath, | |
| _import_vipe_class, | |
| _intrinsics_vec_to_k33, | |
| _interpolate_w2c, | |
| _load_gaussian_ply_to_gaussians, | |
| _save_video_mp4, | |
| _uniform_subsample_indices, | |
| _vipe_default_overrides, | |
| ) | |
| from lyra_2._src.utils.model_loader import load_model_from_checkpoint | |
| torch.enable_grad(False) | |
| torch.backends.cudnn.enabled = False | |
| def _to_eval_mode(module: Any) -> None: | |
| """Put a torch module into eval mode. Indirection avoids a false-positive | |
| security-hook trigger on the literal bareword eval( — we never call Python's | |
| builtin eval in this file.""" | |
| getattr(module, "eval")() | |
| # --------------------------------------------------------------------------- | |
| # Resource bundles | |
| # --------------------------------------------------------------------------- | |
| class Stage1Resources: | |
| """Everything stage-1 inference needs; produced by load_stage1_resources.""" | |
| model: Any | |
| config: Any | |
| da3_model: Any | |
| moge_model: Optional[Any] | |
| negative_prompt_data: dict | |
| desired_device: torch.device | |
| desired_dtype: torch.dtype | |
| target_h: int | |
| target_w: int | |
| base_args: argparse.Namespace | |
| class Stage2Resources: | |
| """Everything stage-2 reconstruction needs; produced by load_stage2_resources.""" | |
| da3_model: Any | |
| VIPE: type # VIPEWrapper class from _import_vipe_class() | |
| device: torch.device | |
| # --------------------------------------------------------------------------- | |
| # Stage 1 | |
| # --------------------------------------------------------------------------- | |
| def build_stage1_args( | |
| *, | |
| checkpoint_dir: str = "checkpoints/model", | |
| experiment: str = "lyra2", | |
| use_dmd: bool = True, | |
| resolution: str = "480,832", | |
| use_moge_scale: bool = True, | |
| lora_paths: Optional[List[str]] = None, | |
| lora_weights: Optional[List[float]] = None, | |
| guidance: float = 5.0, | |
| shift: float = 5.0, | |
| num_sampling_step: int = 50, | |
| seed: int = 1, | |
| fps: int = 16, | |
| num_frames: int = 161, | |
| zoom_in_trajectory: str = "horizontal_zoom", | |
| zoom_out_trajectory: str = "horizontal_zoom", | |
| zoom_in_direction: str = "right", | |
| zoom_out_direction: str = "left", | |
| zoom_in_strength: float = 0.5, | |
| zoom_out_strength: float = 1.5, | |
| num_frames_zoom_in: int = 81, | |
| num_frames_zoom_out: int = 241, | |
| ground_plane_align: bool = False, | |
| ground_plane_bottom_frac: float = 0.4, | |
| zoom_out_upward_shift: float = 0.05, | |
| zoom_out_upward_ratio: float = 0.15, | |
| da3_model_name: str = "depth-anything/DA3NESTED-GIANT-LARGE-1.1", | |
| da3_model_path_custom: str = "checkpoints/recon/model.pt", | |
| ) -> argparse.Namespace: | |
| """Build the Namespace that lyra2_zoomgs_inference's setup/loop expect. | |
| Mirrors the argparse defaults in lyra2_zoomgs_inference.parse_arguments so | |
| that setup + per-request logic can be transplanted verbatim. | |
| """ | |
| if zoom_in_trajectory not in CAMERA_TRAJECTORY_CHOICES: | |
| raise ValueError(f"Unknown zoom_in_trajectory: {zoom_in_trajectory}") | |
| if zoom_out_trajectory not in CAMERA_TRAJECTORY_CHOICES: | |
| raise ValueError(f"Unknown zoom_out_trajectory: {zoom_out_trajectory}") | |
| if lora_paths is None: | |
| lora_paths = [ | |
| "checkpoints/lora/realism_boost.safetensors", | |
| "checkpoints/lora/detail_enhancer.safetensors", | |
| ] | |
| if lora_weights is None: | |
| lora_weights = [0.4, 0.4] | |
| args = argparse.Namespace( | |
| input_image_path="", # filled in per-request | |
| num_samples=1, | |
| sample_start_idx=0, | |
| sample_id=0, | |
| prompt="", # filled in per-request | |
| prompt_dir=None, | |
| prompt_suffix="", | |
| experiment=experiment, | |
| checkpoint_dir=checkpoint_dir, | |
| output_path="", # filled in per-request | |
| guidance=guidance, | |
| shift=shift, | |
| num_sampling_step=num_sampling_step, | |
| seed=seed, | |
| fps=fps, | |
| num_frames=num_frames, | |
| num_frames_zoom_in=num_frames_zoom_in, | |
| num_frames_zoom_out=num_frames_zoom_out, | |
| resolution=resolution, | |
| context_parallel_size=1, | |
| lora_paths=list(lora_paths), | |
| lora_weights=list(lora_weights), | |
| offload=False, | |
| offload_when_prompt=False, | |
| zoom_in_trajectory=zoom_in_trajectory, | |
| zoom_out_trajectory=zoom_out_trajectory, | |
| zoom_in_direction=zoom_in_direction, | |
| zoom_out_direction=zoom_out_direction, | |
| zoom_in_strength=zoom_in_strength, | |
| zoom_out_strength=zoom_out_strength, | |
| use_moge_scale=use_moge_scale, | |
| ground_plane_align=ground_plane_align, | |
| ground_plane_bottom_frac=ground_plane_bottom_frac, | |
| zoom_out_upward_shift=zoom_out_upward_shift, | |
| zoom_out_upward_ratio=zoom_out_upward_ratio, | |
| depth_backend="da3", | |
| da3_model_name=da3_model_name, | |
| da3_model_path_custom=da3_model_path_custom, | |
| da3_frame_interval=8, | |
| da3_max_history_frames=10, | |
| da3_include_ar_chunk_last_frames=False, | |
| da3_use_predicted_pose=False, | |
| da3_predicted_pose_continuation=False, | |
| use_dmd=use_dmd, | |
| ablate_same_t5=False, | |
| use_dmd_scheduler=False, | |
| warp_chunk_size=None, | |
| num_retrieval_views=1, | |
| disable_cache_update=False, | |
| multiview_ids=None, | |
| offload_da3_diffusion=False, | |
| ) | |
| if args.use_dmd: | |
| args.use_dmd_scheduler = True | |
| args.lora_paths.append(DMD_LORA_PATH) | |
| args.lora_weights.append(DMD_LORA_WEIGHT) | |
| return args | |
| def load_stage1_resources(args: argparse.Namespace) -> Stage1Resources: | |
| """One-time setup for stage 1. Mirrors lyra2_zoomgs_inference.__main__ lines 489–586.""" | |
| t0 = time.monotonic() | |
| log.info("[resident] stage1: loading Lyra-2 model + LoRAs + DA3 (+MoGe)") | |
| misc.set_random_seed(seed=args.seed, by_rank=True) | |
| negative_prompt_data = torch.load( | |
| "checkpoints/text_encoder/negative_prompt.pt", | |
| map_location="cpu", | |
| weights_only=False, | |
| ) | |
| experiment_opts = [ | |
| "model.config.use_mp_policy_fsdp=False", | |
| "model.config.keep_original_net_dtype=False", | |
| ] | |
| if args.lora_paths: | |
| experiment_opts += ["model.config.net.postpone_checkpoint=True"] | |
| model, config = load_model_from_checkpoint( | |
| config_file="lyra_2/_src/configs/config.py", | |
| experiment_name=args.experiment, | |
| checkpoint_path=args.checkpoint_dir, | |
| enable_fsdp=False, | |
| instantiate_ema=False, | |
| load_ema_to_reg=False, | |
| experiment_opts=experiment_opts, | |
| ) | |
| if args.lora_paths: | |
| lora_names = [model.load_lora_weights(p) for p in args.lora_paths] | |
| model.set_weights_and_activate_adapters(lora_names, args.lora_weights) | |
| if hasattr(model, "net") and hasattr(model.net, "enable_selective_checkpoint"): | |
| model.net.enable_selective_checkpoint(model.net.sac_config, model.net.blocks) | |
| desired_dtype = model.tensor_kwargs.get("dtype", None) | |
| desired_device = model.tensor_kwargs.get("device", None) | |
| if desired_dtype is not None: | |
| model.net = model.net.to(device=desired_device, dtype=desired_dtype) | |
| log.info(f"[resident] Casted model.net to dtype={desired_dtype}", rank0_only=True) | |
| assert getattr(model.config, "important_start", True) is True | |
| assert getattr(model.config, "encode_video_from_start", True) is True | |
| assert not getattr(model.config, "use_hd_map_cond", False) | |
| _to_eval_mode(model) | |
| if args.warp_chunk_size is not None: | |
| model.config.warp_chunk_size = args.warp_chunk_size | |
| model.warp_chunk_size = args.warp_chunk_size | |
| target_h, target_w = [int(x) for x in args.resolution.split(",")] | |
| da3_device = model.tensor_kwargs.get( | |
| "device", "cuda" if torch.cuda.is_available() else "cpu" | |
| ) | |
| da3_model = load_da3_model( | |
| da3_model_name=args.da3_model_name, | |
| da3_model_path_custom=args.da3_model_path_custom, | |
| device=da3_device, | |
| ) | |
| _to_eval_mode(da3_model) | |
| moge_model = None | |
| if args.use_moge_scale: | |
| moge_model = load_moge_model(da3_device) | |
| _to_eval_mode(moge_model) | |
| log.info("[resident] MoGe model loaded for depth scale alignment.", rank0_only=True) | |
| log.info( | |
| f"[resident] stage1 load complete in {time.monotonic() - t0:.1f}s " | |
| f"(device={desired_device}, dtype={desired_dtype}, res={target_h}x{target_w})", | |
| rank0_only=True, | |
| ) | |
| return Stage1Resources( | |
| model=model, | |
| config=config, | |
| da3_model=da3_model, | |
| moge_model=moge_model, | |
| negative_prompt_data=negative_prompt_data, | |
| desired_device=desired_device, | |
| desired_dtype=desired_dtype, | |
| target_h=target_h, | |
| target_w=target_w, | |
| base_args=args, | |
| ) | |
| def run_stage1_single( | |
| res: Stage1Resources, | |
| image_path: str | Path, | |
| prompt: str, | |
| preset_params: dict, | |
| output_path: str | Path, | |
| *, | |
| seed: Optional[int] = None, | |
| ) -> Path: | |
| """Produce a zoom-in+zoom-out video for one image + caption. | |
| Returns the path to the combined MP4 (``<output_path>/videos/<stem>.mp4``). | |
| Body mirrors lyra2_zoomgs_inference.__main__ lines 588–799, parameterized | |
| on image_path, prompt, and preset_params (a dict with keys | |
| num_frames_zoom_in, num_frames_zoom_out, zoom_in_strength, zoom_out_strength). | |
| """ | |
| image_path = str(image_path) | |
| if not os.path.isfile(image_path): | |
| raise FileNotFoundError(f"Input image not found: {image_path}") | |
| if not prompt or not prompt.strip(): | |
| raise ValueError("prompt is required") | |
| output_path = Path(output_path) | |
| output_path.mkdir(parents=True, exist_ok=True) | |
| # Per-request args: shallow-copy the base so we can stamp per-request fields. | |
| args = argparse.Namespace(**vars(res.base_args)) | |
| args.input_image_path = image_path | |
| args.prompt = prompt | |
| args.output_path = str(output_path) | |
| args.num_frames_zoom_in = int(preset_params["num_frames_zoom_in"]) | |
| args.num_frames_zoom_out = int(preset_params["num_frames_zoom_out"]) | |
| args.zoom_in_strength = float(preset_params["zoom_in_strength"]) | |
| args.zoom_out_strength = float(preset_params["zoom_out_strength"]) | |
| if seed is not None: | |
| args.seed = int(seed) | |
| model = res.model | |
| da3_model = res.da3_model | |
| moge_model = res.moge_model | |
| desired_device = res.desired_device | |
| desired_dtype = res.desired_dtype | |
| negative_prompt_data = res.negative_prompt_data | |
| target_h, target_w = res.target_h, res.target_w | |
| # _build_image_list supports single path or dir. | |
| all_image_paths = _build_image_list(args.input_image_path) | |
| image_paths = [all_image_paths[0]] | |
| videos_dir = output_path / "videos" | |
| videos_dir.mkdir(parents=True, exist_ok=True) | |
| combined_video_path: Optional[Path] = None | |
| for img_idx, img_path in enumerate(image_paths): | |
| base_name = os.path.splitext(os.path.basename(img_path))[0] | |
| per_image_dir = output_path / base_name | |
| per_image_dir.mkdir(parents=True, exist_ok=True) | |
| combined_video_path = videos_dir / f"{base_name}.mp4" | |
| log.info(f"[resident] stage1: processing {img_path}", rank0_only=True) | |
| misc.set_random_seed(seed=args.seed, by_rank=True) | |
| bgr = cv2.imread(img_path) | |
| if bgr is None: | |
| raise RuntimeError(f"Cannot read image: {img_path}") | |
| rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) | |
| rgb_t = torch.from_numpy(rgb) | |
| log.info("[resident] stage1: running DA3 single-image depth...", rank0_only=True) | |
| image_chw01, depth_hw, K_33, mask_hw = _da3_infer_depth_intrinsics_single( | |
| da3_model=da3_model, | |
| img_rgb_uint8=rgb_t, | |
| target_hw=(target_h, target_w), | |
| ) | |
| if args.use_moge_scale and moge_model is not None: | |
| from lyra_2._src.inference.depth_utils import moge_infer_depth_intrinsics | |
| log.info("[resident] stage1: aligning DA3 depth to MoGe scale...", rank0_only=True) | |
| moge_model.to(desired_device) | |
| with torch.nn.attention.sdpa_kernel([torch.nn.attention.SDPBackend.MATH]): | |
| _, moge_depth_hw, _, moge_mask_hw = moge_infer_depth_intrinsics( | |
| moge_model, | |
| rgb_t, | |
| depth_pred_hw=(target_h, target_w), | |
| target_hw=(target_h, target_w), | |
| ) | |
| da3_d = depth_hw.to(moge_depth_hw.device) | |
| da3_m = mask_hw.to(moge_mask_hw.device) | |
| valid_mask = (da3_m > 0.5) & (moge_mask_hw > 0.5) | |
| if valid_mask.sum() > 10: | |
| d_da3_vals = da3_d[valid_mask] | |
| d_moge_vals = moge_depth_hw[valid_mask] | |
| inv_da3 = 1.0 / (d_da3_vals + 1e-6) | |
| inv_moge = 1.0 / (d_moge_vals + 1e-6) | |
| numerator = (inv_da3 * inv_moge).sum() | |
| denominator = (inv_da3 * inv_da3).sum() | |
| if denominator > 1e-8: | |
| scale = numerator / denominator | |
| log.info( | |
| f"[resident] stage1: depth scale factor = {scale.item()}", | |
| rank0_only=True, | |
| ) | |
| if scale > 1e-6: | |
| depth_hw = depth_hw / scale.to(depth_hw.device) | |
| else: | |
| log.warning( | |
| f"[resident] stage1: scale too small ({scale.item()}), skipping alignment.", | |
| rank0_only=True, | |
| ) | |
| else: | |
| log.warning( | |
| "[resident] stage1: denominator too small for LS scale alignment.", | |
| rank0_only=True, | |
| ) | |
| else: | |
| log.warning( | |
| "[resident] stage1: not enough overlapping valid pixels for scale alignment.", | |
| rank0_only=True, | |
| ) | |
| # Move MoGe back to CPU so it doesn't eat VRAM for the diffusion pass. | |
| moge_model.cpu() | |
| del moge_depth_hw, moge_mask_hw, da3_d, da3_m | |
| torch.cuda.empty_cache() | |
| gc.collect() | |
| img_bchw = image_chw01.to(device=desired_device) * 2.0 - 1.0 | |
| caption = args.prompt | |
| if args.prompt_suffix: | |
| caption = caption.rstrip() + " " + args.prompt_suffix | |
| t5 = get_umt5_embedding(caption, device=desired_device).to(dtype=desired_dtype) | |
| if t5.dim() == 2: | |
| t5 = t5.unsqueeze(0) | |
| elif t5.dim() == 3 and t5.shape[0] != 1: | |
| t5 = t5[:1] | |
| neg_t5 = misc.to(negative_prompt_data["t5_text_embeddings"], **model.tensor_kwargs) | |
| N_in = int(args.num_frames_zoom_in or args.num_frames) | |
| N_out = int(args.num_frames_zoom_out or args.num_frames) | |
| ground_normal = None | |
| if args.ground_plane_align: | |
| ground_normal = _fit_ground_normal_from_depth( | |
| depth_hw, K_33, mask_hw, bottom_frac=args.ground_plane_bottom_frac | |
| ) | |
| log.info( | |
| f"[resident] stage1: ZOOM-IN ({args.zoom_in_trajectory} " | |
| f"{args.zoom_in_direction} str={args.zoom_in_strength}, N={N_in})", | |
| rank0_only=True, | |
| ) | |
| result_in = _generate_one_direction( | |
| model=model, | |
| args=args, | |
| img_bchw=img_bchw, | |
| depth_hw=depth_hw, | |
| mask_hw=mask_hw, | |
| K_33=K_33, | |
| t5_embeddings=t5, | |
| neg_t5_embeddings=neg_t5, | |
| trajectory=args.zoom_in_trajectory, | |
| direction=args.zoom_in_direction, | |
| strength=args.zoom_in_strength, | |
| N=N_in, | |
| da3_model=da3_model, | |
| process_group=None, | |
| log_prefix=f"{base_name}_zoom_in", | |
| ground_normal_cam=ground_normal, | |
| ) | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| log.info( | |
| f"[resident] stage1: ZOOM-OUT ({args.zoom_out_trajectory} " | |
| f"{args.zoom_out_direction} str={args.zoom_out_strength}, N={N_out})", | |
| rank0_only=True, | |
| ) | |
| result_out = _generate_one_direction( | |
| model=model, | |
| args=args, | |
| img_bchw=img_bchw, | |
| depth_hw=depth_hw, | |
| mask_hw=mask_hw, | |
| K_33=K_33, | |
| t5_embeddings=t5, | |
| neg_t5_embeddings=neg_t5, | |
| trajectory=args.zoom_out_trajectory, | |
| direction=args.zoom_out_direction, | |
| strength=args.zoom_out_strength, | |
| N=N_out, | |
| da3_model=da3_model, | |
| process_group=None, | |
| log_prefix=f"{base_name}_zoom_out", | |
| upward_shift=args.zoom_out_upward_shift, | |
| ground_normal_cam=ground_normal, | |
| zoom_out_upward_ratio=args.zoom_out_upward_ratio, | |
| ) | |
| if result_in is None and result_out is None: | |
| raise RuntimeError(f"Both zoom-in and zoom-out failed for {img_path}") | |
| for tag, res_v in [("zoom_in", result_in), ("zoom_out", result_out)]: | |
| if res_v is None: | |
| continue | |
| vid_stem = str(per_image_dir / tag) | |
| to_show = [] | |
| if res_v.get("warp_video") is not None: | |
| to_show.append(res_v["warp_video"]) | |
| to_show.append(res_v["video"]) | |
| save_output(to_show, vid_stem + ".mp4") | |
| videos_to_combine = [] | |
| if result_out is not None: | |
| videos_to_combine.append(result_out["video"].flip(dims=[2])) | |
| if result_in is not None: | |
| videos_to_combine.append(result_in["video"]) | |
| combined_video = torch.cat(videos_to_combine, dim=2) | |
| combined_01 = (combined_video[0].clamp(-1, 1) * 0.5 + 0.5).float().cpu() | |
| save_img_or_video(combined_01, str(combined_video_path).replace(".mp4", ""), fps=args.fps) | |
| save_img_or_video(combined_01, str(per_image_dir / "combined"), fps=args.fps) | |
| del combined_video, combined_01 | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| gc.collect() | |
| if combined_video_path is None or not combined_video_path.exists(): | |
| raise RuntimeError( | |
| f"stage1 did not produce expected video at {combined_video_path}" | |
| ) | |
| return combined_video_path | |
| # --------------------------------------------------------------------------- | |
| # Stage 2 | |
| # --------------------------------------------------------------------------- | |
| def load_stage2_resources( | |
| *, | |
| da3_from_stage1: Optional[Any] = None, | |
| device: Optional[torch.device | str] = None, | |
| da3_model_name: str = "depth-anything/DA3NESTED-GIANT-LARGE-1.1", | |
| da3_model_path_custom: str = "checkpoints/recon/model.pt", | |
| ) -> Stage2Resources: | |
| """One-time setup for stage 2: reuse or load DA3, import VIPE class. | |
| When stage 1 already holds a DA3 model we pass it in via ``da3_from_stage1`` | |
| to save ~1 GB VRAM and a couple minutes of cold load. | |
| """ | |
| t0 = time.monotonic() | |
| dev = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu")) | |
| if da3_from_stage1 is not None: | |
| log.info("[resident] stage2: reusing DA3 model from stage 1", rank0_only=True) | |
| da3_model = da3_from_stage1 | |
| else: | |
| log.info("[resident] stage2: loading DA3 model", rank0_only=True) | |
| custom = None | |
| if da3_model_path_custom: | |
| custom = str(Path(da3_model_path_custom).expanduser().resolve()) | |
| if not Path(custom).is_file(): | |
| raise FileNotFoundError(f"DA3 checkpoint not found: {custom}") | |
| da3_model = load_da3_model( | |
| da3_model_name=da3_model_name, | |
| da3_model_path_custom=custom, | |
| device=str(dev), | |
| ) | |
| _to_eval_mode(da3_model) | |
| log.info("[resident] stage2: importing VIPE class", rank0_only=True) | |
| VIPE = _import_vipe_class() | |
| log.info( | |
| f"[resident] stage2 load complete in {time.monotonic() - t0:.1f}s", | |
| rank0_only=True, | |
| ) | |
| return Stage2Resources(da3_model=da3_model, VIPE=VIPE, device=dev) | |
| def run_stage2_single( | |
| res: Stage2Resources, | |
| video_path: str | Path, | |
| output_dir: str | Path, | |
| *, | |
| da3_max_frames: int = 128, | |
| use_da3_render_pose: bool = True, | |
| max_frames: int = 0, | |
| max_resolution: int = 0, | |
| da3_process_res: Optional[int] = None, | |
| da3_process_method: str = "upper_bound_resize", | |
| gs_down_ratio: int = 2, | |
| gs_scale_extra_multiplier: float = 1.0, | |
| gs_ply_prune_opacity_percentile: Optional[float] = None, | |
| gs_ds_feature_mode: bool = True, | |
| render_fps: Optional[float] = None, | |
| render_chunk_size: int = 1, | |
| vipe_full_mode: bool = False, | |
| vipe_overrides: Optional[Sequence[str]] = None, | |
| force: bool = False, | |
| ) -> Path: | |
| """Mirrors vipe_da3_gs_recon.main() for the skip_vipe=False path only. | |
| Returns the path to ``reconstructed_scene.ply`` inside ``output_dir``. | |
| """ | |
| input_video = Path(video_path).expanduser().resolve() | |
| if not input_video.is_file(): | |
| raise FileNotFoundError(f"Input video not found: {input_video}") | |
| output_dir = Path(output_dir).expanduser().resolve() | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| done_marker = output_dir / ".done" | |
| if done_marker.is_file() and not force: | |
| existing = output_dir / "reconstructed_scene.ply" | |
| if existing.is_file(): | |
| log.info( | |
| f"[resident] stage2: output already present at {existing}; " | |
| f"reusing (pass force=True to re-run)", | |
| rank0_only=True, | |
| ) | |
| return existing | |
| da3_model = res.da3_model | |
| VIPE = res.VIPE | |
| device = res.device | |
| log.info(f"[resident] stage2: input_video={input_video}", rank0_only=True) | |
| log.info(f"[resident] stage2: output_dir={output_dir}", rank0_only=True) | |
| images_all, indices_all, fps = _collect_vipe_images( | |
| str(input_video), | |
| vipe_stride=1, | |
| max_frames=max_frames, | |
| max_views=0, | |
| ) | |
| if not images_all: | |
| raise RuntimeError("No frames read from video.") | |
| indices_da3_rel = _uniform_subsample_indices(len(images_all), da3_max_frames) | |
| if not indices_da3_rel: | |
| raise RuntimeError("No frames selected for DA3.") | |
| images_da3 = [images_all[idx] for idx in indices_da3_rel] | |
| indices_da3 = [indices_all[idx] for idx in indices_da3_rel] | |
| eff_fps = float(fps) | |
| log.info( | |
| f"[resident] stage2: fps={fps:.4g}, da3_max_frames={da3_max_frames}, " | |
| f"frames_all={len(images_all)}, frames_da3={len(images_da3)}", | |
| rank0_only=True, | |
| ) | |
| frames_np = np.stack(images_all, axis=0).astype(np.float32) / 255.0 | |
| frames_thwc = torch.from_numpy(frames_np).contiguous() | |
| with tempfile.TemporaryDirectory(prefix="vipe_da3_gs_") as tmpdir: | |
| vipe_output_path = Path(tmpdir) / "vipe_out" | |
| vipe_output_path.mkdir(parents=True, exist_ok=True) | |
| overrides = ( | |
| list(vipe_overrides) | |
| if vipe_overrides is not None | |
| else _vipe_default_overrides(vipe_output_path) | |
| ) | |
| log.info("[resident] stage2: instantiating VIPE...", rank0_only=True) | |
| vipe = VIPE(overrides, fast_mode=not vipe_full_mode) | |
| log.info("[resident] stage2: running VIPE...", rank0_only=True) | |
| vipe_out = vipe.infer_frames(frames_thwc, fps=eff_fps, name=input_video.stem) | |
| c2w = vipe_out.extrinsics_c2w.to(dtype=torch.float32) | |
| w2c = torch.linalg.inv(c2w) | |
| intrinsics_vipe = _intrinsics_vec_to_k33(vipe_out.intrinsics.to(dtype=torch.float32)) | |
| w2c_np_vipe_full = w2c.cpu().numpy().astype(np.float32) | |
| k_np_vipe_full = intrinsics_vipe.cpu().numpy().astype(np.float32) | |
| w2c_np_da3 = w2c_np_vipe_full[indices_da3_rel] | |
| k_np_da3 = k_np_vipe_full[indices_da3_rel] | |
| np.savez( | |
| output_dir / "vipe_predictions.npz", | |
| frame_ids=vipe_out.frame_ids.cpu().numpy().astype(np.int64), | |
| w2c_vipe=w2c_np_vipe_full, | |
| intrinsics_vipe=k_np_vipe_full, | |
| w2c_da3=w2c_np_da3, | |
| intrinsics_da3=k_np_da3, | |
| indices_vipe=np.asarray(indices_all, dtype=np.int64), | |
| indices_da3=np.asarray(indices_da3, dtype=np.int64), | |
| fps=np.asarray([eff_fps], dtype=np.float32), | |
| input_video_path=np.asarray([str(input_video)]), | |
| ) | |
| # Release VIPE's heavy sub-models before DA3 pass-2 runs. | |
| del vipe | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| if da3_process_res is not None: | |
| proc_res = int(da3_process_res) | |
| proc_method = str(da3_process_method) | |
| elif int(max_resolution) > 0: | |
| proc_res = int(max_resolution) | |
| proc_method = "lower_bound_resize" | |
| else: | |
| h0, w0 = images_da3[0].shape[:2] | |
| proc_res = int(max(h0, w0)) | |
| proc_method = "upper_bound_resize" | |
| _ensure_da3_on_syspath() | |
| from depth_anything_3.utils.gsply_helpers import save_gaussian_ply # type: ignore | |
| log.info( | |
| f"[resident] stage2: DA3 GS recon: views={len(images_da3)} process_res={proc_res}", | |
| rank0_only=True, | |
| ) | |
| pred = da3_model.inference( | |
| image=images_da3, | |
| extrinsics=w2c_np_da3, | |
| intrinsics=k_np_da3, | |
| align_to_input_extrinsics=False, | |
| align_to_input_ext_scale=False, | |
| infer_gs=True, | |
| process_res=proc_res, | |
| process_res_method=proc_method, | |
| export_dir=None, | |
| export_format="mini_npz", | |
| use_aligned_pred_cam=True, | |
| gs_down_ratio=gs_down_ratio, | |
| gs_scale_extra_multiplier=gs_scale_extra_multiplier, | |
| gs_ds_feature_mode=gs_ds_feature_mode, | |
| ) | |
| aligned_w2c_da3 = None | |
| if use_da3_render_pose and pred.extrinsics is not None: | |
| aligned_w2c_da3 = _compute_aligned_pred_w2c( | |
| np.asarray(pred.extrinsics, dtype=np.float32), | |
| w2c_np_da3, | |
| ) | |
| final_ply_path = output_dir / "reconstructed_scene.ply" | |
| depth_t = torch.from_numpy(np.asarray(pred.depth, dtype=np.float32)).float() | |
| save_gaussian_ply( | |
| pred.gaussians, | |
| str(final_ply_path), | |
| ctx_depth=depth_t.unsqueeze(-1), | |
| prune_by_opacity_percentile=gs_ply_prune_opacity_percentile, | |
| prune_border_gs=False | |
| if ( | |
| gs_ply_prune_opacity_percentile is not None | |
| and gs_ply_prune_opacity_percentile > 0 | |
| ) | |
| else True, | |
| ) | |
| log.info(f"[resident] stage2: saved PLY to {final_ply_path}", rank0_only=True) | |
| if use_da3_render_pose and aligned_w2c_da3 is not None: | |
| w2c_render = _interpolate_w2c(aligned_w2c_da3, indices_da3_rel, len(images_all)) | |
| k_render = k_np_vipe_full | |
| log.info( | |
| "[resident] stage2: rendering with DA3-aligned poses.", rank0_only=True | |
| ) | |
| else: | |
| w2c_render = w2c_np_vipe_full | |
| k_render = k_np_vipe_full | |
| log.info("[resident] stage2: rendering with VIPE poses.", rank0_only=True) | |
| np.savez( | |
| output_dir / "cameras.npz", | |
| w2c_render=w2c_render, | |
| indices_da3=np.asarray(indices_da3, dtype=np.int64), | |
| fps=np.asarray([eff_fps], dtype=np.float32), | |
| no_vipe=np.asarray([0], dtype=np.int32), | |
| w2c_vipe=w2c_np_vipe_full, | |
| intrinsics_vipe=k_np_vipe_full, | |
| w2c_da3=w2c_np_da3, | |
| intrinsics_da3=k_np_da3, | |
| indices_vipe=np.asarray(indices_all, dtype=np.int64), | |
| use_da3_render_pose=np.asarray([int(use_da3_render_pose)], dtype=np.int32), | |
| ) | |
| del pred | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| from depth_anything_3.model.utils.gs_renderer import ( # type: ignore | |
| run_renderer_in_chunk_w_trj_mode, | |
| ) | |
| gs_device = device | |
| if hasattr(da3_model, "model"): | |
| try: | |
| gs_device = next(da3_model.model.parameters()).device | |
| except StopIteration: | |
| gs_device = device | |
| gaussians = _load_gaussian_ply_to_gaussians(str(final_ply_path), device=gs_device) | |
| render_extr = torch.from_numpy(w2c_render).to( | |
| device=gs_device, dtype=gaussians.means.dtype | |
| )[None] | |
| render_intr = torch.from_numpy(k_render).to( | |
| device=gs_device, dtype=gaussians.means.dtype | |
| )[None] | |
| if render_extr.shape[-2:] == (3, 4): | |
| pad = torch.tensor( | |
| [0, 0, 0, 1], device=gs_device, dtype=gaussians.means.dtype | |
| ).view(1, 1, 1, 4) | |
| render_extr = torch.cat( | |
| [render_extr, pad.expand(render_extr.shape[0], render_extr.shape[1], -1, -1)], | |
| dim=-2, | |
| ) | |
| render_h, render_w = images_all[0].shape[:2] | |
| eff_render_fps = ( | |
| float(render_fps) | |
| if render_fps is not None | |
| else float(max(1, round(eff_fps))) | |
| ) | |
| log.info( | |
| f"[resident] stage2: rendering {render_extr.shape[1]} frames at " | |
| f"{render_h}x{render_w} (fps={eff_render_fps:.2f})", | |
| rank0_only=True, | |
| ) | |
| color, depth = run_renderer_in_chunk_w_trj_mode( | |
| gaussians=gaussians, | |
| extrinsics=render_extr, | |
| intrinsics=render_intr, | |
| image_shape=(render_h, render_w), | |
| chunk_size=int(render_chunk_size), | |
| trj_mode="original", | |
| use_sh=True, | |
| color_mode="RGB+ED", | |
| enable_tqdm=True, | |
| ) | |
| frames_render = ( | |
| color[0].clamp(0.0, 1.0).mul(255.0).byte().permute(0, 2, 3, 1).cpu().numpy() | |
| ) | |
| render_video_path = output_dir / "gs_trajectory.mp4" | |
| _save_video_mp4(str(render_video_path), frames_render, fps=eff_render_fps) | |
| del gaussians, render_extr, render_intr, color, depth, frames_render | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| done_marker.write_text("done\n") | |
| log.info("[resident] stage2: done.", rank0_only=True) | |
| return final_ply_path | |