| """Generate embeddings AND breed predictions for a dataset in a SINGLE model pass per image. |
| |
| This is the batch path to use going forward: when the embedder and breed classifier are the same |
| HF model (``EMBEDDER=hf`` + ``BREED_CLASSIFIER=hf`` on the same repo), each image forwards through |
| the model exactly once, producing both the re-ID embedding and the breed softmax. It replaces |
| running ``scripts.embed_dataset`` and a separate breed pass (which forward every image through the |
| model twice). Idempotent: skips pictures that already have both, so it's safe to re-run / resume. |
| |
| Usage (from the backend dir): |
| python -m scripts.process_dataset --dataset-id 3 # one dataset |
| python -m scripts.process_dataset --all # every dataset |
| python -m scripts.process_dataset --all --limit 5 # quick smoke (first 5 pictures each) |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import time |
|
|
| from sqlalchemy import select |
|
|
| from app.db import SessionLocal, engine |
| from app.ml import get_breed_classifier, get_embedder, hf_breed_name_version |
| from app.models import Base, Dataset, Picture |
| from app.services.datasets import _known_ids, _picture_ids, _unknown_ids |
| from app.services.images import _same_hf_model, embed_and_breed_picture |
|
|
| BATCH = 100 |
|
|
|
|
| def run(dataset_id: int | None, do_all: bool, limit: int | None) -> None: |
| Base.metadata.create_all(bind=engine) |
| db = SessionLocal() |
| try: |
| embedder = get_embedder() |
| print(f"Active embedder: {embedder.name}/{embedder.version} (dim {embedder.dim})") |
| if _same_hf_model(): |
| bn, bv = hf_breed_name_version() |
| print(f"Single-pass mode: one forward per image -> embedding + breed ({bn}/{bv}).") |
| else: |
| c = get_breed_classifier() |
| print( |
| f"Two-model mode: embedder + separate breed classifier {c.name}/{c.version} " |
| "(embedder and breed are not the same HF model)." |
| ) |
|
|
| if do_all: |
| datasets = db.execute(select(Dataset).order_by(Dataset.id)).scalars().all() |
| elif dataset_id is not None: |
| d = db.get(Dataset, dataset_id) |
| datasets = [d] if d else [] |
| if not d: |
| print(f"Dataset #{dataset_id} not found.") |
| else: |
| print("Pass --dataset-id N or --all.") |
| return |
|
|
| for ds in datasets: |
| pic_ids = _picture_ids(db, _known_ids(db, ds.id), _unknown_ids(db, ds.id)) |
| if limit: |
| pic_ids = pic_ids[:limit] |
| print(f"\nDataset #{ds.id} {ds.name!r}: {len(pic_ids)} picture(s) to consider") |
| embedded = breeds = skipped = errors = 0 |
| start = time.time() |
| for i, pid in enumerate(pic_ids, 1): |
| pic = db.get(Picture, pid) |
| if pic is None: |
| continue |
| try: |
| emb, breed = embed_and_breed_picture(db, pic, skip_if_exists=True) |
| embedded += int(emb) |
| breeds += int(breed) |
| skipped += int(not emb and not breed) |
| except Exception as exc: |
| errors += 1 |
| print(f" ! picture {pid}: {exc}") |
| if i % BATCH == 0: |
| db.commit() |
| rate = i / max(time.time() - start, 1e-6) |
| print( |
| f" {i}/{len(pic_ids)} embedded={embedded} breeds={breeds} " |
| f"skipped={skipped} errors={errors} ({rate:.1f} img/s)" |
| ) |
| db.commit() |
| print( |
| f" done: embedded={embedded} breeds={breeds} skipped={skipped} errors={errors} " |
| f"in {time.time() - start:.0f}s" |
| ) |
| finally: |
| db.close() |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser( |
| description="Embed + predict breeds for a dataset in one model pass per image." |
| ) |
| parser.add_argument("--dataset-id", type=int, default=None) |
| parser.add_argument("--all", action="store_true", help="Process every dataset") |
| parser.add_argument("--limit", type=int, default=None, help="Only the first N pictures (smoke test)") |
| args = parser.parse_args() |
| run(args.dataset_id, args.all, args.limit) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|