Meridian / recam /geometry.py
yycc's picture
Meridian
9f57754
Raw History Blame Contribute Delete
7.99 kB
# Copyright 2026 Viggle AI. Licensed under the Apache License, Version 2.0 (see LICENSE-CODE).
# SPDX-License-Identifier: Apache-2.0
"""Geometry: one solo VGGT-Omega pass per clip, then the clip's point cloud rendered from any camera.
frames [F,H,W,3] u8 --letterbox--> [F,FULL,FULL,3] --to_input--> [F,3,RES,RES] --reconstruct--> S
S + an authored camera path --warp--> [F,h,w,3] u8 render with mid-grey holes, + coverage
`S` is per-frame depth / confidence-pruned `keep` / extrinsics / intrinsics in VGGT's 512 px frame; every
frame is reconstructed on its own (the model sees the whole clip, but no target footage exists at
inference), and cameras are authored in the *source's* gauge: the source camera of the window's first
frame is the world origin and one unit is the median depth around the picture centre (`gauge`).
"""
import os
import sys
import torch
import torch.nn.functional as F
import torchvision.transforms.v2.functional as TF
FULL = 1280 # side of the square the source is letterboxed into (the training corpus's native side)
RES = 512 # VGGT input side
CONF_PCT, EDGE_RTOL = 2.0, 0.30 # drop the 2 % least confident points and every 3x3 depth edge wider than 30 %
HOLE = 128 # uncovered pixels: mid grey. There is no mask channel; the prompt says grey is a hole.
NUM_FRAMES = 73
LENGTHS = (73, 90, 107, 124, 141, 158, 175, 243) # 17k+5, the lengths `assets/` has a prompt embed for
def vggt(ckpt=None, repo=None, device="cuda"):
"""Load VGGT-Omega (1B, 512). Not redistributed here -- see README: clone facebookresearch/vggt-omega and
download `vggt_omega_1b_512.pt` from the gated `facebook/VGGT-Omega`. Defaults: `$VGGT_OMEGA_DIR`,
`$VGGT_OMEGA_CKPT`."""
repo = repo or os.environ.get("VGGT_OMEGA_DIR")
ckpt = ckpt or os.environ.get("VGGT_OMEGA_CKPT") or (repo and f"{repo}/checkpoints/vggt_omega_1b_512.pt")
if not ckpt:
raise RuntimeError("VGGT-Omega not found: set $VGGT_OMEGA_DIR (and $VGGT_OMEGA_CKPT if the .pt lives elsewhere) "
"or pass --vggt-repo / --vggt; see README, Install")
if repo:
sys.path.insert(0, repo)
from vggt_omega.models import VGGTOmega
model = VGGTOmega().eval()
model.load_state_dict(torch.load(ckpt, map_location="cpu"))
return model.to(device)
def reconstruct(model, imgs):
"""[N,3,512,512] in [0,1] -> dict(extr, intr, depth, keep, c2w). `keep` drops depth edges and low confidence."""
from vggt_omega.utils.pose_enc import encoding_to_camera
with torch.inference_mode():
pred = model(imgs)
extr, intr = encoding_to_camera(pred["pose_enc"], (RES, RES))
depth = pred["depth"][0, ..., 0].float()
conf = pred["depth_conf"][0].float().clone()
mx = F.max_pool2d(depth[None], 3, 1, 1)[0]
mn = -F.max_pool2d(-depth[None], 3, 1, 1)[0]
conf[(mx - mn) / depth.abs().clamp(min=1e-6) > EDGE_RTOL] = 0.0
keep = torch.isfinite(depth) & torch.isfinite(conf) & (conf > 1e-5)
q = torch.stack([torch.quantile(conf[i][keep[i]], CONF_PCT / 100) if keep[i].any()
else conf.new_zeros(()) for i in range(len(conf))])
keep &= conf >= q[:, None, None]
w2c = torch.eye(4, device=depth.device).repeat(len(depth), 1, 1)
w2c[:, :3, :4] = extr[0].float()
return dict(extr=extr[0].float().clone(), intr=intr[0].float().clone(), depth=depth.clone(),
keep=keep.clone(), c2w=torch.linalg.inv(w2c))
def unproject(depth, extr, intr):
n, h, w = depth.shape
yy, xx = torch.meshgrid(torch.arange(h, device=depth.device, dtype=torch.float32),
torch.arange(w, device=depth.device, dtype=torch.float32), indexing="ij")
fx, fy = intr[:, 0, 0][:, None, None], intr[:, 1, 1][:, None, None]
cx, cy = intr[:, 0, 2][:, None, None], intr[:, 1, 2][:, None, None]
cam = torch.stack([(xx - cx) / fx * depth, (yy - cy) / fy * depth, depth], -1)
return torch.einsum("sij,shwj->shwi", extr[:, :3, :3].transpose(1, 2),
cam - extr[:, :3, 3][:, None, None, :])
def upsample(depth, keep, n=FULL):
"""512 depth/keep -> n x n. Bilinear depth; a hi-res pixel survives only where all parents did,
which kills bilinear's flying pixels across a depth discontinuity."""
d = F.interpolate(depth[:, None], size=(n, n), mode="bilinear", align_corners=False)[:, 0]
k = F.interpolate(keep[:, None].float(), size=(n, n), mode="bilinear", align_corners=False)[:, 0] > 0.999
return d, k
def scale_k(K, f):
"""`K` for an image resized by `f` under the pixel-centre convention every resampler here uses
(`x_hi = f*(x_lo + 0.5) - 0.5`), so the principal point picks up `0.5*(f-1)` on top of the scaling."""
out = K * f
out[..., 2, 2] = 1.0
out[..., :2, 2] += 0.5 * (f - 1)
return out
def canvas_k(K, box, f):
"""`K` (VGGT 512-space, `[...,3,3]`) -> the crop box `box` resized by `f`, i.e. canvas pixels."""
x0, y0 = box[:2]
out = scale_k(K, FULL / RES)
out[..., 0, 2] -= x0
out[..., 1, 2] -= y0
return scale_k(out, f)
def render_hw(P, C, extr, intr, H, W, splat=1):
"""Z-buffer point splat (3x3 per point). Returns the image and the covered pixels."""
cam = P @ extr[:3, :3].T + extr[:3, 3]
m = cam[:, 2] > 1e-6
cam, C = cam[m], C[m]
z = cam[:, 2]
u = cam[:, 0] / z * intr[0, 0] + intr[0, 2]
v = cam[:, 1] / z * intr[1, 1] + intr[1, 2]
offs = [(dx, dy) for dy in (-splat, 0, splat) for dx in (-splat, 0, splat)]
x = torch.cat([(u + dx).round() for dx, _ in offs]).long()
y = torch.cat([(v + dy).round() for _, dy in offs]).long()
z = z.repeat(len(offs))
C = C.repeat(len(offs), 1)
k = (x >= 0) & (x < W) & (y >= 0) & (y < H)
idx, z, C = (y[k] * W + x[k]), z[k], C[k]
zbuf = torch.full((H * W,), float("inf"), device=P.device)
zbuf.scatter_reduce_(0, idx, z, "amin", include_self=True)
win = z == zbuf[idx]
img = torch.full((H * W, 3), HOLE, dtype=torch.uint8, device=P.device)
img[idx[win]] = C[win]
cov = torch.zeros(H * W, dtype=torch.bool, device=P.device)
cov[idx] = True
return img.view(H, W, 3), cov.view(H, W)
def warp(S, w2c, intr_t, src_full, box, canvas):
"""The source clip's per-frame point cloud, seen from the target camera `w2c [F,4,4]` with intrinsics
`intr_t [F,3,3]` (VGGT 512-space). `box = (x0, y0, w, h, f)` is the crop box in the letterboxed frame and
its resize factor to `canvas = (w, h)`. One frame at a time: a 1280x1280 unproject is 19 MB of points."""
x0, y0, cw, ch, f = box
w, h = canvas
imgs, covs = [], []
for ti in range(len(src_full)):
n = src_full.shape[1]
d, k = upsample(S["depth"][ti : ti + 1], S["keep"][ti : ti + 1], n)
Kf = scale_k(S["intr"][ti], FULL / RES)
P = unproject(d, S["extr"][ti : ti + 1], Kf[None])[0][y0 : y0 + ch, x0 : x0 + cw]
kb = k[0, y0 : y0 + ch, x0 : x0 + cw]
Kc = canvas_k(intr_t[ti], box, f)
img, cov = render_hw(P[kb], src_full[ti, y0 : y0 + ch, x0 : x0 + cw][kb], w2c[ti, :3], Kc, h, w)
imgs.append(img)
covs.append(cov)
return torch.stack(imgs), torch.stack(covs)
def resize_u8(x, size, mode=TF.InterpolationMode.BILINEAR, chunk=16):
"""[F,H,W,3] uint8 -> [F,size[0],size[1],3] uint8, antialiased, chunked to bound memory."""
out = []
for i in range(0, len(x), chunk):
y = TF.resize(x[i : i + chunk].permute(0, 3, 1, 2).float(), list(size), mode, antialias=True)
out.append(y.clamp_(0, 255).round_().to(torch.uint8).permute(0, 2, 3, 1))
return torch.cat(out)
def to_input(x, chunk=16):
"""[F,FULL,FULL,3] uint8 -> [F,3,RES,RES] float in [0,1]. Bicubic, matching VGGT's own preprocessing."""
out = []
for i in range(0, len(x), chunk):
y = x[i : i + chunk].permute(0, 3, 1, 2).float().div(255)
out.append(TF.resize(y, [RES, RES], TF.InterpolationMode.BICUBIC, antialias=True).clamp_(0, 1))
return torch.cat(out)