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