moge-2-p150 / code /examples /quickstart.py
changh95's picture
Python API: warm-up so the first real call is fast (warmup_variants, model.warmup()), quiet logs, install extras
524ca6d verified
Raw History Blame Contribute Delete
4.02 kB
# SPDX-License-Identifier: Apache-2.0
"""Quickstart: the README / model-card Python snippet on the demo image, with the outputs saved.
python code/examples/quickstart.py [--image PATH] [--out moge_output] [--device-id 0]
Runs from any directory: the default image is the repo's ``media/source.png``, found relative to this
file. A relative ``--image`` or ``--out`` path is relative to the current directory.
Writes ``<out>/depth.png`` and ``<out>/normal.png`` (colorized), ``<out>/result.npz`` (points, depth,
normal, mask, intrinsics, metric_scale; the server's npz layout) and ``<out>/result.json`` (scalars).
"""
from __future__ import annotations
import argparse
import json
import os
import numpy as np
from PIL import Image
HERE = os.path.dirname(os.path.abspath(__file__))
DEMO = os.path.normpath(os.path.join(HERE, "..", "..", "media", "source.png"))
def colorize(out):
"""(depth RGB, normal RGB) uint8 images; the upstream moge.utils.vis colouring when matplotlib is present."""
try:
from moge.utils.vis import colorize_depth, colorize_normal # vendored upstream helpers (need matplotlib)
return colorize_depth(out.depth, mask=out.mask), colorize_normal(out.normal, mask=out.mask)
except ImportError:
d = np.where(out.mask, 1.0 / out.depth, np.nan)
lo, hi = np.nanquantile(d, 0.001), np.nanquantile(d, 0.99)
g = (np.nan_to_num((d - lo) / max(hi - lo, 1e-12), nan=0.0).clip(0, 1) * 255).astype(np.uint8)
n = ((out.normal * [0.5, -0.5, -0.5] + 0.5).clip(0, 1) * 255).astype(np.uint8)
return np.repeat(g[..., None], 3, -1), n
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--image", default=DEMO)
ap.add_argument("--out", default="moge_output")
ap.add_argument("--device-id", type=int, default=0)
args = ap.parse_args()
if not os.path.isfile(args.image):
ap.error(f"{args.image} not found; pass --image <your image> (the demo image is media/source.png of the repo)")
# ---- the model-card snippet -------------------------------------------------------------
from tt_moge import MoGeModel
with MoGeModel.from_pretrained(device_id=args.device_id) as model: # weights from the HF cache
out = model(args.image) # path, PIL image, numpy array or torch tensor
print(out) # MoGeOutput(1920x1080, valid ..., median depth ... m, ...)
depth = out.depth # (H, W) float32, metres; inf where out.mask is False
points = out.points # (H, W, 3) float32, metres, camera space (x right, y down, z forward)
normal = out.normal # (H, W, 3) float32, unit normals
K = out.intrinsics_pixels # (3, 3) camera matrix in pixels
# ------------------------------------------------------------------------------------------
os.makedirs(args.out, exist_ok=True)
depth_rgb, normal_rgb = colorize(out)
Image.fromarray(depth_rgb).save(os.path.join(args.out, "depth.png"))
Image.fromarray(normal_rgb).save(os.path.join(args.out, "normal.png"))
out.save_npz(os.path.join(args.out, "result.npz"))
valid = depth[out.mask]
summary = {
"image": os.path.abspath(args.image), "width": out.width, "height": out.height,
"metric_scale": out.metric_scale, "fov_x_deg": out.fov_x, "fov_y_deg": out.fov_y,
"mask_coverage": float(out.mask.mean()),
"depth_m": {"min": float(valid.min()), "median": float(np.median(valid)), "max": float(valid.max())},
"intrinsics": out.intrinsics.tolist(), "intrinsics_pixels": K.tolist(),
"points_shape": list(points.shape), "normal_shape": list(normal.shape),
}
with open(os.path.join(args.out, "result.json"), "w") as f:
json.dump(summary, f, indent=1)
print(json.dumps(summary, indent=1))
print(f"wrote depth.png, normal.png, result.npz, result.json to {os.path.abspath(args.out)}/")
if __name__ == "__main__":
main()