mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
452 lines
18 KiB
Python
452 lines
18 KiB
Python
#!/usr/bin/env python3
|
|
|
|
import torch
|
|
import torch.nn.functional as F
|
|
import numpy as np
|
|
from PIL import Image
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Optional
|
|
from pathlib import Path
|
|
|
|
import os
|
|
from tqdm import tqdm
|
|
import json
|
|
|
|
from spirulae_splat.splat.cuda import (
|
|
undistort_image,
|
|
distort_image,
|
|
warp_image_wide_to_pinhole,
|
|
warp_image_pinhole_to_wide,
|
|
warp_depth_pinhole_to_wide,
|
|
warp_points_pinhole_to_wide,
|
|
depth_to_normal,
|
|
depth_normal_loss,
|
|
depth_to_points,
|
|
)
|
|
|
|
from typing import Tuple, Literal
|
|
|
|
from io import StringIO
|
|
from contextlib import redirect_stdout
|
|
|
|
|
|
@dataclass
|
|
class Config:
|
|
dataset_dir: str
|
|
max_size: int = 1600 # result not always better at high res
|
|
|
|
# Overrides for the per-image geometry decisions. None => decide
|
|
# automatically (legacy behavior). Set from the CLI flags.
|
|
# warp_to_pinhole: whether to split a wide fisheye into several pinhole
|
|
# faces before running the model (vs. undistorting to a single
|
|
# pinhole). Auto: only for fisheye that would lose too much FOV under
|
|
# a single undistortion, and only for metric3d.
|
|
# ray_depth: whether the saved depth stores ray depth (Euclidean distance
|
|
# along the camera ray) rather than linear (z) depth. Auto: ray depth
|
|
# exactly when the wide->pinhole split is used. Must match the
|
|
# trainer's `input_depth_is_ray_depth`.
|
|
warp_to_pinhole: Optional[bool] = None
|
|
ray_depth: Optional[bool] = None
|
|
|
|
|
|
def expand_white_area(boolean_image, offset):
|
|
"""
|
|
Expands the white area in a boolean image by a given offset using dilation via max pooling.
|
|
|
|
Args:
|
|
boolean_image (torch.Tensor): Input boolean image tensor (C, H, W or B, C, H, W).
|
|
White pixels should be True/1 and black pixels False/0.
|
|
offset (int): The number of pixels to expand the white areas by.
|
|
This will result in a kernel size of (2 * offset + 1).
|
|
|
|
Returns:
|
|
torch.Tensor: The image with expanded white areas.
|
|
"""
|
|
if offset <= 0:
|
|
return boolean_image
|
|
|
|
# Ensure the input is a float tensor with a batch dimension (B, C, H, W) for max_pool2d
|
|
if boolean_image.dim() == 3:
|
|
image_with_batch = boolean_image.unsqueeze(0)
|
|
else:
|
|
image_with_batch = boolean_image.clone()
|
|
|
|
# Convert to float for pooling operation
|
|
image_float = image_with_batch.float()
|
|
|
|
# Kernel size for dilation is (2 * offset + 1)
|
|
kernel_size = 2 * offset + 1
|
|
|
|
# Use max pooling with appropriate padding to implement dilation
|
|
# A pixel becomes "white" if any pixel in its neighborhood (defined by the kernel) was white
|
|
expanded_image_float = F.max_pool2d(
|
|
image_float,
|
|
kernel_size=kernel_size,
|
|
stride=1,
|
|
padding=offset
|
|
)
|
|
|
|
# Convert the result back to the original boolean type
|
|
# Values > 0 will become True/1, and 0 will become False/0
|
|
expanded_image_bool = expanded_image_float > 0.5
|
|
|
|
if boolean_image.dim() == 3:
|
|
return expanded_image_bool.squeeze(0)
|
|
else:
|
|
return expanded_image_bool
|
|
|
|
|
|
def linear_to_ray_depth(depth: torch.Tensor, intrins_fxfycxcy) -> torch.Tensor:
|
|
"""Convert linear (z) depth to ray depth in a pinhole frame.
|
|
|
|
ray = z * sqrt(x_n^2 + y_n^2 + 1), with x_n=(px+0.5-cx)/fx, y_n=(py+0.5-cy)/fy.
|
|
`depth` is [B, H, W, 1]; the intrinsics are for its own resolution.
|
|
"""
|
|
fx, fy, cx, cy = intrins_fxfycxcy
|
|
_, H, W, _ = depth.shape
|
|
ys, xs = torch.meshgrid(
|
|
torch.arange(H, device=depth.device, dtype=torch.float32),
|
|
torch.arange(W, device=depth.device, dtype=torch.float32),
|
|
indexing='ij',
|
|
)
|
|
xn = (xs + 0.5 - cx) / fx
|
|
yn = (ys + 0.5 - cy) / fy
|
|
sec = torch.sqrt(xn * xn + yn * yn + 1.0)
|
|
return depth * sec[None, :, :, None]
|
|
|
|
|
|
def process_image(
|
|
model_type: Literal['metric3d', 'da3', 'da3+lang-sam', 'metric3d+da3+lang-sam'],
|
|
models: Tuple,
|
|
intrins: dict,
|
|
image_path: str,
|
|
depth_save_path: Optional[str]=None,
|
|
normal_save_path: Optional[str]=None
|
|
):
|
|
if (depth_save_path is None or os.path.exists(depth_save_path)) and \
|
|
(normal_save_path is None or os.path.exists(normal_save_path)):
|
|
return
|
|
|
|
image = Image.open(image_path).convert("RGB")
|
|
w, h = image.size
|
|
sc = Config.max_size / max(w, h)
|
|
if sc < 1.0:
|
|
w, h = int(sc*w+0.5), int(sc*h+0.5)
|
|
image = image.resize((w, h))
|
|
|
|
image_original = image
|
|
image = torch.from_numpy(np.array(image))[None].float().cuda() / 255.0
|
|
|
|
sw = image.shape[2] / intrins['w']
|
|
sh = image.shape[1] / intrins['h']
|
|
intrins = (
|
|
intrins['camera_model'],
|
|
(intrins['fl_x']*sw, intrins['fl_y']*sh, intrins['cx']*sw, intrins['cy']*sh),
|
|
tuple(intrins.get(key, 0.0) for key in "k1 k2 k3 k4 p1 p2 sx1 sy1 b1 b2".split())
|
|
)
|
|
# is_ray_depth = (intrins[0].lower() == "fisheye")
|
|
|
|
# Decide whether to split a wide fisheye into several pinhole faces
|
|
# ("warp to pinhole") vs. undistorting to a single pinhole. Auto (the
|
|
# legacy heuristic): only for a fisheye that would lose too much FOV under
|
|
# a single undistortion, and only for metric3d. Config.warp_to_pinhole
|
|
# overrides the decision when not None.
|
|
if Config.warp_to_pinhole is None:
|
|
do_warp = (
|
|
intrins[0].lower() == "fisheye"
|
|
and model_type == "metric3d" # TODO: add support for other models
|
|
and distort_image(torch.ones_like(image[..., :1]), *intrins).mean().item() <= 0.75
|
|
)
|
|
else:
|
|
do_warp = Config.warp_to_pinhole
|
|
|
|
# Whether the saved depth stores ray depth (Euclidean distance along the
|
|
# camera ray) rather than linear (z) depth. Auto: ray depth exactly when we
|
|
# split into pinhole faces (each face keeps the wide capture's ray
|
|
# semantics); linear when we undistort to a single pinhole. Must match the
|
|
# trainer's `input_depth_is_ray_depth`.
|
|
is_ray_depth = do_warp if Config.ray_depth is None else Config.ray_depth
|
|
|
|
axes = None
|
|
if not do_warp:
|
|
image = undistort_image(image, *intrins)
|
|
else:
|
|
# split into multiple images for very wide fisheye, to minimize loss of fov
|
|
r2, r3, r6 = 2**0.5, 3**0.5, 6**0.5
|
|
a0 = [1/r2, 1/r6, 1/r3]
|
|
a1 = [-1/r2, 1/r6, 1/r3]
|
|
a2 = [0, -2/r6, 1/r3]
|
|
axes = torch.Tensor([
|
|
# [a0, a1, a2],
|
|
# [a1, a2, a0],
|
|
# [a2, a0, a1],
|
|
[[1,0,0],[0,1,0],[0,0,1]],
|
|
[[0,1,0],[0,0,1],[1,0,0]],
|
|
[[-1,0,0],[0,0,1],[0,1,0]],
|
|
[[0,-1,0],[0,0,1],[-1,0,0]],
|
|
[[1,0,0],[0,0,1],[0,-1,0]],
|
|
]).float().cuda()
|
|
axes[:, 0:2] *= 1.2
|
|
|
|
i2 = 0.5**0.5
|
|
a = 1.27 # in radians
|
|
sa, ca = np.sin(a), np.cos(a)
|
|
axes = torch.Tensor([
|
|
[[1,0,0],[0,1,0],[0,0,1]],
|
|
[[0,1,0],[-ca,0,sa],[sa,0,ca]],
|
|
[[-1,0,0],[0,-ca,sa],[0,sa,ca]],
|
|
[[0,-1,0],[ca,0,sa],[-sa,0,ca]],
|
|
[[1,0,0],[0,ca,sa],[0,-sa,ca]],
|
|
# [[-i2,i2,0],[-i2*ca,-i2*ca,sa],[i2*sa,i2*sa,ca]],
|
|
# [[-i2,-i2,0],[i2*ca,-i2*ca,sa],[-i2*sa,i2*sa,ca]],
|
|
# [[i2,-i2,0],[i2*ca,i2*ca,sa],[-i2*sa,-i2*sa,ca]],
|
|
# [[i2,i2,0],[-i2*ca,i2*ca,sa],[i2*sa,-i2*sa,ca]],
|
|
]).float().cuda()
|
|
axes[:, 0:2] *= 1.25
|
|
|
|
if axes is not None:
|
|
original_shape = image.shape
|
|
target_size = int(np.sqrt(original_shape[1]*original_shape[2]/len(axes)))
|
|
image = warp_image_wide_to_pinhole(image, *intrins, axes, target_size, target_size)[0]
|
|
|
|
# pred_depth and pred_normal are distorted, pred_sky is original
|
|
pred_depth, pred_normal, pred_sky = None, None, None
|
|
|
|
if "metric3d" in model_type:
|
|
model = models[model_type.split('+').index("metric3d")]
|
|
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
|
pred_depth, _, output_dict = model.inference({'input': image.permute((0, 3, 1, 2))})
|
|
pred_normal = output_dict['prediction_normal'][:, :3]
|
|
|
|
if axes is not None:
|
|
pred_depth = warp_depth_pinhole_to_wide(pred_depth.permute(0, 2, 3, 1)[None], *intrins, axes, *original_shape[1:3], is_ray_depth=is_ray_depth)
|
|
pred_normal = warp_points_pinhole_to_wide(pred_normal.permute(0, 2, 3, 1)[None], *intrins, axes, *original_shape[1:3])
|
|
pred_depth = pred_depth.permute(0, 3, 1, 2)
|
|
pred_normal = pred_normal.permute(0, 3, 1, 2)
|
|
|
|
pred_depth = torch.nn.functional.interpolate(
|
|
pred_depth, size=(h, w), mode='bilinear', align_corners=False
|
|
).permute(0, 2, 3, 1) # [h, w, 1]
|
|
pred_normal = torch.nn.functional.interpolate(
|
|
pred_normal, size=(h, w), mode='bilinear', align_corners=False
|
|
)[:, :3].permute(0, 2, 3, 1) # [h, w, 3]
|
|
|
|
pad = Config.max_size // 200 + 1 # fix bad values in case undistortion gives black border that messes up model
|
|
sky_depth = torch.amax(pred_depth[0, pad:-pad, pad:-pad]).item()
|
|
pred_depth = pred_depth.clip(max=sky_depth)
|
|
|
|
# for fine sky only (depth is worse than metric3d)
|
|
if "da3" in model_type:
|
|
model = models[model_type.split('+').index("da3")]
|
|
with redirect_stdout(StringIO()):
|
|
outputs = model.inference([
|
|
# image_original,
|
|
Image.fromarray((255*image[0].clip(0,1)).to('cpu',torch.uint8).numpy())
|
|
])
|
|
depth, sky = outputs.depth[-1], outputs.sky[-1]
|
|
# sky_original = outputs.sky[0] # doesn't work well with fisheye circle
|
|
sky = torch.from_numpy(sky)[None][None].cuda()
|
|
sky = torch.nn.functional.interpolate(
|
|
sky.float(), size=(h, w), mode='bilinear', align_corners=False
|
|
).permute(0, 2, 3, 1) > 0.5 # [h, w, 1]
|
|
|
|
assert pred_depth is not None
|
|
sky_diffused = expand_white_area(sky, pad) # TODO: might hurt distant details
|
|
sky_depth = torch.quantile(pred_depth[~sky_diffused], 0.9999).item()
|
|
pred_depth = torch.where(sky, sky_depth*torch.ones_like(pred_depth), pred_depth.clip(max=sky_depth))
|
|
|
|
# for rough sky (can be inaccurate / miss high frequency details)
|
|
# TODO: SAM-3
|
|
if "lang-sam" in model_type:
|
|
model = models[model_type.split('+').index("lang-sam")]
|
|
box_threshold = 0.3
|
|
text_threshold = 0.25
|
|
with redirect_stdout(StringIO()):
|
|
results = model.predict([image_original], ["sky"], box_threshold, text_threshold)
|
|
for output in results:
|
|
masks = output['masks']
|
|
if not np.any(masks):
|
|
continue
|
|
masks = np.any(masks, axis=0)
|
|
if pred_sky is None:
|
|
pred_sky = masks
|
|
else:
|
|
pred_sky |= masks
|
|
if pred_sky is not None:
|
|
pred_sky = torch.from_numpy(pred_sky)[None][None].cuda()
|
|
pred_sky = torch.nn.functional.interpolate(
|
|
pred_sky.float(), size=(h, w), mode='bilinear', align_corners=False
|
|
).permute(0, 2, 3, 1)[0] > 0.5 # [h, w, 1]
|
|
|
|
if depth_save_path is not None:
|
|
pred_depth /= sky_depth
|
|
if axes is None:
|
|
# Undistort path: metric3d gives linear (z) depth in the pinhole
|
|
# frame. When ray depth is requested, convert before distorting
|
|
# back to the input frame (the split path already handles this via
|
|
# warp_depth_pinhole_to_wide's is_ray_depth).
|
|
if is_ray_depth:
|
|
pred_depth = linear_to_ray_depth(pred_depth, intrins[1])
|
|
pred_depth = distort_image(pred_depth, *intrins)
|
|
pred_depth = pred_depth[0]
|
|
if pred_sky is not None:
|
|
pred_depth = torch.where(
|
|
pred_sky & (pred_depth == 0.0),
|
|
torch.ones_like(pred_depth),
|
|
pred_depth
|
|
)
|
|
pred_depth = torch.clip(65535*pred_depth, 0, 65535).cpu().numpy().astype(np.uint16)
|
|
os.makedirs(Path(depth_save_path).parent, exist_ok=True)
|
|
Image.fromarray(pred_depth.squeeze(-1), mode='I;16').save(depth_save_path)
|
|
|
|
if normal_save_path is not None:
|
|
# TODO: probably also need to distort values of normals
|
|
pred_normal = 0.5 + 0.5 * pred_normal / torch.norm(pred_normal, dim=-1, keepdim=True)
|
|
if axes is None:
|
|
pred_normal = distort_image(pred_normal, *intrins) # TODO: normalize?
|
|
pred_normal = pred_normal[0]
|
|
pred_normal = torch.clip(255*pred_normal, 0, 255).to(torch.uint8).cpu().numpy()
|
|
os.makedirs(Path(normal_save_path).parent, exist_ok=True)
|
|
Image.fromarray(pred_normal).save(normal_save_path)
|
|
|
|
|
|
def update_intrins(in_obj, out_obj):
|
|
out_obj = {**out_obj}
|
|
|
|
if 'camera_model' in in_obj:
|
|
out_obj['camera_model'] = {
|
|
'SIMPLE_PINHOLE': "pinhole",
|
|
'PINHOLE': "pinhole",
|
|
'SIMPLE_RADIAL': "pinhole",
|
|
'SIMPLE_RADIAL_FISHEYE': "fisheye",
|
|
'RADIAL': "pinhole",
|
|
'RADIAL_FISHEYE': "fisheye",
|
|
'OPENCV': "pinhole",
|
|
'OPENCV_FISHEYE': "fisheye",
|
|
'THIN_PRISM_FISHEYE': "fisheye",
|
|
'FISHEYE': "fisheye",
|
|
}[in_obj['camera_model']]
|
|
elif 'camera_model' not in out_obj:
|
|
out_obj['camera_model'] = 'pinhole'
|
|
|
|
for key in 'w h fl_x fl_y cx cy k1 k2 k3 k4 p1 p2 sx1 sy1 b1 b2'.split():
|
|
if key in in_obj:
|
|
out_obj[key] = in_obj[key]
|
|
elif key not in out_obj:
|
|
out_obj[key] = 0.0
|
|
|
|
return out_obj
|
|
|
|
|
|
def process_dir(dataset_dir: str, include_normal: bool, include_sky: bool):
|
|
image_dir = os.path.join(dataset_dir, "images")
|
|
|
|
def with_ext(file_path, new_ext):
|
|
base_name, _ = os.path.splitext(file_path)
|
|
return base_name + '.' + new_ext
|
|
|
|
image_filenames = []
|
|
with open(os.path.join(dataset_dir, "transforms.json")) as fp:
|
|
content = json.load(fp)
|
|
global_intrins = update_intrins(content, {})
|
|
for frame in content['frames']:
|
|
file_path = frame['file_path'].lstrip('./')
|
|
assert file_path.startswith("images")
|
|
image_filename = file_path[len("images")+1:]
|
|
depth_filename = with_ext(image_filename, 'png')
|
|
normal_filename = with_ext(image_filename, 'png')
|
|
image_filenames.append((
|
|
image_filename, depth_filename, normal_filename,
|
|
update_intrins(frame, global_intrins)
|
|
))
|
|
frame["depth_file_path"] = os.path.join("depths", depth_filename)
|
|
if include_normal:
|
|
frame["normal_file_path"] = os.path.join("normals", normal_filename)
|
|
elif 'normal_file_path' in frame and False:
|
|
del frame['normal_file_path']
|
|
if len(image_filenames) == 0:
|
|
print("No image found")
|
|
return
|
|
image_filenames.sort()
|
|
with open(os.path.join(dataset_dir, "transforms.json"), 'w') as fp:
|
|
json.dump(content, fp, indent=4)
|
|
|
|
depth_dir = os.path.join(dataset_dir, "depths")
|
|
normal_dir = os.path.join(dataset_dir, "normals")
|
|
os.makedirs(depth_dir, exist_ok=True)
|
|
if include_normal:
|
|
os.makedirs(normal_dir, exist_ok=True)
|
|
|
|
# Metric3D v2 for good depth/normal
|
|
print("Loading Metric3D v2 model...")
|
|
try:
|
|
model_metric3d = torch.hub.load('yvanyin/metric3d', 'metric3d_vit_large', pretrain=True)
|
|
except ImportError:
|
|
print("mmengine not found. Please run `pip install mmcv`.")
|
|
exit(0)
|
|
model_metric3d = model_metric3d.eval().cuda().bfloat16()
|
|
print("Metric3D v2 model loaded")
|
|
|
|
if include_sky:
|
|
# DA2 for good sky
|
|
print("Loading Depth Anything 3 model...")
|
|
try:
|
|
from depth_anything_3.api import DepthAnything3
|
|
except ImportError:
|
|
print("Depth Anything v3 not found. Please install following https://huggingface.co/depth-anything/DA3MONO-LARGE.")
|
|
exit(0)
|
|
model_da3 = DepthAnything3.from_pretrained("depth-anything/da3mono-large").cuda()
|
|
print("Depth Anything 3 model loaded")
|
|
|
|
# lang-sam for rough sky; TODO: SAM-3
|
|
print("Loading lang-sam model...")
|
|
try:
|
|
from lang_sam.lang_sam import LangSAM
|
|
except ImportError:
|
|
print("lang-sam not found. Please install https://github.com/luca-medeiros/lang-segment-anything.")
|
|
exit(0)
|
|
model_langsam = LangSAM("sam2.1_hiera_large", device="cuda")
|
|
models = (model_metric3d, model_da3, model_langsam)
|
|
print("lang-sam model loaded")
|
|
|
|
else:
|
|
models = (model_metric3d,)
|
|
|
|
for (image_filename, depth_filename, normal_filename, intrins) \
|
|
in tqdm(image_filenames, "Predicting depths"):
|
|
process_image(
|
|
"metric3d+da3+lang-sam" if include_sky else "metric3d",
|
|
models, intrins,
|
|
os.path.join(image_dir, image_filename),
|
|
os.path.join(depth_dir, depth_filename),
|
|
os.path.join(normal_dir, normal_filename),
|
|
)
|
|
|
|
|
|
def _tristate(value):
|
|
return {"auto": None, "yes": True, "no": False}[value]
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import argparse
|
|
parser = argparse.ArgumentParser(
|
|
description="Generate depth and normal maps.")
|
|
parser.add_argument("dataset_dir", nargs=1, help="Path to the dataset folder.")
|
|
# parser.add_argument("--normal", action="store_true", help="Whether to predict normal in addition to depth.")
|
|
parser.add_argument("--sky", action="store_true", help="Whether to predict sky for full images. Useful for highly distorted fisheye images.")
|
|
parser.add_argument(
|
|
"--warp-to-pinhole", choices=["auto", "yes", "no"], default="auto",
|
|
help="Whether to split a wide fisheye into pinhole faces before running "
|
|
"the model. 'auto' uses the FOV-loss heuristic (default).")
|
|
parser.add_argument(
|
|
"--ray-depth", choices=["auto", "yes", "no"], default="auto",
|
|
help="Whether the saved depth stores ray depth (vs. linear/z depth). "
|
|
"'auto' picks ray depth iff the wide->pinhole split is used "
|
|
"(default). Must match the trainer's input_depth_is_ray_depth.")
|
|
args = parser.parse_args()
|
|
Config.warp_to_pinhole = _tristate(args.warp_to_pinhole)
|
|
Config.ray_depth = _tristate(args.ray_depth)
|
|
# process_dir(args.dataset_dir[0], args.normal, args.sky)
|
|
process_dir(args.dataset_dir[0], True, args.sky)
|