lyra2-explorable-scene / resident_inference.py
cezar-hapiko's picture
Refactor: load Lyra-2 model once at startup, reuse across requests
328801b
Raw History Blame Contribute Delete
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
# ---------------------------------------------------------------------------
@dataclass
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
@dataclass
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