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()
|