File size: 6,309 Bytes
edbf290
8b5a68d
 
 
 
 
 
62b7002
 
 
 
 
 
fc49781
62b7002
8b5a68d
62b7002
 
 
e9802f8
8b5a68d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fc49781
8b5a68d
 
fc49781
 
 
 
 
 
 
8b5a68d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fc49781
8b5a68d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62b7002
 
 
8b5a68d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62b7002
8b5a68d
 
 
 
fc49781
 
8b5a68d
 
62b7002
8b5a68d
 
 
 
 
fc49781
 
8b5a68d
 
 
 
 
 
 
 
 
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
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
"""Run the 2D floor-plan component detector from Hugging Face Hub or local files.

Example:
    python inference.py plan.png --output-dir results/plan
"""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import numpy as np
import torch
from huggingface_hub import hf_hub_download
from PIL import Image, ImageDraw
from rfdetr import RFDETRMedium


DEFAULT_REPO_ID = "OsamaMo/2dplan2strct"
COLORS = {
    "wall": "#e63946",
    "room": "#457b9d",
    "door": "#f4a261",
    "window": "#2a9d8f",
}


class FloorPlanDetector:
    """Load RF-DETR once, then run it on one or more floor-plan images."""

    def __init__(
        self,
        repo_id: str = DEFAULT_REPO_ID,
        *,
        revision: str | None = None,
        model_dir: Path | None = None,
        device: str = "auto",
    ) -> None:
        self.repo_id = repo_id
        if device not in {"auto", "cpu", "cuda"}:
            raise ValueError("device must be 'auto', 'cpu', or 'cuda'")
        self.device = "cuda" if device == "auto" and torch.cuda.is_available() else device
        if self.device == "auto":
            self.device = "cpu"
        if self.device == "cuda" and not torch.cuda.is_available():
            raise RuntimeError("CUDA was requested, but no CUDA GPU is available")
        if model_dir is None:
            config_path = Path(hf_hub_download(repo_id=repo_id, filename="config.json", revision=revision))
            config = json.loads(config_path.read_text(encoding="utf-8"))
            weights = hf_hub_download(repo_id=repo_id, filename=config["checkpoint"], revision=revision)
        else:
            model_dir = Path(model_dir)
            config = json.loads((model_dir / "config.json").read_text(encoding="utf-8"))
            weights = str(model_dir / config["checkpoint"])

        if config["variant"] != "RFDETRMedium":
            raise ValueError(f"Unsupported model variant: {config['variant']}")
        self.class_names = config["class_names"]
        self.model = RFDETRMedium(
            pretrain_weights=weights,
            resolution=config["resolution"],
            num_classes=config["num_classes"],
            device=self.device,
        )
        if self.model.class_names and self.model.class_names != self.class_names:
            raise ValueError("Checkpoint class names do not match config.json")

    def predict(self, image_path: str | Path, *, threshold: float = 0.35) -> dict:
        """Return boxes in original-image pixel coordinates, sorted by confidence."""
        if not 0 <= threshold <= 1:
            raise ValueError("threshold must be between 0 and 1")
        image_path = Path(image_path)
        with Image.open(image_path) as source:
            image = np.array(source.convert("RGB"))
        height, width = image.shape[:2]
        result = self.model.predict(image, threshold=threshold)
        detections = [
            {
                "label": self.class_names[int(class_id)],
                "score": float(score),
                "box_xyxy": [float(value) for value in box],
            }
            for box, score, class_id in zip(result.xyxy, result.confidence, result.class_id)
        ]
        detections.sort(key=lambda item: item["score"], reverse=True)
        return {
            "model": self.repo_id,
            "image": str(image_path),
            "width": width,
            "height": height,
            "threshold": threshold,
            "detections": detections,
        }


def save_prediction(prediction: dict, image_path: str | Path, output_dir: str | Path) -> tuple[Path, Path]:
    """Write a JSON result and an annotated PNG without changing the source image."""
    output_dir = Path(output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)
    json_path = output_dir / "predictions.json"
    image_out = output_dir / "annotated.png"
    json_path.write_text(json.dumps(prediction, indent=2) + "\n", encoding="utf-8")

    with Image.open(image_path) as source:
        image = source.convert("RGB")
    draw = ImageDraw.Draw(image)
    line_width = max(2, min(image.size) // 300)
    for item in reversed(prediction["detections"]):
        color = COLORS.get(item["label"], "#ffffff")
        box = item["box_xyxy"]
        draw.rectangle(box, outline=color, width=line_width)
        label = f"{item['label']} {item['score']:.2f}"
        text_box = draw.textbbox((box[0], box[1]), label)
        text_height = text_box[3] - text_box[1]
        label_y = max(0, box[1] - text_height - 4)
        draw.rectangle((box[0], label_y, box[0] + text_box[2] - text_box[0] + 4, label_y + text_height + 4), fill=color)
        draw.text((box[0] + 2, label_y + 2), label, fill="white")
    image.save(image_out)
    return json_path, image_out


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("image", type=Path, help="PNG or JPEG floor-plan image")
    parser.add_argument("--repo-id", default=DEFAULT_REPO_ID, help="Hugging Face model repository")
    parser.add_argument("--revision", help="Optional Hub commit SHA or tag for reproducible downloads")
    parser.add_argument("--model-dir", type=Path, help="Use previously downloaded config and weights offline")
    parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto",
                        help="Run on CPU or CUDA; auto uses CUDA when available")
    parser.add_argument("--threshold", type=float, default=0.35, help="Confidence threshold from 0 to 1")
    parser.add_argument("--output-dir", type=Path, default=Path("output/prediction"))
    args = parser.parse_args()
    if not args.image.is_file():
        parser.error(f"image does not exist: {args.image}")
    if not 0 <= args.threshold <= 1:
        parser.error("--threshold must be between 0 and 1")

    detector = FloorPlanDetector(args.repo_id, revision=args.revision, model_dir=args.model_dir,
                                 device=args.device)
    prediction = detector.predict(args.image, threshold=args.threshold)
    json_path, image_path = save_prediction(prediction, args.image, args.output_dir)
    print(f"{len(prediction['detections'])} detections")
    print(f"JSON: {json_path}")
    print(f"Annotated image: {image_path}")


if __name__ == "__main__":
    main()