File size: 2,509 Bytes
5a64372
0f4c637
053108d
4186a52
75ef19f
 
ddc147e
 
053108d
 
8087e5f
cfa43ed
053108d
0f4c637
 
 
 
 
 
 
 
 
053108d
75ef19f
 
 
 
 
 
 
 
053108d
75ef19f
 
 
 
 
 
053108d
75ef19f
 
 
8cc593a
75ef19f
 
053108d
75ef19f
 
 
 
 
 
 
 
053108d
75ef19f
 
 
053108d
75ef19f
 
053108d
75ef19f
 
053108d
75ef19f
 
 
 
 
053108d
 
 
 
 
 
 
 
 
b247ae9
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
import gradio as gr
import json, os, sys, glob, shutil, subprocess
from PIL import Image

if not os.path.isdir('/tmp/HorizonNet'):
    os.system("git clone https://github.com/sunset1995/HorizonNet.git /tmp/HorizonNet")
sys.path.insert(0, "/tmp/HorizonNet")

from huggingface_hub import hf_hub_download

CKPT = hf_hub_download(repo_id="gum-tech/horizonnet-resnet50-rnn", filename="resnet50_rnn__st3d.pth")
print("Model checkpoint at:", CKPT)

def run(cmd, cwd="/tmp/HorizonNet"):
    result = subprocess.run(
        cmd, cwd=cwd, shell=False,
        stdout=subprocess.PIPE, stderr=subprocess.STDOUT
    )
    log = result.stdout.decode("utf-8", errors="replace")
    print(f"=== {cmd[0]} (exit {result.returncode}) ===\n{log}\n===")
    return result.returncode, log

def predict(image):
    try:
        for d in ["/tmp/hn_input", "/tmp/hn_pre", "/tmp/hn_out"]:
            shutil.rmtree(d, ignore_errors=True)
            os.makedirs(d)

        img_path = "/tmp/hn_input/room.png"
        image.convert("RGB").resize((1024, 512)).save(img_path)
        print("Image saved to", img_path)

        # Step 1: preprocess
        code, log = run([
            "python", "preprocess.py",
            "--img_glob", img_path,
            "--output_dir", "/tmp/hn_pre"
        ])

        aligned = glob.glob("/tmp/hn_pre/*_aligned_rgb.png")
        if not aligned:
            return json.dumps({"error": "Preprocessing failed", "code": code, "log": log[-1000:]})

        aligned_path = aligned[0]
        print("Aligned image:", aligned_path)

        # Step 2: inference
        code, log = run([
            "python", "inference.py",
            "--pth", CKPT,
            "--img_glob", aligned_path,
            "--output_dir", "/tmp/hn_out",
            "--no_cuda"
        ])

        out_jsons = glob.glob("/tmp/hn_out/*.json")
        if not out_jsons:
            return json.dumps({"error": "Inference failed", "code": code, "log": log[-1000:]})

        with open(out_jsons[0]) as f:
            result = json.load(f)

        print("Result:", result)
        return json.dumps(result)

    except Exception as e:
        import traceback
        tb = traceback.format_exc()
        print("=== EXCEPTION ===\n", tb)
        return json.dumps({"error": str(e), "traceback": tb})

demo = gr.Interface(
    fn=predict,
    inputs=gr.Image(type="pil"),
    outputs=gr.Text(),
    title="HorizonNet API",
    description="Upload an equirectangular panorama to extract room layout corners."
)

demo.launch()