aipluto-backend / scripts /seed_filler.py
andreskoenig's picture
deploy from cli
d32533a verified
Raw
History Blame Contribute Delete
4.43 kB
"""Seed ~200 distractor sightings using the Dog CEO public API (dog.ceo).
The API serves images from the Stanford Dogs dataset via CDN — no HF account,
no dataset scripts, no large downloads. Each image is fetched individually.
Run inside the backend container:
python -m scripts.seed_filler
"""
import argparse
import io
import logging
import random
import time
from collections import defaultdict
import requests
from PIL import Image
from sqlalchemy import select
from tqdm import tqdm
from app.config import get_settings
from app.models import Sighting
from scripts._common import insert_sighting, log, open_session
_API_BATCH = 50 # dog.ceo supports up to 50 random images per request
_API_URL = "https://dog.ceo/api/breeds/image/random/{n}"
def _fetch_urls(total: int) -> list[tuple[str, str]]:
"""Return list of (image_url, breed) from dog.ceo. Breed extracted from URL path."""
results: list[tuple[str, str]] = []
while len(results) < total:
n = min(_API_BATCH, total - len(results))
resp = requests.get(_API_URL.format(n=n), timeout=15)
resp.raise_for_status()
for url in resp.json().get("message", []):
# URL format: https://images.dog.ceo/breeds/{breed}/{file}
parts = url.rstrip("/").split("/")
breed = parts[-2] if len(parts) >= 2 else "unknown"
results.append((url, breed))
return results[:total]
def _download_image(url: str, timeout: int = 20) -> Image.Image | None:
try:
resp = requests.get(url, timeout=timeout)
resp.raise_for_status()
return Image.open(io.BytesIO(resp.content)).convert("RGB")
except Exception as exc: # noqa: BLE001
log.warning("Failed to download %s: %s", url, exc)
return None
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--force",
action="store_true",
help="Re-seed even if filler rows already exist.",
)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(name)s: %(message)s")
settings = get_settings()
rng = random.Random(args.seed)
session = open_session()
try:
existing = session.scalar(
select(Sighting).where(Sighting.source == "filler").limit(1)
)
if existing and not args.force:
log.info("Filler rows already present — skipping (pass --force to re-seed).")
return
if existing and args.force:
deleted = session.query(Sighting).filter(
Sighting.source == "filler"
).delete()
session.commit()
log.info("Deleted %d existing filler rows.", deleted)
target_count = settings.seed_filler_count
# Fetch extra URLs upfront to account for download failures / detection misses.
fetch_count = int(target_count * 1.5)
log.info("Fetching %d dog image URLs from dog.ceo...", fetch_count)
urls = _fetch_urls(fetch_count)
rng.shuffle(urls)
per_breed: dict[str, int] = defaultdict(int)
max_per_breed = max(1, target_count // settings.seed_filler_breeds * 2)
inserted = 0
with tqdm(total=target_count, desc="seed:filler") as pbar:
for url, breed in urls:
if inserted >= target_count:
break
if per_breed[breed] >= max_per_breed:
continue
image = _download_image(url)
if image is None:
continue
row = insert_sighting(
session,
image,
source="filler",
identity=None,
rng=rng,
prefix="filler",
)
if row is None:
time.sleep(0.05) # brief pause after detection miss
continue
per_breed[breed] += 1
inserted += 1
pbar.update(1)
if inserted % 25 == 0:
session.commit()
session.commit()
log.info(
"Inserted %d filler sightings across %d breeds.",
inserted,
len(per_breed),
)
finally:
session.close()
if __name__ == "__main__":
main()