File size: 3,177 Bytes
1bb5a3e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
High-level inference for the JKTSV DINOv3 geolocation model.

Example
-------
    from inference import GeoTagPredictor

    predictor = GeoTagPredictor("nadh0708/JKTSV-modelD")   # HF repo id
    # ...or a local checkpoint:
    predictor = GeoTagPredictor("model/modelD_40e_2.pth")

    result = predictor.predict("street.jpg")
    # {'lat': -6.21, 'lon': 106.84}

    results = predictor.predict(["a.jpg", "b.jpg"])         # batched
    # [{'lat': ..., 'lon': ...}, ...]

The model was trained on Google Street View perspective crops of Jakarta roads
(8 headings, 0-315 deg). Predictions are only meaningful for Jakarta street
imagery.
"""

from __future__ import annotations

from typing import Union

import torch
from PIL import Image
from torchvision import transforms

from modeling_geotag import DinoGeoRegressor

_IMAGENET_MEAN = [0.485, 0.456, 0.406]
_IMAGENET_STD = [0.229, 0.224, 0.225]

ImageInput = Union[str, Image.Image]


class GeoTagPredictor:
    def __init__(
        self,
        model_id_or_path: str,
        device: str | torch.device | None = None,
        filename: str = "pytorch_model.bin",
    ):
        self.device = torch.device(
            device or ("cuda" if torch.cuda.is_available() else "cpu")
        )
        self.model = DinoGeoRegressor.from_pretrained(
            model_id_or_path, filename=filename, device=self.device
        )
        self.transform = transforms.Compose([
            transforms.Resize((224, 224)),
            transforms.ToTensor(),
            transforms.Normalize(mean=_IMAGENET_MEAN, std=_IMAGENET_STD),
        ])

    def _load(self, image: ImageInput) -> Image.Image:
        if isinstance(image, Image.Image):
            return image.convert("RGB")
        return Image.open(image).convert("RGB")

    @torch.no_grad()
    def predict(
        self, images: Union[ImageInput, list[ImageInput]]
    ) -> Union[dict, list[dict]]:
        """Predict (lat, lon) for one image or a list of images."""
        single = not isinstance(images, (list, tuple))
        batch = [images] if single else list(images)

        pixel_values = torch.stack([self.transform(self._load(im)) for im in batch])
        pixel_values = pixel_values.to(self.device)

        lonlat = self.model.predict_lonlat(pixel_values).cpu()
        results = [
            {"lat": float(lat), "lon": float(lon)}
            for lon, lat in lonlat.tolist()
        ]
        return results[0] if single else results


def _cli() -> None:
    import argparse
    import json

    parser = argparse.ArgumentParser(description="Geolocate Jakarta street imagery.")
    parser.add_argument("model", help="HF repo id or local checkpoint path")
    parser.add_argument("images", nargs="+", help="image file path(s)")
    parser.add_argument("--device", default=None)
    parser.add_argument(
        "--filename", default="pytorch_model.bin",
        help="weights filename inside the HF repo",
    )
    args = parser.parse_args()

    predictor = GeoTagPredictor(args.model, device=args.device, filename=args.filename)
    out = predictor.predict(args.images)
    print(json.dumps(out, indent=2))


if __name__ == "__main__":
    _cli()