imag vector generation with dimentionality reduction
This commit is contained in:
@@ -5,9 +5,12 @@ import logging
|
||||
import requests
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.schemas import AuthConfigOut, HealthOut
|
||||
from app.api.schemas import AuthConfigOut, HealthOut, ImageVectorsOut
|
||||
from app.infrastructure.security import auth_config_summary
|
||||
from app.infrastructure.settings import OLLAMA_BASE_URL, OLLAMA_MODEL_NAME, EMBEDDINGS_MODEL
|
||||
from app.infrastructure.settings import (
|
||||
ENABLE_IMAGE_VECTORS, OLLAMA_BASE_URL, OLLAMA_MODEL_NAME, EMBEDDINGS_MODEL,
|
||||
)
|
||||
from app.services import image_embedder
|
||||
from app.services.vector_store import _connect # internal, but handy for a connectivity probe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -56,4 +59,8 @@ def health() -> HealthOut:
|
||||
ollama_model=OLLAMA_MODEL_NAME,
|
||||
embeddings_model=EMBEDDINGS_MODEL,
|
||||
auth=AuthConfigOut(**auth_config_summary()),
|
||||
# Same idea as `auth`: img_vector staying NULL after a deploy has one
|
||||
# usual cause (the model file is not in the image), and it must be
|
||||
# visible from outside the container. status() never loads the model.
|
||||
image_vectors=ImageVectorsOut(enabled=ENABLE_IMAGE_VECTORS, **image_embedder.status()),
|
||||
)
|
||||
|
||||
@@ -2,11 +2,28 @@ from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Query
|
||||
from fastapi import APIRouter, File, Form, HTTPException, Query, UploadFile
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from app.api.schemas import SearchOut, SourceProductOut
|
||||
from app.infrastructure.settings import SEARCH_DEFAULT_TOP_K, SEARCH_MAX_TOP_K
|
||||
from app.api.routers.brands import _row_to_product_out
|
||||
from app.api.schemas import (
|
||||
ImageMatchOut,
|
||||
ImageSearchOut,
|
||||
ImageVectorSearchRequest,
|
||||
SearchOut,
|
||||
SourceProductOut,
|
||||
)
|
||||
from app.infrastructure.settings import (
|
||||
IMAGE_SEARCH_DEFAULT_MIN_SCORE,
|
||||
IMAGE_SEARCH_DEFAULT_TOP_K,
|
||||
IMAGE_SEARCH_MAX_TOP_K,
|
||||
IMAGE_VECTOR_MAX_BYTES,
|
||||
SEARCH_DEFAULT_TOP_K,
|
||||
SEARCH_MAX_TOP_K,
|
||||
)
|
||||
from app.services import image_embedder
|
||||
from app.services.catalog_search import search_catalog
|
||||
from app.services.image_match import ImageSearchResult, InvalidVectorError, search_by_vector
|
||||
|
||||
router = APIRouter(tags=["search"])
|
||||
|
||||
@@ -46,3 +63,100 @@ def catalog_search_endpoint(
|
||||
detected_brand=result.detected_brand,
|
||||
detected_category=result.detected_category,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Search by image
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public like GET /search. Two ways in, one ranking: the Nearle app embeds
|
||||
# the cropped photo on-device with the same MobileNetV3 model that filled
|
||||
# img_vector and POSTs the 1024 floats; anything without the model POSTs the
|
||||
# photo and this API embeds it (bounded: IMAGE_VECTOR_MAX_BYTES, one
|
||||
# inference at a time behind the embedder's lock).
|
||||
|
||||
def _to_image_search_out(result: ImageSearchResult) -> ImageSearchOut:
|
||||
matches = []
|
||||
for row in result.rows:
|
||||
card = _row_to_product_out(row, row.get("brand") or "")
|
||||
matches.append(ImageMatchOut(
|
||||
**card.model_dump(),
|
||||
score=round(float(row["score"]), 4),
|
||||
text_overlap=float(row.get("text_overlap", 0.0)),
|
||||
))
|
||||
return ImageSearchOut(
|
||||
results=matches,
|
||||
total=len(matches),
|
||||
detected_brand=result.detected_brand,
|
||||
scoped_to_brand=result.scoped_to_brand,
|
||||
scope_fallback=result.scope_fallback,
|
||||
min_score=result.min_score,
|
||||
top_k=result.top_k,
|
||||
query_text=result.query_text,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/search/image-vector", response_model=ImageSearchOut)
|
||||
def image_vector_search_endpoint(body: ImageVectorSearchRequest) -> ImageSearchOut:
|
||||
"""Products that look like the photo whose embedding is `vector`.
|
||||
|
||||
`vector` is the 1024-float, L2-normalised MobileNetV3-Small embedding the
|
||||
app computes on-device. `score` on each result is cosine similarity
|
||||
(1 - pgvector distance). Optional `text` - the OCR read of the label -
|
||||
narrows the search to the brand it names and picks the right pack size
|
||||
among products that share one photo. An explicit `brand` is a hard
|
||||
filter; a brand recognised from `text` falls back to every brand when it
|
||||
finds nothing (`scope_fallback`).
|
||||
"""
|
||||
try:
|
||||
result = search_by_vector(
|
||||
body.vector, text=body.text, brand=body.brand, category=body.category,
|
||||
top_k=body.top_k, min_score=body.min_score,
|
||||
)
|
||||
except InvalidVectorError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc))
|
||||
return _to_image_search_out(result)
|
||||
|
||||
|
||||
@router.post("/search/image", response_model=ImageSearchOut)
|
||||
async def image_search_endpoint(
|
||||
file: UploadFile = File(..., description="The product photo (JPEG/PNG/WebP), ideally cropped to the pack"),
|
||||
text: Optional[str] = Form(None, max_length=500, description="OCR text read off the label"),
|
||||
brand: Optional[str] = Form(None, max_length=120),
|
||||
category: Optional[str] = Form(None, max_length=120),
|
||||
top_k: int = Form(IMAGE_SEARCH_DEFAULT_TOP_K, ge=1, le=IMAGE_SEARCH_MAX_TOP_K),
|
||||
min_score: float = Form(IMAGE_SEARCH_DEFAULT_MIN_SCORE, ge=-1.0, le=1.0),
|
||||
) -> ImageSearchOut:
|
||||
"""Same as /search/image-vector, but the API embeds the photo itself.
|
||||
|
||||
503 when this deployment has no embedding model (GET /api/health ->
|
||||
image_vectors.model_present says so); send a vector instead.
|
||||
"""
|
||||
content = await file.read()
|
||||
if not content:
|
||||
raise HTTPException(status_code=400, detail="The uploaded image is empty.")
|
||||
if len(content) > IMAGE_VECTOR_MAX_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"Image is {len(content) / 1_048_576:.1f} MB; the limit is "
|
||||
f"{IMAGE_VECTOR_MAX_BYTES // 1_048_576} MB. Crop or downscale it.",
|
||||
)
|
||||
if not await run_in_threadpool(image_embedder.available):
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="The image embedding model is not available on this deployment. "
|
||||
"Embed the photo client-side and POST the vector to /api/search/image-vector.",
|
||||
)
|
||||
vector = await run_in_threadpool(image_embedder.embedding_for_bytes, content)
|
||||
if vector is None:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail="Could not decode the image (unsupported format, corrupt data, or too many pixels).",
|
||||
)
|
||||
try:
|
||||
result = await run_in_threadpool(
|
||||
search_by_vector, vector, text=text, brand=brand, category=category,
|
||||
top_k=top_k, min_score=min_score,
|
||||
)
|
||||
except InvalidVectorError as exc: # cannot happen for a model output, but the route must not 500
|
||||
raise HTTPException(status_code=422, detail=str(exc))
|
||||
return _to_image_search_out(result)
|
||||
|
||||
@@ -1,9 +1,16 @@
|
||||
"""Pydantic request/response models for the FastAPI layer."""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
from app.infrastructure.settings import (
|
||||
IMAGE_SEARCH_DEFAULT_MIN_SCORE,
|
||||
IMAGE_SEARCH_DEFAULT_TOP_K,
|
||||
IMAGE_SEARCH_MAX_TOP_K,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -120,6 +127,17 @@ class AuthConfigOut(BaseModel):
|
||||
api_keys_source: str = "default"
|
||||
|
||||
|
||||
class ImageVectorsOut(BaseModel):
|
||||
"""Why img_vector is (or is not) being filled. Reported without loading
|
||||
the model. `model_present=false` after a deploy means the .tflite was not
|
||||
shipped in the image - the one failure this feature absorbs silently."""
|
||||
enabled: bool = True
|
||||
model_path: str = ""
|
||||
model_present: bool = False
|
||||
runtime_importable: bool = False
|
||||
state: str = "unknown"
|
||||
|
||||
|
||||
class HealthOut(BaseModel):
|
||||
status: str
|
||||
database: bool
|
||||
@@ -127,6 +145,9 @@ class HealthOut(BaseModel):
|
||||
ollama_model: str
|
||||
embeddings_model: str
|
||||
auth: AuthConfigOut
|
||||
# Defaulted so a client of this schema still validates against a
|
||||
# deployment predating the field.
|
||||
image_vectors: ImageVectorsOut = Field(default_factory=ImageVectorsOut)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -214,6 +235,52 @@ class SuggestOut(BaseModel):
|
||||
suggestions: List[SuggestionOut]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Image search (POST /api/search/image-vector, POST /api/search/image)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ImageVectorSearchRequest(BaseModel):
|
||||
"""A phone photo's embedding, as the Nearle app computes it on-device."""
|
||||
vector: List[float] = Field(
|
||||
..., min_length=1024, max_length=1024,
|
||||
description="L2-normalised MobileNetV3-Small embedding, 1024 floats",
|
||||
)
|
||||
text: Optional[str] = Field(None, max_length=500, description="OCR text read off the label")
|
||||
brand: Optional[str] = Field(None, max_length=120, description="Restrict to one brand (no fallback)")
|
||||
category: Optional[str] = Field(None, max_length=120)
|
||||
top_k: int = Field(IMAGE_SEARCH_DEFAULT_TOP_K, ge=1, le=IMAGE_SEARCH_MAX_TOP_K)
|
||||
min_score: float = Field(IMAGE_SEARCH_DEFAULT_MIN_SCORE, ge=-1.0, le=1.0,
|
||||
description="Drop matches with cosine similarity below this")
|
||||
|
||||
@field_validator("vector")
|
||||
@classmethod
|
||||
def _finite_and_nonzero(cls, v: List[float]) -> List[float]:
|
||||
if not all(math.isfinite(x) for x in v):
|
||||
raise ValueError("vector contains NaN or infinite values")
|
||||
if math.sqrt(sum(x * x for x in v)) < 1e-6:
|
||||
raise ValueError("vector is all zeros")
|
||||
return v
|
||||
|
||||
|
||||
class ImageMatchOut(ProductOut):
|
||||
"""One catalog product that looks like the photo: the product card plus
|
||||
how close it is. `score` is cosine similarity (1 - pgvector distance);
|
||||
`text_overlap` is the label-text tie-break weight, 0 when no text was sent."""
|
||||
score: float
|
||||
text_overlap: float = 0.0
|
||||
|
||||
|
||||
class ImageSearchOut(BaseModel):
|
||||
results: List[ImageMatchOut]
|
||||
total: int
|
||||
detected_brand: Optional[str] = None
|
||||
scoped_to_brand: bool = False
|
||||
scope_fallback: bool = False
|
||||
min_score: float = 0.0
|
||||
top_k: int = 0
|
||||
query_text: Optional[str] = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RAG chat
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user