Implement detect_objects Celery task: - Runs YOLOv8n on 640px thumbnail via ONNX Runtime - Creates Tag(kind=object) rows for each COCO class detected - Writes photo_tags associations with confidence, bbox, and source - Wipes previous detections per source model on re-run No new tables/migrations — uses the unified Tag model from PR3. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
243 lines
8.2 KiB
Python
243 lines
8.2 KiB
Python
"""
|
|
Celery tasks for the vision pipeline — embedding, OCR, object detection,
|
|
face recognition.
|
|
|
|
All tasks run on the dedicated `vision` queue with limited concurrency
|
|
(memory-bound CPU inference). They read thumbnails generated by
|
|
generate_thumbnails, so they MUST run after thumbs complete.
|
|
"""
|
|
import asyncio
|
|
import logging
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
from celery import shared_task
|
|
from sqlalchemy import select, delete, text
|
|
from PIL import Image
|
|
|
|
from app.database import AsyncSessionLocal
|
|
from app.models import Photo
|
|
from app.models.embeddings import Embedding
|
|
from app.config import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _load_thumb(photo_id: str, size: str = "medium") -> np.ndarray | None:
|
|
"""Load a thumbnail as an RGB numpy array."""
|
|
thumb_path = Path(f"/data/thumbs/{photo_id}/{size}.webp")
|
|
if not thumb_path.exists():
|
|
logger.warning("Thumbnail not found: %s", thumb_path)
|
|
return None
|
|
img = Image.open(thumb_path).convert("RGB")
|
|
return np.array(img)
|
|
|
|
|
|
@shared_task(name='embed_photo', queue='vision')
|
|
def embed_photo(photo_id: str):
|
|
"""Generate CLIP embedding for a photo and store in pgvector."""
|
|
if not settings.vision.enabled:
|
|
return {'status': 'skipped', 'reason': 'vision disabled'}
|
|
return asyncio.run(_embed_photo_async(photo_id))
|
|
|
|
|
|
async def _embed_photo_async(photo_id: str):
|
|
image = _load_thumb(photo_id, "medium") # 640px
|
|
if image is None:
|
|
return {'status': 'error', 'message': 'thumbnail not found'}
|
|
|
|
from app.services.vision.registry import registry
|
|
embedder = registry.get_embedder()
|
|
vector = embedder.embed_image(image)
|
|
|
|
model_name = settings.vision.embedder.name
|
|
|
|
async with AsyncSessionLocal() as session:
|
|
# Upsert: delete existing then insert
|
|
await session.execute(
|
|
delete(Embedding).where(
|
|
Embedding.photo_id == photo_id,
|
|
Embedding.model == model_name,
|
|
)
|
|
)
|
|
emb = Embedding(
|
|
photo_id=photo_id,
|
|
model=model_name,
|
|
vector=vector.tolist(),
|
|
)
|
|
session.add(emb)
|
|
await session.commit()
|
|
|
|
logger.info("Embedded photo %s with %s", photo_id, model_name)
|
|
return {'status': 'success', 'photo_id': photo_id}
|
|
|
|
|
|
@shared_task(name='vision_fanout', queue='vision')
|
|
def vision_fanout(photo_id: str):
|
|
"""Dispatch all enabled vision tasks for a photo."""
|
|
if not settings.vision.enabled:
|
|
return {'status': 'skipped', 'reason': 'vision disabled'}
|
|
|
|
embed_photo.delay(photo_id)
|
|
|
|
if settings.vision.ocr.enabled:
|
|
ocr_photo.delay(photo_id)
|
|
if settings.vision.detector.enabled:
|
|
detect_objects.delay(photo_id)
|
|
if settings.vision.faces.enabled:
|
|
extract_faces.delay(photo_id)
|
|
|
|
return {'status': 'dispatched', 'photo_id': photo_id}
|
|
|
|
|
|
@shared_task(name='ocr_photo', queue='vision')
|
|
def ocr_photo(photo_id: str):
|
|
"""Run OCR on a photo and store text regions."""
|
|
if not settings.vision.enabled or not settings.vision.ocr.enabled:
|
|
return {'status': 'skipped', 'reason': 'OCR disabled'}
|
|
return asyncio.run(_ocr_photo_async(photo_id))
|
|
|
|
|
|
async def _ocr_photo_async(photo_id: str):
|
|
image = _load_thumb(photo_id, "large") # 1280px for better OCR accuracy
|
|
if image is None:
|
|
return {'status': 'error', 'message': 'thumbnail not found'}
|
|
|
|
from app.services.vision.registry import registry
|
|
ocr_engine = registry.get_ocr()
|
|
results = ocr_engine.run(image)
|
|
|
|
if not results:
|
|
logger.info("No OCR text found for photo %s", photo_id)
|
|
return {'status': 'success', 'photo_id': photo_id, 'regions': 0}
|
|
|
|
from app.models.ocr_text import OCRText
|
|
|
|
async with AsyncSessionLocal() as session:
|
|
# Delete existing OCR results for this photo (re-run safe)
|
|
await session.execute(
|
|
delete(OCRText).where(OCRText.photo_id == photo_id)
|
|
)
|
|
for r in results:
|
|
session.add(OCRText(
|
|
photo_id=photo_id,
|
|
text=r.text,
|
|
language=r.language,
|
|
confidence=r.confidence,
|
|
bbox=r.bbox,
|
|
))
|
|
await session.commit()
|
|
|
|
logger.info("OCR: %d text regions for photo %s", len(results), photo_id)
|
|
return {'status': 'success', 'photo_id': photo_id, 'regions': len(results)}
|
|
|
|
|
|
@shared_task(name='detect_objects', queue='vision')
|
|
def detect_objects(photo_id: str):
|
|
"""Detect objects in a photo, create Tag(kind=object) rows, and
|
|
link via photo_tags with confidence/bbox/source."""
|
|
if not settings.vision.enabled or not settings.vision.detector.enabled:
|
|
return {'status': 'skipped', 'reason': 'detection disabled'}
|
|
return asyncio.run(_detect_objects_async(photo_id))
|
|
|
|
|
|
async def _detect_objects_async(photo_id: str):
|
|
image = _load_thumb(photo_id, "medium") # 640px
|
|
if image is None:
|
|
return {'status': 'error', 'message': 'thumbnail not found'}
|
|
|
|
from app.services.vision.registry import registry
|
|
detector = registry.get_detector()
|
|
detections = detector.detect(image)
|
|
|
|
if not detections:
|
|
logger.info("No objects detected for photo %s", photo_id)
|
|
return {'status': 'success', 'photo_id': photo_id, 'objects': 0}
|
|
|
|
from app.models.tags import Tag, photo_tags
|
|
|
|
source_name = "vision:yolov8n"
|
|
|
|
async with AsyncSessionLocal() as session:
|
|
# Wipe previous detection results for this photo from this model
|
|
await session.execute(
|
|
delete(photo_tags).where(
|
|
photo_tags.c.photo_id == photo_id,
|
|
photo_tags.c.source == source_name,
|
|
)
|
|
)
|
|
|
|
for det in detections:
|
|
# Find or create the object tag
|
|
result = await session.execute(
|
|
select(Tag).where(Tag.name == det.label, Tag.kind == 'object')
|
|
)
|
|
tag = result.scalar_one_or_none()
|
|
if not tag:
|
|
tag = Tag(name=det.label, kind='object', source=source_name)
|
|
session.add(tag)
|
|
await session.flush() # get tag.id
|
|
|
|
# Insert photo_tags association with ML metadata
|
|
await session.execute(
|
|
photo_tags.insert().values(
|
|
photo_id=photo_id,
|
|
tag_id=tag.id,
|
|
confidence=det.confidence,
|
|
bbox=det.bbox,
|
|
source=source_name,
|
|
)
|
|
)
|
|
|
|
await session.commit()
|
|
|
|
labels = [d.label for d in detections]
|
|
logger.info("Detected %d objects in photo %s: %s", len(detections), photo_id, labels)
|
|
return {'status': 'success', 'photo_id': photo_id, 'objects': len(detections)}
|
|
|
|
|
|
@shared_task(name='extract_faces', queue='vision')
|
|
def extract_faces(photo_id: str):
|
|
"""Detect faces and extract embeddings — implemented in PR7."""
|
|
return {'status': 'not_implemented'}
|
|
|
|
|
|
@shared_task(name='backfill_vision')
|
|
def backfill_vision(task: str | None = None, limit: int | None = None):
|
|
"""Queue vision tasks for photos that haven't been processed yet."""
|
|
return asyncio.run(_backfill_vision_async(task, limit))
|
|
|
|
|
|
async def _backfill_vision_async(task: str | None, limit: int | None):
|
|
model_name = settings.vision.embedder.name
|
|
|
|
async with AsyncSessionLocal() as session:
|
|
# Find photos without embeddings
|
|
stmt = text("""
|
|
SELECT p.id FROM photos p
|
|
LEFT JOIN embeddings e ON e.photo_id = p.id AND e.model = :model
|
|
WHERE e.photo_id IS NULL
|
|
AND p.processing_status = 'completed'
|
|
ORDER BY p.created_at DESC
|
|
""")
|
|
if limit:
|
|
stmt = text(str(stmt) + f" LIMIT {limit}")
|
|
|
|
result = await session.execute(stmt, {"model": model_name})
|
|
photo_ids = [row[0] for row in result.fetchall()]
|
|
|
|
count = 0
|
|
for pid in photo_ids:
|
|
if task == 'embed' or task is None:
|
|
embed_photo.delay(pid)
|
|
if task == 'ocr' or task is None:
|
|
ocr_photo.delay(pid)
|
|
if task == 'detect' or task is None:
|
|
detect_objects.delay(pid)
|
|
if task == 'faces' or task is None:
|
|
extract_faces.delay(pid)
|
|
count += 1
|
|
|
|
logger.info("Backfill queued %d photos for vision processing", count)
|
|
return {'status': 'queued', 'count': count}
|