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