File size: 2,591 Bytes
ba7051a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a74422
 
ba7051a
2a74422
ba7051a
2a74422
 
 
ba7051a
 
 
 
2a74422
ba7051a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-FileCopyrightText: (c) 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0

"""Run MoGe-2 on a Blackhole p150a and save depth / normal visualizations.

    python -m scripts.make_demo --image data/example.jpg --out media

Produces ``<out>/source.png``, ``<out>/depth.png`` and ``<out>/normal.png``.
Depth/normal colorization uses the upstream ``moge.utils.vis`` helpers.
"""
from __future__ import annotations

import argparse
import os

import numpy as np
import torch
from PIL import Image

import tt_moge  # noqa: F401
from tt_moge.reference.load_pretrained import load_moge2
from tt_moge.tt.ttnn_moge import TtMoGe

NUM_TOKENS = int(os.environ.get("MOGE_NUM_TOKENS", "1800"))


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--image", default=os.environ.get("MOGE_IMAGE", "data/example.jpg"))
    ap.add_argument("--out", default="media")
    ap.add_argument("--device-id", type=int, default=int(os.environ.get("MOGE_DEVICE", "0")))
    args = ap.parse_args()

    os.makedirs(args.out, exist_ok=True)
    torch.set_grad_enabled(False)

    im = Image.open(args.image).convert("RGB")
    im.save(os.path.join(args.out, "source.png"))
    arr = np.asarray(im, dtype=np.float32) / 255.0
    image = torch.from_numpy(arr).permute(2, 0, 1)[None].float()

    import ttnn  # noqa: F401
    from tt_moge.device import close_device, open_device

    # the p150 configuration (ETH dispatch, 1 CQ, 12x10); MOGE_DISPATCH=worker [MOGE_2CQ=1]: Galaxy-only opt-in
    two_cq = bool(int(os.environ.get("MOGE_2CQ", "0")))
    device, _info = open_device(args.device_id, dispatch=os.environ.get("MOGE_DISPATCH", "eth"),
                                grid=os.environ.get("MOGE_GRID", "12x10"), l1_small_size=32768,
                                trace_region_size=1500000000, num_command_queues=2 if two_cq else 1)
    try:
        model = TtMoGe(ref_moge_model=load_moge2(), device=device)
        out = model(image, num_tokens=NUM_TOKENS)
    finally:
        close_device(device)

    from moge.utils.vis import colorize_depth, colorize_normal

    depth = out["points"][0, ..., 2].cpu().numpy()
    mask = out["mask"][0].cpu().numpy() > 0.5 if "mask" in out else None
    depth_vis = colorize_depth(depth, mask=mask)
    Image.fromarray(depth_vis).save(os.path.join(args.out, "depth.png"))

    if "normal" in out:
        normal_vis = colorize_normal(out["normal"][0].cpu().numpy())
        Image.fromarray(normal_vis).save(os.path.join(args.out, "normal.png"))

    print(f"wrote source/depth/normal to {args.out}/")


if __name__ == "__main__":
    main()