Download code/scripts/make_demo.py from changh95/moge-2-p150: direct link, hf CLI and curl.
- Browser
- Download file 2.59 kB
-
https://huggingface.co/changh95/moge-2-p150/resolve/main/code/scripts/make_demo.py
- Command line
-
hf download hf://changh95/moge-2-p150/code/scripts/make_demo.py
-
curl -L -o make_demo.py https://huggingface.co/changh95/moge-2-p150/resolve/main/code/scripts/make_demo.py
2.59 kB
| # 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() | |