File size: 4,020 Bytes
ea621a3
 
 
524ca6d
 
 
 
ea621a3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
# 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()