image vector dimensionality reduction
This commit is contained in:
11
.env.example
11
.env.example
@@ -314,12 +314,17 @@ USE_PLAYWRIGHT_FALLBACK=true
|
||||
|
||||
MIN_IMAGE_BYTES=3000
|
||||
|
||||
# img_vector: a 32x32 RGB pixel thumbnail of each product's primary image,
|
||||
# stored on its brand table and computed on a background thread after every
|
||||
# catalog write (app/services/image_vector.py). Existing rows are filled by
|
||||
# img_vector: a MobileNetV3-Small embedding (1024 floats, L2-normalised) of
|
||||
# each product's primary image, stored on its brand table and computed on a
|
||||
# background thread after every catalog write (app/services/image_vector.py,
|
||||
# model in app/services/image_embedder.py). Existing rows are filled by
|
||||
# `python -m scripts.backfill_image_vectors --all --apply`. Set false to stop
|
||||
# the API process from downloading images at all.
|
||||
ENABLE_IMAGE_VECTORS=true
|
||||
# Where the .tflite lives. The default is inside app/ (shipped with the image
|
||||
# and never volume-mounted); only override to point at a different file.
|
||||
#IMAGE_EMBED_MODEL_PATH=/app/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite
|
||||
#IMAGE_EMBED_NUM_THREADS=2
|
||||
|
||||
# Product SKU: try a live web search for a real marketplace product ID
|
||||
# (Amazon ASIN, Flipkart PID, etc.) before falling back to an internal SKU.
|
||||
|
||||
1
.gitattributes
vendored
1
.gitattributes
vendored
@@ -15,6 +15,7 @@
|
||||
# to the wildcard above cannot silently start mangling them.
|
||||
*.joblib binary
|
||||
*.db binary
|
||||
*.tflite binary
|
||||
|
||||
# Windows-only scripts genuinely need CRLF. None are tracked in this repo
|
||||
# today (start_app.bat lives above it), but the rule belongs with the policy.
|
||||
|
||||
@@ -381,27 +381,41 @@ USE_PLAYWRIGHT_FALLBACK = _bool("USE_PLAYWRIGHT_FALLBACK", "true")
|
||||
MIN_IMAGE_BYTES = int(os.getenv("MIN_IMAGE_BYTES", "3000"))
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# img_vector - a 32x32 RGB pixel thumbnail of each product's primary image,
|
||||
# stored on its brand table (app/services/image_vector.py)
|
||||
# img_vector - a MobileNetV3-Small image embedding (1024 floats, L2-normalised)
|
||||
# of each product's primary image, stored on its brand table
|
||||
# (app/services/image_vector.py owns the column, image_embedder.py the model)
|
||||
# ---------------------------------------------------------------------------
|
||||
# Computed off the request path by one bounded worker thread after every
|
||||
# catalog write, and by scripts/backfill_image_vectors.py for existing rows.
|
||||
# Default on: the work is one small download per row just written, never on
|
||||
# a request. tests/conftest.py pins it OFF so the suite never dials the
|
||||
# database named in a developer's .env.
|
||||
# Default on: the work is one small download and one ~30ms inference per row
|
||||
# just written, never on a request. tests/conftest.py pins it OFF so the
|
||||
# suite never dials the database named in a developer's .env.
|
||||
ENABLE_IMAGE_VECTORS = _bool("ENABLE_IMAGE_VECTORS", "true")
|
||||
# A download past this many bytes is abandoned - a wrong URL to a video must
|
||||
# not fill the container.
|
||||
IMAGE_VECTOR_MAX_BYTES = int(os.getenv("IMAGE_VECTOR_MAX_BYTES", str(8 * 1024 * 1024)))
|
||||
# Refused before decoding when the header claims more pixels than this
|
||||
# (decompression-bomb guard; 40MP is well past any product photo).
|
||||
IMAGE_VECTOR_MAX_PIXELS = int(os.getenv("IMAGE_VECTOR_MAX_PIXELS", "40000000"))
|
||||
# (decompression-bomb guard; 25MP is well past any product photo, and the
|
||||
# embedder decodes at full resolution - ~75MB of RGB at this cap).
|
||||
IMAGE_VECTOR_MAX_PIXELS = int(os.getenv("IMAGE_VECTOR_MAX_PIXELS", "25000000"))
|
||||
IMAGE_VECTOR_TIMEOUT_SECONDS = float(os.getenv("IMAGE_VECTOR_TIMEOUT_SECONDS", "15"))
|
||||
# Minimum gap between two requests to the same image host.
|
||||
IMAGE_VECTOR_HOST_PAUSE_SECONDS = float(os.getenv("IMAGE_VECTOR_HOST_PAUSE_SECONDS", "0.5"))
|
||||
# Writes waiting for the worker; beyond this a write's rows are left for the
|
||||
# backfill script rather than queued.
|
||||
IMAGE_VECTOR_QUEUE_MAX = int(os.getenv("IMAGE_VECTOR_QUEUE_MAX", "64"))
|
||||
# The TFLite embedder (mobilenet_v3_small_embedder.tflite, input [1,224,224,3]
|
||||
# float32, output [1,1024]). Lives under app/, NOT data/: /app/data is a named
|
||||
# volume on every deployment (see BUNDLED_ASSETS_DIR), and a file added to
|
||||
# the image under a mounted path is invisible on any volume that already
|
||||
# exists. app/ is copied into the image and never mounted.
|
||||
IMAGE_EMBED_MODEL_PATH = _dir(
|
||||
"IMAGE_EMBED_MODEL_PATH",
|
||||
_BACKEND_ROOT / "app" / "services" / "models" / "mobilenet" / "mobilenet_v3_small_embedder.tflite",
|
||||
)
|
||||
# Intra-op threads for one inference. 2 on the prod host; inference itself is
|
||||
# serialised by a lock (a TFLite interpreter is not thread-safe).
|
||||
IMAGE_EMBED_NUM_THREADS = int(os.getenv("IMAGE_EMBED_NUM_THREADS", "2"))
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# USDA FoodData Central - nutrition for loose, unbranded commodities
|
||||
|
||||
250
app/services/image_embedder.py
Normal file
250
app/services/image_embedder.py
Normal file
@@ -0,0 +1,250 @@
|
||||
"""The image embedding model behind `img_vector`: MobileNetV3-Small, TFLite.
|
||||
|
||||
WHAT IT PRODUCES
|
||||
----------------
|
||||
One 1024-float vector per image, L2-normalised, from
|
||||
`mobilenet_v3_small_embedder.tflite` (input [1, 224, 224, 3] float32, output
|
||||
[1, 1024]). It is the same model, runtime and post-processing a colleague uses
|
||||
for their product images, which is the whole point: two vectors from the same
|
||||
photo must agree, and `<=>` (cosine distance) between our rows and theirs must
|
||||
mean something.
|
||||
|
||||
THIS FILE OWNS TWO THINGS AND NOTHING ELSE
|
||||
------------------------------------------
|
||||
1. `preprocess()` - bytes to the input tensor. It is the ONLY place the
|
||||
resize / crop / scaling decisions live, because those decisions are what
|
||||
make the vectors comparable. See the banner on that function.
|
||||
2. The interpreter - one per process, created lazily on first use, and every
|
||||
call into it serialised by one lock. A TFLite interpreter is not
|
||||
thread-safe, and this process has exactly two callers: the single
|
||||
image-vector worker thread and the backfill script's download pool.
|
||||
|
||||
Which column gets the vector, when, and for which rows is
|
||||
app/services/image_vector.py's business, not this file's.
|
||||
|
||||
FAILURE IS SILENT BY DESIGN
|
||||
---------------------------
|
||||
The runtime (`ai-edge-litert`) is imported here, not at module import, and
|
||||
the model file is opened here, not at boot. If either is missing, `available()`
|
||||
returns False after ONE warning, `embed()` returns None, and the catalog keeps
|
||||
writing rows with a NULL img_vector. A missing wheel or a model left out of an
|
||||
image must never turn into a failed product write.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import threading
|
||||
import warnings
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from app.infrastructure.settings import (
|
||||
IMAGE_EMBED_MODEL_PATH,
|
||||
IMAGE_EMBED_NUM_THREADS,
|
||||
IMAGE_VECTOR_MAX_PIXELS,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
EMBED_DIM = 1024
|
||||
IMAGE_EMBED_INPUT_SIZE = 224
|
||||
_INPUT_SHAPE = (1, IMAGE_EMBED_INPUT_SIZE, IMAGE_EMBED_INPUT_SIZE, 3)
|
||||
_OUTPUT_SHAPE = (1, EMBED_DIM)
|
||||
|
||||
_lock = threading.Lock()
|
||||
_interpreter: Any = None
|
||||
_input_index: Optional[int] = None
|
||||
_output_index: Optional[int] = None
|
||||
_disabled_reason: Optional[str] = None
|
||||
_warned = False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# bytes -> input tensor
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def preprocess(image_bytes: bytes) -> Optional[np.ndarray]:
|
||||
"""Decode `image_bytes` into the model's input tensor, or None.
|
||||
|
||||
==== PREPROCESSING CONTRACT ==============================================
|
||||
This is the DEFAULT recipe, written before the colleague's own code was
|
||||
available. When theirs arrives, replace the body of this function with a
|
||||
line-for-line port and delete this banner. The three decisions that decide
|
||||
whether two systems' vectors agree are:
|
||||
* resize method (default: BILINEAR)
|
||||
* squash vs crop (default: squash the whole image to 224x224,
|
||||
no aspect-ratio preservation, no centre crop)
|
||||
* value scaling (default: raw 0..255 as float32 - what a Keras
|
||||
MobileNetV3 export expects, since the graph
|
||||
carries its own Rescaling layer)
|
||||
Fixed in any variant: EXIF orientation applied (browsers apply it, so the
|
||||
stored vector describes what a person sees), alpha flattened on white,
|
||||
RGB channel order, NHWC layout, batch of one.
|
||||
==========================================================================
|
||||
|
||||
No `Image.draft()` here, deliberately: a DCT-downscaled JPEG decode is not
|
||||
bit-comparable with a full decode followed by a resize, and comparability
|
||||
is the entire point. Never raises - pytest runs warnings as errors, and a
|
||||
DecompressionBombWarning is one of the things the broad except absorbs.
|
||||
"""
|
||||
if not image_bytes:
|
||||
return None
|
||||
try:
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
im = Image.open(io.BytesIO(image_bytes))
|
||||
if im.width * im.height > IMAGE_VECTOR_MAX_PIXELS:
|
||||
logger.debug("image rejected: %dx%d exceeds pixel cap", im.width, im.height)
|
||||
return None
|
||||
im = ImageOps.exif_transpose(im) or im
|
||||
|
||||
if "A" in im.getbands() or "transparency" in im.info:
|
||||
rgba = im.convert("RGBA")
|
||||
white = Image.new("RGBA", rgba.size, (255, 255, 255, 255))
|
||||
im = Image.alpha_composite(white, rgba)
|
||||
im = im.convert("RGB")
|
||||
im = im.resize((IMAGE_EMBED_INPUT_SIZE, IMAGE_EMBED_INPUT_SIZE), Image.Resampling.BILINEAR)
|
||||
|
||||
arr = np.asarray(im, dtype=np.float32) # (224, 224, 3), 0..255
|
||||
arr = np.expand_dims(arr, axis=0) # (1, 224, 224, 3)
|
||||
if arr.shape != _INPUT_SHAPE:
|
||||
logger.debug("image rejected: tensor shape %s", arr.shape)
|
||||
return None
|
||||
return np.ascontiguousarray(arr)
|
||||
except Exception as exc: # noqa: BLE001 - see docstring
|
||||
logger.debug("image undecodable: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def l2_normalize(vec: np.ndarray) -> Optional[np.ndarray]:
|
||||
"""Unit-length copy of `vec`, or None for a (near-)zero vector."""
|
||||
arr = np.asarray(vec, dtype=np.float32).reshape(-1)
|
||||
norm = float(np.linalg.norm(arr))
|
||||
if not np.isfinite(norm) or norm < 1e-12:
|
||||
return None
|
||||
return arr / norm
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the interpreter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _disable(reason: str) -> None:
|
||||
"""Record why the embedder is off and say so once. Caller holds _lock."""
|
||||
global _disabled_reason, _interpreter, _warned
|
||||
_disabled_reason = reason
|
||||
_interpreter = None
|
||||
if not _warned:
|
||||
_warned = True
|
||||
logger.warning("image embedder disabled: %s (img_vector stays NULL until fixed)", reason)
|
||||
|
||||
|
||||
def _load() -> None:
|
||||
"""Create and validate the interpreter. Caller holds _lock. Never raises."""
|
||||
global _interpreter, _input_index, _output_index
|
||||
path = IMAGE_EMBED_MODEL_PATH
|
||||
if not path.is_file():
|
||||
_disable(f"model file not found at {path}")
|
||||
return
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
# The wheel's import-time deprecation chatter would become a hard
|
||||
# error under pytest's filterwarnings=error.
|
||||
warnings.simplefilter("ignore")
|
||||
from ai_edge_litert.interpreter import Interpreter
|
||||
except Exception as exc: # noqa: BLE001 - ImportError or a broken wheel
|
||||
_disable(f"ai-edge-litert is not importable ({exc})")
|
||||
return
|
||||
try:
|
||||
interp = Interpreter(model_path=str(path), num_threads=max(1, IMAGE_EMBED_NUM_THREADS))
|
||||
interp.allocate_tensors()
|
||||
inputs = interp.get_input_details()
|
||||
outputs = interp.get_output_details()
|
||||
except Exception as exc: # noqa: BLE001 - unreadable / malformed model
|
||||
_disable(f"could not load {path.name} ({exc})")
|
||||
return
|
||||
|
||||
if len(inputs) != 1 or tuple(int(d) for d in inputs[0]["shape"]) != _INPUT_SHAPE \
|
||||
or np.dtype(inputs[0]["dtype"]) != np.float32:
|
||||
got = [(list(i["shape"]), np.dtype(i["dtype"]).name) for i in inputs]
|
||||
_disable(f"{path.name} input is {got}, expected [1,224,224,3] float32")
|
||||
return
|
||||
if len(outputs) != 1 or tuple(int(d) for d in outputs[0]["shape"]) != _OUTPUT_SHAPE:
|
||||
got = [list(o["shape"]) for o in outputs]
|
||||
_disable(f"{path.name} output is {got}, expected [1,1024]")
|
||||
return
|
||||
|
||||
_interpreter = interp
|
||||
_input_index = int(inputs[0]["index"])
|
||||
_output_index = int(outputs[0]["index"])
|
||||
logger.info("image embedder ready: %s (%d thread(s))", path.name, IMAGE_EMBED_NUM_THREADS)
|
||||
|
||||
|
||||
def _ensure_loaded() -> bool:
|
||||
"""Caller holds _lock."""
|
||||
if _interpreter is None and _disabled_reason is None:
|
||||
_load()
|
||||
return _interpreter is not None
|
||||
|
||||
|
||||
def available() -> bool:
|
||||
"""True when the model can be used. Loads it on first call."""
|
||||
with _lock:
|
||||
return _ensure_loaded()
|
||||
|
||||
|
||||
def describe() -> str:
|
||||
"""One line for script banners: where the model is and whether it works."""
|
||||
with _lock:
|
||||
ok = _ensure_loaded()
|
||||
if ok:
|
||||
return f"ready - {IMAGE_EMBED_MODEL_PATH} ({IMAGE_EMBED_NUM_THREADS} thread(s), {EMBED_DIM}-d)"
|
||||
return f"disabled - {_disabled_reason}"
|
||||
|
||||
|
||||
def embed(preprocessed: np.ndarray) -> Optional[np.ndarray]:
|
||||
"""Raw model output for one preprocessed tensor, as (1024,) float32, or None.
|
||||
|
||||
Serialised on the module lock: the interpreter's tensors are shared
|
||||
state, and two threads calling set_tensor/invoke at once corrupt both.
|
||||
"""
|
||||
if preprocessed is None or tuple(preprocessed.shape) != _INPUT_SHAPE:
|
||||
return None
|
||||
with _lock:
|
||||
if not _ensure_loaded():
|
||||
return None
|
||||
try:
|
||||
_interpreter.set_tensor(_input_index, np.asarray(preprocessed, dtype=np.float32))
|
||||
_interpreter.invoke()
|
||||
out = _interpreter.get_tensor(_output_index)
|
||||
return np.array(out, dtype=np.float32, copy=True).reshape(-1)[:EMBED_DIM]
|
||||
except Exception as exc: # noqa: BLE001 - one bad tensor must not stop the worker
|
||||
logger.debug("inference failed: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def embedding_for_bytes(image_bytes: bytes) -> Optional[List[float]]:
|
||||
"""bytes -> preprocess -> embed -> L2-normalise -> 1024 floats, or None."""
|
||||
tensor = preprocess(image_bytes)
|
||||
if tensor is None:
|
||||
return None
|
||||
raw = embed(tensor)
|
||||
if raw is None or raw.shape != (EMBED_DIM,):
|
||||
return None
|
||||
unit = l2_normalize(raw)
|
||||
if unit is None:
|
||||
return None
|
||||
return [float(v) for v in unit]
|
||||
|
||||
|
||||
def _reset() -> None:
|
||||
"""Tests only: forget the interpreter and any recorded failure."""
|
||||
global _interpreter, _input_index, _output_index, _disabled_reason, _warned
|
||||
with _lock:
|
||||
_interpreter = None
|
||||
_input_index = None
|
||||
_output_index = None
|
||||
_disabled_reason = None
|
||||
_warned = False
|
||||
@@ -1,18 +1,26 @@
|
||||
"""A pixel thumbnail of every product's PRIMARY image, stored on its brand table.
|
||||
"""An image embedding of every product's PRIMARY image, stored on its brand table.
|
||||
|
||||
WHAT `img_vector` IS
|
||||
--------------------
|
||||
The image the product card shows - `image_url`, else `image_urls[0]`, which is
|
||||
exactly the order `frontend/src/components/ProductCard.jsx` walks - decoded,
|
||||
EXIF-rotated, flattened onto white, resized to 32x32 RGB and unrolled into
|
||||
3072 uint8 values. It is a literal downscaled picture, not a learned embedding:
|
||||
two rows with the same photo have (near-)identical vectors, and a row whose
|
||||
image changed has a different one. Stored as pgvector `vector(3072)` so that a
|
||||
`<=>` comparison later needs no migration; 0..255 is exact in float4.
|
||||
exactly the order `frontend/src/components/ProductCard.jsx` walks - run through
|
||||
MobileNetV3-Small (`app/services/image_embedder.py`): 1024 floats, L2-normalised,
|
||||
stored as pgvector `vector(1024)` with an hnsw cosine index. It is a learned
|
||||
embedding, not a picture: two photos of the same product from different angles
|
||||
land near each other, and `<=>` between rows is a visual-similarity search. The
|
||||
model, runtime and post-processing are the same a colleague uses for their
|
||||
product images, so our vectors and theirs are directly comparable.
|
||||
|
||||
(Until 2026-09-17 this column held a 32x32 pixel thumbnail as `vector(3072)`.
|
||||
`vector_store._ensure_img_vector_type` drops and re-creates a column of the
|
||||
wrong dimension on the next write, discarding those vectors; the backfill
|
||||
script fills the new ones.)
|
||||
|
||||
`img_vector_src` records the URL the vector was computed from. That is what
|
||||
makes recomputation idempotent: a row is due when it has no vector, or when
|
||||
its current primary URL differs from the one the vector came from.
|
||||
its current primary URL differs from the one the vector came from. A change
|
||||
to the model or its preprocessing is neither; `backfill_rows(force=True)` /
|
||||
`--force` recomputes every row for that.
|
||||
|
||||
WHY THIS MODULE, AND NOT THE UPSERT
|
||||
-----------------------------------
|
||||
@@ -31,14 +39,12 @@ and an overflowing queue simply leaves rows for
|
||||
`scripts/backfill_image_vectors.py`. Same shape as app/core/batch_worker.py,
|
||||
for the same reasons.
|
||||
|
||||
Pillow, not OpenCV: it is already a dependency, `Image.draft()` decodes a JPEG
|
||||
at 1/8 scale for a target this small, and it honours EXIF orientation on
|
||||
in-memory bytes. OpenCV would add ~130MB on disk and always decode at full
|
||||
resolution, on a container that already holds torch.
|
||||
The worker is also the only thread in the API process that runs the model,
|
||||
which matters because a TFLite interpreter is not thread-safe; the embedder
|
||||
serialises every call on its own lock regardless.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
@@ -50,16 +56,15 @@ from app.infrastructure.settings import (
|
||||
ENABLE_IMAGE_VECTORS,
|
||||
IMAGE_VECTOR_HOST_PAUSE_SECONDS,
|
||||
IMAGE_VECTOR_MAX_BYTES,
|
||||
IMAGE_VECTOR_MAX_PIXELS,
|
||||
IMAGE_VECTOR_QUEUE_MAX,
|
||||
IMAGE_VECTOR_TIMEOUT_SECONDS,
|
||||
MIN_IMAGE_BYTES,
|
||||
)
|
||||
from app.services import image_embedder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
IMG_VECTOR_SIZE = 32
|
||||
IMG_VECTOR_DIM = IMG_VECTOR_SIZE * IMG_VECTOR_SIZE * 3 # 3072
|
||||
IMG_VECTOR_DIM = image_embedder.EMBED_DIM # 1024
|
||||
IMG_VECTOR_COLUMNS = ("img_vector", "img_vector_src")
|
||||
|
||||
_IMAGE_EXTENSIONS = (".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp")
|
||||
@@ -93,54 +98,12 @@ def primary_image_url(row: Mapping[str, Any]) -> Optional[str]:
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bytes -> vector
|
||||
# ---------------------------------------------------------------------------
|
||||
# The model, its preprocessing and its lock live in image_embedder.py; this
|
||||
# module only decides which bytes go in and where the result goes.
|
||||
|
||||
def pixel_vector(image_bytes: bytes) -> Optional[List[int]]:
|
||||
"""32x32 RGB pixels of `image_bytes` as 3072 ints in 0..255, or None.
|
||||
|
||||
None for anything Pillow cannot turn into a picture: corrupt or truncated
|
||||
data, an unknown format, or a header claiming more than
|
||||
IMAGE_VECTOR_MAX_PIXELS (checked before decoding, so a decompression bomb
|
||||
costs a header read and nothing more). Never raises - pytest runs with
|
||||
warnings as errors, and a DecompressionBombWarning is one of the things
|
||||
the broad except is there to absorb.
|
||||
"""
|
||||
if not image_bytes:
|
||||
return None
|
||||
try:
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
im = Image.open(io.BytesIO(image_bytes))
|
||||
if im.width * im.height > IMAGE_VECTOR_MAX_PIXELS:
|
||||
logger.debug("image rejected: %dx%d exceeds pixel cap", im.width, im.height)
|
||||
return None
|
||||
# JPEG only: ask the decoder for a DCT-scaled decode no smaller than
|
||||
# 4x the target. A no-op for every other format.
|
||||
try:
|
||||
im.draft("RGB", (IMG_VECTOR_SIZE * 4, IMG_VECTOR_SIZE * 4))
|
||||
except Exception: # noqa: BLE001 - draft is an optimisation, not a requirement
|
||||
pass
|
||||
im = ImageOps.exif_transpose(im) or im
|
||||
|
||||
has_alpha = "A" in im.getbands() or "transparency" in im.info
|
||||
if has_alpha:
|
||||
rgba = im.convert("RGBA")
|
||||
white = Image.new("RGBA", rgba.size, (255, 255, 255, 255))
|
||||
im = Image.alpha_composite(white, rgba)
|
||||
im = im.convert("RGB")
|
||||
im = im.resize((IMG_VECTOR_SIZE, IMG_VECTOR_SIZE), Image.Resampling.BILINEAR)
|
||||
values = list(im.tobytes())
|
||||
if len(values) != IMG_VECTOR_DIM:
|
||||
logger.debug("image rejected: %d values, expected %d", len(values), IMG_VECTOR_DIM)
|
||||
return None
|
||||
return values
|
||||
except Exception as exc: # noqa: BLE001 - see docstring
|
||||
logger.debug("image undecodable: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def to_pg(vector: Iterable[int]) -> str:
|
||||
def to_pg(vector: Iterable[float]) -> str:
|
||||
"""The text form pgvector accepts, identical to what the upsert uses for `embedding`."""
|
||||
return "[" + ",".join(str(int(v)) for v in vector) + "]"
|
||||
return "[" + ",".join(str(float(v)) for v in vector) + "]"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -230,12 +193,12 @@ def download_image_bytes(url: str, timeout: Optional[float] = None) -> Optional[
|
||||
return None
|
||||
|
||||
|
||||
def vector_for_url(url: str, timeout: Optional[float] = None) -> Tuple[Optional[List[int]], str]:
|
||||
"""(vector or None, the URL it was computed from)."""
|
||||
def vector_for_url(url: str, timeout: Optional[float] = None) -> Tuple[Optional[List[float]], str]:
|
||||
"""(L2-normalised 1024-float embedding or None, the URL it was computed from)."""
|
||||
data = download_image_bytes(url, timeout=timeout)
|
||||
if data is None:
|
||||
return None, url
|
||||
return pixel_vector(data), url
|
||||
return image_embedder.embedding_for_bytes(data), url
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -264,6 +227,7 @@ def rows_needing_vectors(
|
||||
*,
|
||||
recompute_stale: bool = True,
|
||||
limit: Optional[int] = None,
|
||||
force: bool = False,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Rows that have a primary image and no current vector for it.
|
||||
|
||||
@@ -271,17 +235,27 @@ def rows_needing_vectors(
|
||||
`image_url` was rewritten by a later upsert is stale and is included when
|
||||
`recompute_stale` is on. Rows with no usable primary are dropped here, so
|
||||
the caller never spends a request on them.
|
||||
|
||||
`force` takes every row with a primary image, current or not. That is the
|
||||
switch for a model or preprocessing change: the URL did not move, so
|
||||
nothing else would notice that every stored vector is now the wrong one.
|
||||
"""
|
||||
sql = (
|
||||
f"SELECT image_id, product_name, image_url, image_urls, img_vector_src, "
|
||||
f"(img_vector IS NULL) AS missing FROM {table_name} "
|
||||
f"WHERE (img_vector IS NULL OR img_vector_src IS DISTINCT FROM "
|
||||
f"COALESCE(NULLIF(image_url, ''), image_urls[1]))"
|
||||
f"(img_vector IS NULL) AS missing FROM {table_name}"
|
||||
)
|
||||
clauses: List[str] = []
|
||||
params: List[Any] = []
|
||||
if not force:
|
||||
clauses.append(
|
||||
"(img_vector IS NULL OR img_vector_src IS DISTINCT FROM "
|
||||
"COALESCE(NULLIF(image_url, ''), image_urls[1]))"
|
||||
)
|
||||
if image_ids is not None:
|
||||
sql += " AND image_id = ANY(%s)"
|
||||
clauses.append("image_id = ANY(%s)")
|
||||
params.append(list(image_ids))
|
||||
if clauses:
|
||||
sql += " WHERE " + " AND ".join(clauses)
|
||||
sql += " ORDER BY updated_at DESC"
|
||||
if limit is not None:
|
||||
sql += " LIMIT %s"
|
||||
@@ -295,16 +269,17 @@ def rows_needing_vectors(
|
||||
primary = primary_image_url(row)
|
||||
if not primary:
|
||||
continue
|
||||
if not row.get("missing") and not recompute_stale:
|
||||
continue
|
||||
if not row.get("missing") and row.get("img_vector_src") == primary:
|
||||
continue
|
||||
if not force:
|
||||
if not row.get("missing") and not recompute_stale:
|
||||
continue
|
||||
if not row.get("missing") and row.get("img_vector_src") == primary:
|
||||
continue
|
||||
row["primary_url"] = primary
|
||||
out.append(row)
|
||||
return out
|
||||
|
||||
|
||||
def store_vector(cur, table_name: str, image_id: str, vector: List[int], src: str) -> None:
|
||||
def store_vector(cur, table_name: str, image_id: str, vector: List[float], src: str) -> None:
|
||||
"""Write one vector. Deliberately leaves `updated_at` alone: the catalog
|
||||
listing orders on it, and a background enrichment must not reshuffle it."""
|
||||
cur.execute(
|
||||
@@ -321,11 +296,15 @@ def backfill_rows(
|
||||
limit: Optional[int] = None,
|
||||
dry_run: bool = False,
|
||||
pause_seconds: Optional[float] = None,
|
||||
force: bool = False,
|
||||
) -> Dict[str, int]:
|
||||
"""Compute and store vectors for one brand, sequentially. Never raises.
|
||||
|
||||
Returns counts: candidates, computed, failed, skipped (no columns yet).
|
||||
`pause_seconds` is the gap between two requests to the same host.
|
||||
Returns counts: candidates, computed, failed, skipped (1 when the table
|
||||
has no img_vector columns yet, or the embedder is unavailable - in both
|
||||
cases nothing is downloaded). `pause_seconds` is the gap between two
|
||||
requests to the same host; `force` recomputes rows that already have a
|
||||
current vector.
|
||||
"""
|
||||
from app.services.vector_store import _connect, _table_name
|
||||
|
||||
@@ -343,8 +322,14 @@ def backfill_rows(
|
||||
logger.info("image vectors: %s has no img_vector columns yet - skipped", table_name)
|
||||
totals["skipped"] = 1
|
||||
return totals
|
||||
if not image_embedder.available():
|
||||
# available() already warned once with the reason. Rows stay
|
||||
# NULL and are picked up by the next write or the script.
|
||||
logger.info("image vectors: embedder unavailable - %s left for later", table_name)
|
||||
totals["skipped"] = 1
|
||||
return totals
|
||||
rows = rows_needing_vectors(
|
||||
cur, table_name, image_ids, recompute_stale=recompute_stale, limit=limit
|
||||
cur, table_name, image_ids, recompute_stale=recompute_stale, limit=limit, force=force
|
||||
)
|
||||
totals["candidates"] = len(rows)
|
||||
for row in rows:
|
||||
|
||||
@@ -208,13 +208,19 @@ def get_brand_table_ddl(brand: str) -> str:
|
||||
-- Vector embedding for search
|
||||
embedding vector(384),
|
||||
|
||||
-- 32x32 RGB pixel thumbnail of the PRIMARY image (image_url, else
|
||||
-- image_urls[0]) as 3072 uint8 values, and the URL it was computed
|
||||
-- from. Written ONLY by app/services/image_vector.py via targeted
|
||||
-- UPDATEs - never by upsert_brand_products, for the same reason as
|
||||
-- nutrition_score above. Not indexed: pgvector caps vector indexes
|
||||
-- at 2000 dims, and an exact scan is fine at this catalog's size.
|
||||
img_vector vector(3072),
|
||||
-- MobileNetV3-Small embedding of the PRIMARY image (image_url, else
|
||||
-- image_urls[0]): 1024 floats, L2-normalised, and the URL it was
|
||||
-- computed from. Written ONLY by app/services/image_vector.py via
|
||||
-- targeted UPDATEs - never by upsert_brand_products, for the same
|
||||
-- reason as nutrition_score above.
|
||||
--
|
||||
-- Its hnsw index is deliberately NOT declared here. This DDL runs
|
||||
-- before _ensure_columns on every write, and a table still carrying
|
||||
-- the previous vector(3072) column would make CREATE INDEX fail
|
||||
-- (2000-dim cap) and take the whole write with it. The index is
|
||||
-- created by _ensure_img_vector_type, after the column is known to
|
||||
-- be vector(1024).
|
||||
img_vector vector(1024),
|
||||
img_vector_src TEXT,
|
||||
|
||||
-- Timestamps
|
||||
@@ -257,6 +263,82 @@ def _connect() -> Optional[psycopg.Connection]:
|
||||
# fail the write - see the loop in _ensure_columns.
|
||||
_TYPE_OPTIONAL_COLUMNS = frozenset({"img_vector"})
|
||||
|
||||
# The embedder's output width. col_defs above must declare img_vector as
|
||||
# exactly vector(IMG_VECTOR_DIMS) - a test pins the two together.
|
||||
IMG_VECTOR_DIMS = 1024
|
||||
|
||||
|
||||
def _ensure_img_vector_type(cur, table_name: str) -> None:
|
||||
"""Self-heal img_vector to vector(IMG_VECTOR_DIMS) and give it an hnsw index.
|
||||
|
||||
The column first shipped as vector(3072) (a pixel thumbnail) and was
|
||||
re-specified as a 1024-float model embedding the next day. The old data
|
||||
is meaningless under the new definition, so a column of any other
|
||||
dimension is dropped and re-added, NULL, and `img_vector_src` is cleared
|
||||
so the backfill sees every row as missing. This is the one place in the
|
||||
schema code that discards data - by design, and only for this column.
|
||||
|
||||
The dimension is read from pg_attribute.atttypmod (pgvector stores the
|
||||
declared width there; -1 for a bare `vector`) with fetchall(), so the
|
||||
recording cursors the tests and the migrate dry run use - which answer []
|
||||
- simply take the "nothing to do" path.
|
||||
|
||||
Never raises: every statement is its own try, and _connect() is
|
||||
autocommit, so a refusal poisons nothing.
|
||||
"""
|
||||
try:
|
||||
cur.execute(
|
||||
"SELECT a.atttypmod FROM pg_attribute a "
|
||||
"JOIN pg_class c ON c.oid = a.attrelid "
|
||||
"JOIN pg_namespace n ON n.oid = c.relnamespace "
|
||||
"WHERE n.nspname = 'public' AND c.relname = %s "
|
||||
"AND a.attname = 'img_vector' AND NOT a.attisdropped",
|
||||
(table_name,),
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
except Exception as e: # noqa: BLE001 - a failed probe is a skipped heal, not a failed write
|
||||
logger.warning(f"Could not read img_vector dimension on {table_name}: {e}")
|
||||
return
|
||||
if not rows:
|
||||
return
|
||||
dims = rows[0][0]
|
||||
if dims != IMG_VECTOR_DIMS:
|
||||
try:
|
||||
cur.execute(f"ALTER TABLE {table_name} DROP COLUMN img_vector")
|
||||
# IF NOT EXISTS, although it cannot exist here: migrate_brand_schema
|
||||
# recognises ADD COLUMN statements by that exact phrase.
|
||||
cur.execute(
|
||||
f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS img_vector vector({IMG_VECTOR_DIMS})"
|
||||
)
|
||||
cur.execute(f"UPDATE {table_name} SET img_vector_src = NULL")
|
||||
logger.warning(
|
||||
f"Retyped {table_name}.img_vector from vector({dims}) to vector({IMG_VECTOR_DIMS}); "
|
||||
f"stored vectors discarded, backfill needed"
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 - e.g. two first writes raced; the next one re-probes
|
||||
logger.warning(f"Could not retype {table_name}.img_vector: {e}")
|
||||
return
|
||||
# Probe before CREATE INDEX IF NOT EXISTS rather than relying on it: the
|
||||
# migrate script's dry run records every non-SELECT statement, and an
|
||||
# unconditional CREATE would make every current table look like work.
|
||||
# A cursor that cannot answer (fetchall() == []) is treated as "no index".
|
||||
try:
|
||||
cur.execute(
|
||||
"SELECT 1 FROM pg_indexes WHERE schemaname = 'public' AND tablename = %s AND indexname = %s",
|
||||
(table_name, f"idx_{table_name}_img_vector"),
|
||||
)
|
||||
if cur.fetchall():
|
||||
return
|
||||
except Exception as e: # noqa: BLE001 - fall through and let CREATE IF NOT EXISTS decide
|
||||
logger.debug(f"Index probe on {table_name} failed, creating anyway: {e}")
|
||||
try:
|
||||
cur.execute(
|
||||
f"CREATE INDEX IF NOT EXISTS idx_{table_name}_img_vector "
|
||||
f"ON {table_name} USING hnsw (img_vector vector_cosine_ops)"
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 - pgvector < 0.5 has no hnsw; an exact scan still works
|
||||
logger.warning(f"hnsw index on {table_name}.img_vector skipped: {e}")
|
||||
|
||||
|
||||
def _ensure_columns(cur, table_name: str) -> None:
|
||||
"""Add missing columns and relax legacy NOT NULL constraints for smooth schema migration."""
|
||||
@@ -318,7 +400,7 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
# Written only by app/services/image_vector.py - see get_brand_table_ddl.
|
||||
# Plain string literals, always: scripts/migrate_brand_schema.py reads
|
||||
# this dict back out of the source with ast.literal_eval.
|
||||
"img_vector": "vector(3072)",
|
||||
"img_vector": "vector(1024)",
|
||||
"img_vector_src": "TEXT",
|
||||
}
|
||||
cur.execute(
|
||||
@@ -330,13 +412,14 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
existing = {row[0] for row in col_info}
|
||||
|
||||
# 1. Add missing columns
|
||||
have_img_vector = "img_vector" in existing
|
||||
for col, col_type in col_defs.items():
|
||||
if col not in existing:
|
||||
if col in _TYPE_OPTIONAL_COLUMNS:
|
||||
# vector(3072) needs pgvector >= 0.4 (the 16000-dim ceiling).
|
||||
# The production server's version is not pinned anywhere, so
|
||||
# a refusal here is logged and the write goes on without the
|
||||
# column; image_vector.table_has_columns() skips the table.
|
||||
# The vector type needs pgvector; the production server's
|
||||
# version is not pinned anywhere, so a refusal here is logged
|
||||
# and the write goes on without the column -
|
||||
# image_vector.table_has_columns() skips the table.
|
||||
# Safe to continue because _connect() is autocommit.
|
||||
try:
|
||||
cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}")
|
||||
@@ -346,11 +429,18 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
f"(pgvector too old for this type?): {e}"
|
||||
)
|
||||
continue
|
||||
if col == "img_vector":
|
||||
have_img_vector = True
|
||||
else:
|
||||
cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}")
|
||||
if col not in ("created_at", "updated_at"):
|
||||
logger.info(f"Added missing column '{col}' to {table_name}")
|
||||
|
||||
# 1b. A pre-existing img_vector of the wrong width is replaced, and the
|
||||
# column gets its index. Only when the column is actually there.
|
||||
if have_img_vector:
|
||||
_ensure_img_vector_type(cur, table_name)
|
||||
|
||||
# 2. Relax legacy NOT NULL constraints on columns not present in standard insert
|
||||
#
|
||||
# `nutrition_score` and `health_score` are listed even though the INSERT
|
||||
@@ -651,8 +741,8 @@ def upsert_brand_products(brand: str, products: List[Dict[str, Any]], cleanup: b
|
||||
# in) are simply ignored here.
|
||||
#
|
||||
# img_vector and img_vector_src are omitted for the same reason:
|
||||
# no writer reaching this statement can carry a pixel vector, so
|
||||
# naming them would only ever NULL them. app/services/image_vector.py
|
||||
# no writer reaching this statement can carry an image embedding,
|
||||
# so naming them would only ever NULL them. app/services/image_vector.py
|
||||
# owns them and is scheduled below, after the write.
|
||||
cur.executemany(
|
||||
f"""
|
||||
@@ -1682,7 +1772,7 @@ def _table_exists(cur, table_name: str) -> bool:
|
||||
# ---------------------------------------------------------------------------
|
||||
# Every product read used to be `SELECT *`, which dragged the vector(384)
|
||||
# embedding over the wire on every listing: measured at 4.7KB of a 7.2KB row
|
||||
# (see enrichment/catalog_consensus.py). img_vector is 12KB more per row, and
|
||||
# (see enrichment/catalog_consensus.py). img_vector is 4KB more per row, and
|
||||
# the "all products" browse pulls the whole catalog in one call. Nothing that
|
||||
# reads these dicts uses either column - the API mappers copy fields by name,
|
||||
# and the one exporter that looked at `embedding` (brand_sync) was already
|
||||
|
||||
@@ -64,6 +64,11 @@ lxml>=4.9.3
|
||||
python-slugify>=8.0.4
|
||||
ddgs>=9.14.4
|
||||
Pillow>=10.0.0
|
||||
# TFLite runtime for the img_vector image embedder (app/services/image_embedder.py).
|
||||
# Runtime only - no TensorFlow - about 20MB; wheels exist for the container's
|
||||
# cp311 manylinux and for cp314 Windows. Imported lazily on first use, so a
|
||||
# missing wheel means "no image vectors", never a failed boot.
|
||||
ai-edge-litert>=2.2.0
|
||||
aiofiles>=23.2.1
|
||||
|
||||
# Playwright (Python) - last-resort image-search fallback only.
|
||||
|
||||
@@ -4,11 +4,12 @@ Fill `img_vector` for the rows that already have a primary image.
|
||||
|
||||
WHY THIS EXISTS
|
||||
---------------
|
||||
New writes get their pixel vector from the worker `upsert_brand_products`
|
||||
New writes get their image embedding from the worker `upsert_brand_products`
|
||||
schedules (app/services/image_vector.py). Rows written before the column
|
||||
existed never pass through that hook, and rows whose download failed at the
|
||||
time (a CDN that 403'd, a timeout) are left NULL on purpose. This script is
|
||||
the second chance for both.
|
||||
the second chance for both - and, with `--force`, the way to recompute every
|
||||
row after the model or its preprocessing changes.
|
||||
|
||||
WHAT IT DOES, PER TABLE
|
||||
-----------------------
|
||||
@@ -19,8 +20,10 @@ WHAT IT DOES, PER TABLE
|
||||
primary. `image_vector.primary_image_url` decides what "primary" means,
|
||||
the same way the product card does.
|
||||
3. Downloads each image with a small thread pool, no more than one request
|
||||
in flight per host and a pause between them, decodes it to 32x32 RGB and
|
||||
writes the 3072 values with a targeted UPDATE. `updated_at` is untouched.
|
||||
in flight per host and a pause between them, runs it through the
|
||||
MobileNetV3 embedder (inference is serialised on the embedder's lock; the
|
||||
downloads are what run in parallel) and writes the 1024 L2-normalised
|
||||
values with a targeted UPDATE. `updated_at` is untouched.
|
||||
|
||||
USAGE
|
||||
-----
|
||||
@@ -28,10 +31,12 @@ USAGE
|
||||
python -m scripts.backfill_image_vectors --brands Amul --apply
|
||||
python -m scripts.backfill_image_vectors --all --apply
|
||||
python -m scripts.backfill_image_vectors --all --recompute-stale --apply
|
||||
python -m scripts.backfill_image_vectors --all --force --apply # model changed
|
||||
|
||||
A dry run downloads nothing; it reports how many rows would be attempted and
|
||||
shows the first few URLs. No backup is written: nothing else writes this
|
||||
column, and `--recompute-stale` repairs any row.
|
||||
A dry run downloads nothing and needs no model; it reports how many rows would
|
||||
be attempted and shows the first few URLs. `--apply` refuses to start when the
|
||||
embedder cannot load. No backup is written: nothing else writes this column,
|
||||
and `--force` rebuilds any row.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -53,6 +58,7 @@ from app.infrastructure.settings import ( # noqa: E402
|
||||
DB_NAME,
|
||||
IMAGE_VECTOR_HOST_PAUSE_SECONDS,
|
||||
)
|
||||
from app.services import image_embedder # noqa: E402
|
||||
from app.services import image_vector as iv # noqa: E402
|
||||
from app.services.vector_store import _connect # noqa: E402
|
||||
from scripts.repair_brand_images import _brand_tables # noqa: E402
|
||||
@@ -76,7 +82,7 @@ class _HostPacer:
|
||||
with self._lock:
|
||||
return self._busy.setdefault(host, threading.Lock())
|
||||
|
||||
def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[int]], str]:
|
||||
def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[float]], str]:
|
||||
host = urlparse(url).netloc.lower()
|
||||
with self._host_lock(host):
|
||||
with self._lock:
|
||||
@@ -97,6 +103,7 @@ def backfill_table(
|
||||
*,
|
||||
apply: bool,
|
||||
recompute_stale: bool,
|
||||
force: bool,
|
||||
limit: Optional[int],
|
||||
workers: int,
|
||||
timeout: Optional[float],
|
||||
@@ -110,14 +117,15 @@ def backfill_table(
|
||||
entry["skipped"] = True
|
||||
return entry
|
||||
|
||||
rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit)
|
||||
rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit, force=force)
|
||||
entry["candidates"] = len(rows)
|
||||
if not rows:
|
||||
logger.info("%s: nothing to do", table)
|
||||
return entry
|
||||
|
||||
stale = sum(1 for r in rows if not r.get("missing"))
|
||||
logger.info("%s: %d row(s) to compute (%d missing, %d stale)", table, len(rows), len(rows) - stale, stale)
|
||||
logger.info("%s: %d row(s) to compute (%d missing, %d %s)", table, len(rows), len(rows) - stale, stale,
|
||||
"forced" if force else "stale")
|
||||
if not apply:
|
||||
for r in rows[:5]:
|
||||
logger.info(" %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:90])
|
||||
@@ -125,7 +133,7 @@ def backfill_table(
|
||||
logger.info(" ... and %d more", len(rows) - 5)
|
||||
return entry
|
||||
|
||||
results: Dict[str, Tuple[Optional[List[int]], str]] = {}
|
||||
results: Dict[str, Tuple[Optional[List[float]], str]] = {}
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
futures = {
|
||||
pool.submit(pacer.fetch, r["primary_url"], timeout): r["image_id"] for r in rows
|
||||
@@ -161,6 +169,8 @@ def main() -> int:
|
||||
parser.add_argument("--limit", type=int, default=None, help="at most N rows per table")
|
||||
parser.add_argument("--recompute-stale", action="store_true",
|
||||
help="also redo rows whose vector came from a URL that is no longer the primary")
|
||||
parser.add_argument("--force", action="store_true",
|
||||
help="redo EVERY row with a primary image - after a model or preprocessing change")
|
||||
parser.add_argument("--workers", type=int, default=3, help=f"download threads (1-{MAX_WORKERS}, default 3)")
|
||||
parser.add_argument("--timeout", type=float, default=None, help="per-request timeout in seconds")
|
||||
parser.add_argument("--pause", type=float, default=IMAGE_VECTOR_HOST_PAUSE_SECONDS,
|
||||
@@ -172,6 +182,13 @@ def main() -> int:
|
||||
|
||||
logger.info("Target database: %s / %s", DB_HOST, DB_NAME)
|
||||
logger.info("Mode: %s", "APPLY - this writes" if args.apply else "DRY RUN - nothing is downloaded or written")
|
||||
if args.apply:
|
||||
# Fail before the first download, not after: a missing model would
|
||||
# otherwise cost every image fetch and write nothing.
|
||||
logger.info("Embedder: %s", image_embedder.describe())
|
||||
if not image_embedder.available():
|
||||
logger.error("The embedder cannot run - fix the above before --apply.")
|
||||
return 2
|
||||
|
||||
conn = _connect()
|
||||
if conn is None:
|
||||
@@ -186,7 +203,7 @@ def main() -> int:
|
||||
for _suffix, table in _brand_tables(cur, brands):
|
||||
report.append(backfill_table(
|
||||
cur, conn, table,
|
||||
apply=args.apply, recompute_stale=args.recompute_stale,
|
||||
apply=args.apply, recompute_stale=args.recompute_stale, force=args.force,
|
||||
limit=args.limit, workers=workers, timeout=args.timeout, pacer=pacer,
|
||||
))
|
||||
finally:
|
||||
|
||||
@@ -33,6 +33,13 @@ WHAT IT WILL NOT DO
|
||||
`barcode_last_updated` and REAL for the tax figures rather than the
|
||||
types those values look like they want. Type drift is REPORTED here,
|
||||
never silently "fixed".
|
||||
|
||||
THE ONE EXCEPTION is `img_vector`. A column of any dimension other than
|
||||
the embedder's (vector_store.IMG_VECTOR_DIMS) is DROPPED and re-added
|
||||
empty, `img_vector_src` is cleared, and the column gets its hnsw index -
|
||||
see `_ensure_img_vector_type`. That discards every stored vector on the
|
||||
table, on purpose: a vector of the wrong width is not data, and the
|
||||
backfill script rebuilds them. Those statements are shown as `~` lines.
|
||||
* It will not create a brand table that does not exist.
|
||||
* It will not touch `nutrition_facts` or any non-brand table.
|
||||
|
||||
@@ -108,13 +115,13 @@ _TYPE_ALIASES = {
|
||||
"JSONB": {"jsonb"},
|
||||
"DOUBLE PRECISION": {"double precision"},
|
||||
"vector(384)": {"USER-DEFINED"},
|
||||
"vector(3072)": {"USER-DEFINED"},
|
||||
"vector(1024)": {"USER-DEFINED"},
|
||||
}
|
||||
|
||||
# vector(3072) needs the 16000-dimension ceiling pgvector raised to in 0.4.0.
|
||||
# Below that, `_ensure_columns` logs and skips the column rather than failing
|
||||
# the write, and this script says so up front.
|
||||
_MIN_PGVECTOR_FOR_IMG_VECTOR = (0, 4, 0)
|
||||
# The hnsw index on img_vector needs pgvector 0.5.0. Below that,
|
||||
# `_ensure_img_vector_type` logs and skips the index (an exact scan still
|
||||
# works), and this script says so up front.
|
||||
_MIN_PGVECTOR_FOR_IMG_VECTOR = (0, 5, 0)
|
||||
|
||||
|
||||
def _pgvector_version(cur) -> Optional[str]:
|
||||
@@ -184,8 +191,8 @@ def main() -> int:
|
||||
pgv = _pgvector_version(cur)
|
||||
logger.info("pgvector : %s", pgv or "unknown (extension not found?)")
|
||||
if pgv and _version_tuple(pgv) < _MIN_PGVECTOR_FOR_IMG_VECTOR:
|
||||
logger.warning("pgvector %s is older than %s: img_vector vector(3072) will be "
|
||||
"skipped on this server (every other column still applies).",
|
||||
logger.warning("pgvector %s is older than %s: the hnsw index on img_vector will be "
|
||||
"skipped on this server (the column itself still applies).",
|
||||
pgv, ".".join(map(str, _MIN_PGVECTOR_FOR_IMG_VECTOR)))
|
||||
logger.info("")
|
||||
|
||||
@@ -198,6 +205,7 @@ def main() -> int:
|
||||
report: List[Dict[str, Any]] = []
|
||||
total_missing = 0
|
||||
total_drift = 0
|
||||
total_vec = 0
|
||||
|
||||
try:
|
||||
with conn.cursor() as cur:
|
||||
@@ -213,6 +221,21 @@ def main() -> int:
|
||||
_ensure_columns(recorder, table)
|
||||
adds = [s for s in recorder.statements if "ADD COLUMN" in s]
|
||||
|
||||
# The img_vector retype/index (see the docstring's one
|
||||
# exception). A DROP means the ADD that follows is a re-add,
|
||||
# not a missing column - keep it out of that list.
|
||||
vec_ops = [
|
||||
s for s in recorder.statements
|
||||
if "img_vector" in s and (
|
||||
("DROP COLUMN" in s) or s.startswith("UPDATE") or s.startswith("CREATE INDEX")
|
||||
)
|
||||
]
|
||||
if any("DROP COLUMN img_vector" in s for s in vec_ops):
|
||||
readd = [s for s in adds if "IF NOT EXISTS img_vector " in s]
|
||||
adds = [s for s in adds if s not in readd]
|
||||
vec_ops = [s for s in vec_ops if "DROP COLUMN" in s] + readd + \
|
||||
[s for s in vec_ops if "DROP COLUMN" not in s]
|
||||
|
||||
# Type drift: a column that exists but whose type is not what
|
||||
# col_defs would have created. Reported, never altered.
|
||||
drift = []
|
||||
@@ -228,12 +251,15 @@ def main() -> int:
|
||||
"table": table,
|
||||
"missing_columns": [s.split("ADD COLUMN IF NOT EXISTS ")[1] for s in adds],
|
||||
"type_drift": drift,
|
||||
"img_vector": vec_ops,
|
||||
}
|
||||
report.append(entry)
|
||||
total_missing += len(adds)
|
||||
total_drift += len(drift)
|
||||
if vec_ops:
|
||||
total_vec += 1
|
||||
|
||||
if apply and adds:
|
||||
if apply and (adds or vec_ops):
|
||||
_ensure_columns(cur, table)
|
||||
|
||||
if apply:
|
||||
@@ -246,12 +272,14 @@ def main() -> int:
|
||||
return 0
|
||||
|
||||
width = max(len(e["table"]) for e in report)
|
||||
changed = [e for e in report if e["missing_columns"] or e["type_drift"]]
|
||||
changed = [e for e in report if e["missing_columns"] or e["type_drift"] or e["img_vector"]]
|
||||
for entry in sorted(changed, key=lambda e: -len(e["missing_columns"])):
|
||||
logger.info("%-*s %d column(s) missing", width, entry["table"],
|
||||
len(entry["missing_columns"]))
|
||||
for col in entry["missing_columns"]:
|
||||
logger.info("%-*s + %s", width, "", col)
|
||||
for stmt in entry["img_vector"]:
|
||||
logger.info("%-*s ~ %s", width, "", stmt)
|
||||
for d in entry["type_drift"]:
|
||||
logger.info("%-*s ! %s is %s, col_defs declares %s (NOT changed)",
|
||||
width, "", d["column"], d["actual"], d["declared"])
|
||||
@@ -261,7 +289,11 @@ def main() -> int:
|
||||
len(report), len(report) - len(changed))
|
||||
logger.info("%d column(s) %s, %d type mismatch(es) reported",
|
||||
total_missing, "added" if apply else "would be added", total_drift)
|
||||
if not apply and total_missing:
|
||||
if total_vec:
|
||||
logger.info("%d table(s) %s img_vector retype/index (stored vectors on a retyped "
|
||||
"table are discarded - run scripts/backfill_image_vectors afterwards)",
|
||||
total_vec, "had" if apply else "need")
|
||||
if not apply and (total_missing or total_vec):
|
||||
logger.info("")
|
||||
logger.info("Re-run with --apply to write these changes.")
|
||||
return 0
|
||||
|
||||
72
tests/test_image_embedder_model.py
Normal file
72
tests/test_image_embedder_model.py
Normal file
@@ -0,0 +1,72 @@
|
||||
"""The real MobileNetV3 embedder, when its runtime and model file are present.
|
||||
|
||||
Everything in tests/test_image_vector.py runs against a fake interpreter so
|
||||
the suite never depends on a 20MB wheel or a binary that is not in git. This
|
||||
file is the one place the actual model is exercised, and it skips - not
|
||||
fails - when either piece is missing, so a checkout without the .tflite
|
||||
stays green.
|
||||
|
||||
What it pins is the contract the column relies on, not the model's opinions:
|
||||
1024 values, unit length, deterministic for identical bytes, and different
|
||||
for different pictures.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import math
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
pytest.importorskip("ai_edge_litert")
|
||||
|
||||
from app.services import image_embedder as emb # noqa: E402
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not emb.IMAGE_EMBED_MODEL_PATH.is_file(),
|
||||
reason=f"model file not present at {emb.IMAGE_EMBED_MODEL_PATH}",
|
||||
)
|
||||
|
||||
|
||||
def _png(colour) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
Image.new("RGB", (160, 120), colour).save(buf, "PNG")
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fresh():
|
||||
emb._reset()
|
||||
yield
|
||||
emb._reset()
|
||||
|
||||
|
||||
def _cosine(a, b) -> float:
|
||||
return sum(x * y for x, y in zip(a, b))
|
||||
|
||||
|
||||
def test_the_model_loads_and_reports_its_shapes():
|
||||
assert emb.available() is True
|
||||
assert emb.describe().startswith("ready")
|
||||
|
||||
|
||||
def test_an_embedding_is_1024_unit_length_floats():
|
||||
vec = emb.embedding_for_bytes(_png((200, 30, 30)))
|
||||
|
||||
assert vec is not None and len(vec) == 1024
|
||||
assert abs(math.sqrt(sum(v * v for v in vec)) - 1.0) < 1e-4
|
||||
assert all(math.isfinite(v) for v in vec)
|
||||
|
||||
|
||||
def test_identical_bytes_give_identical_vectors():
|
||||
data = _png((10, 120, 200))
|
||||
|
||||
assert emb.embedding_for_bytes(data) == emb.embedding_for_bytes(data)
|
||||
|
||||
|
||||
def test_different_pictures_give_different_vectors():
|
||||
red = emb.embedding_for_bytes(_png((220, 20, 20)))
|
||||
blue = emb.embedding_for_bytes(_png((20, 20, 220)))
|
||||
|
||||
assert red is not None and blue is not None
|
||||
assert -1.0 <= _cosine(red, blue) < 0.999
|
||||
@@ -1,19 +1,27 @@
|
||||
"""img_vector: the pixel thumbnail of each product's primary image.
|
||||
"""img_vector: the MobileNetV3 embedding of each product's primary image.
|
||||
|
||||
Four things have to hold, and each has broken independently for a sibling
|
||||
Five things have to hold, and each has broken independently for a sibling
|
||||
column before, so each gets its own tests here:
|
||||
|
||||
1. The vector is what the card shows: `image_url`, else `image_urls[0]`,
|
||||
EXIF-rotated, flattened onto white, 32x32 RGB - and anything Pillow cannot
|
||||
decode is None, never an exception (pytest runs warnings as errors).
|
||||
2. The columns exist on every brand table via `_ensure_columns`, and a server
|
||||
that refuses `vector(3072)` loses only that column, not the write.
|
||||
3. The upsert never names them (the nutrition_score rule), and the product
|
||||
1. The tensor is what the card shows: `image_url`, else `image_urls[0]`,
|
||||
EXIF-rotated, flattened onto white, RGB, 224x224 - and anything Pillow
|
||||
cannot decode is None, never an exception (pytest runs warnings as
|
||||
errors). The tests are written to survive a change of value scaling, since
|
||||
that part of `preprocess` is meant to be replaced by the colleague's code.
|
||||
2. The model is one lazily-loaded interpreter behind one lock; without the
|
||||
runtime or the file it says so once and returns None forever after.
|
||||
3. The columns exist on every brand table via `_ensure_columns`, a column of
|
||||
the wrong width is dropped and re-added, the index is created there and
|
||||
never in the DDL, and a server that refuses the type loses only that
|
||||
column, not the write.
|
||||
4. The upsert never names them (the nutrition_score rule), and the product
|
||||
readers never select them - the browse endpoint pulls the whole catalog.
|
||||
4. The hook is off the write path: it enqueues and returns, it is a no-op
|
||||
5. The hook is off the write path: it enqueues and returns, it is a no-op
|
||||
when disabled, and a full queue is a log line, not an error.
|
||||
|
||||
No database is involved anywhere; cursors are recorders.
|
||||
No database and no model file are involved anywhere; cursors are recorders
|
||||
and the interpreter is a fake. tests/test_image_embedder_model.py runs the
|
||||
real model when it is present.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -23,9 +31,11 @@ import re
|
||||
import threading
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from app.services import image_embedder as emb
|
||||
from app.services import image_vector as iv
|
||||
from app.services import vector_store
|
||||
|
||||
@@ -46,65 +56,98 @@ def _jpeg(im: Image.Image, **kw) -> bytes:
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def _pixel(vec: List[int], x: int, y: int) -> tuple:
|
||||
i = (y * iv.IMG_VECTOR_SIZE + x) * 3
|
||||
return tuple(vec[i:i + 3])
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fresh_column_cache():
|
||||
def _fresh_state():
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
emb._reset()
|
||||
yield
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
emb._reset()
|
||||
|
||||
|
||||
class FakeInterpreter:
|
||||
"""Stands in for ai_edge_litert.Interpreter: records the input, returns a
|
||||
fixed (1, 1024) output."""
|
||||
|
||||
def __init__(self, output=None):
|
||||
self.inputs: List[np.ndarray] = []
|
||||
self.output = output if output is not None else np.arange(1, 1025, dtype=np.float32).reshape(1, 1024)
|
||||
|
||||
def set_tensor(self, index, value):
|
||||
self.inputs.append(np.array(value, copy=True))
|
||||
|
||||
def invoke(self):
|
||||
pass
|
||||
|
||||
def get_tensor(self, index):
|
||||
return self.output
|
||||
|
||||
|
||||
def _install_fake(monkeypatch, output=None) -> FakeInterpreter:
|
||||
fake = FakeInterpreter(output)
|
||||
monkeypatch.setattr(emb, "_interpreter", fake)
|
||||
monkeypatch.setattr(emb, "_input_index", 0)
|
||||
monkeypatch.setattr(emb, "_output_index", 0)
|
||||
return fake
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. bytes -> vector
|
||||
# 1. bytes -> input tensor (scaling-invariant on purpose)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_solid_png_becomes_3072_values_of_that_colour():
|
||||
vec = iv.pixel_vector(_png(Image.new("RGB", (200, 120), (255, 0, 0))))
|
||||
def test_a_solid_png_becomes_a_224_tensor_dominated_by_its_colour():
|
||||
t = emb.preprocess(_png(Image.new("RGB", (200, 120), (255, 0, 0))))
|
||||
|
||||
assert vec is not None
|
||||
assert len(vec) == iv.IMG_VECTOR_DIM == 3072
|
||||
assert set(vec[0::3]) == {255}
|
||||
assert set(vec[1::3]) == {0}
|
||||
assert set(vec[2::3]) == {0}
|
||||
assert t is not None
|
||||
assert t.shape == (1, 224, 224, 3) and t.dtype == np.float32
|
||||
r, g, b = t[0, :, :, 0], t[0, :, :, 1], t[0, :, :, 2]
|
||||
assert np.all(r == t.max()) and np.all(g == t.min()) and np.all(b == t.min())
|
||||
assert t.max() > t.min()
|
||||
|
||||
|
||||
def test_a_solid_jpeg_is_within_lossy_tolerance():
|
||||
vec = iv.pixel_vector(_jpeg(Image.new("RGB", (300, 300), (40, 120, 200))))
|
||||
def test_the_default_recipe_is_raw_0_to_255():
|
||||
"""Pinned separately from the invariant tests above: this is the one
|
||||
assertion that is EXPECTED to change when the colleague's preprocessing
|
||||
replaces the default. Update it deliberately, not by accident."""
|
||||
t = emb.preprocess(_png(Image.new("RGB", (10, 10), (255, 128, 0))))
|
||||
|
||||
assert vec is not None and len(vec) == 3072
|
||||
for channel, expected in enumerate((40, 120, 200)):
|
||||
assert all(abs(v - expected) <= 3 for v in vec[channel::3])
|
||||
assert float(t[0, 0, 0, 0]) == 255.0 and float(t[0, 0, 0, 2]) == 0.0
|
||||
|
||||
|
||||
def test_a_solid_jpeg_is_uniform_within_lossy_tolerance():
|
||||
t = emb.preprocess(_jpeg(Image.new("RGB", (300, 300), (40, 120, 200))))
|
||||
|
||||
assert t is not None
|
||||
for c in range(3):
|
||||
chan = t[0, :, :, c]
|
||||
assert chan.max() - chan.min() <= 3 * (t.max() / 255.0) + 1e-6
|
||||
|
||||
|
||||
def test_exif_orientation_is_applied_because_the_browser_applies_it():
|
||||
"""Left half red, right half blue, tagged 'rotate 90 CW to display'.
|
||||
|
||||
Without the transpose, the bottom-left pixel is red (it is the left half).
|
||||
With it, the left half has become the top half and bottom-left is blue.
|
||||
Without the transpose the bottom-left pixel is red (the left half). With
|
||||
it, the left half has become the top half and bottom-left is blue.
|
||||
"""
|
||||
im = Image.new("RGB", (64, 32), (255, 0, 0))
|
||||
im.paste((0, 0, 255), (32, 0, 64, 32))
|
||||
exif = Image.Exif()
|
||||
exif[0x0112] = 6
|
||||
vec = iv.pixel_vector(_jpeg(im, exif=exif.tobytes()))
|
||||
t = emb.preprocess(_jpeg(im, exif=exif.tobytes()))
|
||||
|
||||
assert vec is not None
|
||||
r, g, b = _pixel(vec, 0, iv.IMG_VECTOR_SIZE - 1)
|
||||
assert b > 200 and r < 60, (r, g, b)
|
||||
r, g, b = _pixel(vec, 0, 0)
|
||||
assert r > 200 and b < 60, (r, g, b)
|
||||
assert t is not None
|
||||
bottom_left = t[0, 223, 0]
|
||||
top_left = t[0, 0, 0]
|
||||
assert bottom_left[2] > bottom_left[0], bottom_left # blue dominates
|
||||
assert top_left[0] > top_left[2], top_left # red dominates
|
||||
|
||||
|
||||
def test_transparency_is_flattened_onto_white_not_black():
|
||||
im = Image.new("RGBA", (50, 50), (0, 0, 0, 0)) # fully transparent
|
||||
vec = iv.pixel_vector(_png(im))
|
||||
t = emb.preprocess(_png(Image.new("RGBA", (50, 50), (0, 0, 0, 0))))
|
||||
|
||||
assert vec is not None
|
||||
assert set(vec) == {255}
|
||||
assert t is not None
|
||||
assert np.all(t == t.max()) # white: every channel at the top of the scale
|
||||
assert t.max() > 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["P", "L", "LA", "CMYK", "I;16"])
|
||||
@@ -114,10 +157,9 @@ def test_every_pillow_mode_a_product_photo_could_arrive_in_decodes(mode):
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "TIFF" if mode in ("CMYK", "I;16") else "PNG")
|
||||
|
||||
vec = iv.pixel_vector(buf.getvalue())
|
||||
t = emb.preprocess(buf.getvalue())
|
||||
|
||||
assert vec is not None and len(vec) == 3072
|
||||
assert all(0 <= v <= 255 for v in vec)
|
||||
assert t is not None and t.shape == (1, 224, 224, 3) and np.all(np.isfinite(t))
|
||||
|
||||
|
||||
def test_the_first_frame_of_an_animated_gif_is_used():
|
||||
@@ -125,12 +167,12 @@ def test_the_first_frame_of_an_animated_gif_is_used():
|
||||
buf = io.BytesIO()
|
||||
frames[0].save(buf, "GIF", save_all=True, append_images=frames[1:])
|
||||
|
||||
assert iv.pixel_vector(buf.getvalue()) is not None
|
||||
assert emb.preprocess(buf.getvalue()) is not None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("data", [b"", b"not an image", b"\x89PNG\r\n\x1a\n" + b"\x00" * 40])
|
||||
def test_undecodable_bytes_are_none_not_an_exception(data):
|
||||
assert iv.pixel_vector(data) is None
|
||||
assert emb.preprocess(data) is None
|
||||
|
||||
|
||||
def test_a_truncated_file_is_none():
|
||||
@@ -138,18 +180,113 @@ def test_a_truncated_file_is_none():
|
||||
whole = _png(noisy)
|
||||
assert len(whole) > 400, "need a file big enough to cut"
|
||||
|
||||
assert iv.pixel_vector(whole[:200]) is None
|
||||
assert emb.preprocess(whole[:200]) is None
|
||||
|
||||
|
||||
def test_a_header_claiming_too_many_pixels_is_refused_before_decoding(monkeypatch):
|
||||
monkeypatch.setattr(iv, "IMAGE_VECTOR_MAX_PIXELS", 100)
|
||||
monkeypatch.setattr(emb, "IMAGE_VECTOR_MAX_PIXELS", 100)
|
||||
|
||||
assert iv.pixel_vector(_png(Image.new("RGB", (20, 20)))) is None
|
||||
assert iv.pixel_vector(_png(Image.new("RGB", (10, 10)))) is not None
|
||||
assert emb.preprocess(_png(Image.new("RGB", (20, 20)))) is None
|
||||
assert emb.preprocess(_png(Image.new("RGB", (10, 10)))) is not None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1a. the model: one interpreter, one lock, silent when absent
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_l2_normalize_gives_a_unit_vector_and_refuses_zero():
|
||||
unit = emb.l2_normalize(np.array([3.0, 4.0], dtype=np.float32))
|
||||
assert unit is not None and abs(float(np.linalg.norm(unit)) - 1.0) < 1e-6
|
||||
assert emb.l2_normalize(np.zeros(4, dtype=np.float32)) is None
|
||||
|
||||
|
||||
def test_embedding_for_bytes_is_1024_unit_floats_from_the_model_output(monkeypatch):
|
||||
fake = _install_fake(monkeypatch)
|
||||
|
||||
vec = iv.vector_for_url # noqa: F841 - the public path below is what scripts call
|
||||
out = emb.embedding_for_bytes(_png(Image.new("RGB", (30, 30), (0, 255, 0))))
|
||||
|
||||
assert out is not None and len(out) == emb.EMBED_DIM == iv.IMG_VECTOR_DIM == 1024
|
||||
assert abs(sum(v * v for v in out) - 1.0) < 1e-4
|
||||
assert all(isinstance(v, float) for v in out)
|
||||
assert len(fake.inputs) == 1 and fake.inputs[0].shape == (1, 224, 224, 3)
|
||||
assert fake.inputs[0].dtype == np.float32
|
||||
|
||||
|
||||
def test_embed_rejects_a_tensor_of_the_wrong_shape(monkeypatch):
|
||||
fake = _install_fake(monkeypatch)
|
||||
|
||||
assert emb.embed(np.zeros((1, 32, 32, 3), dtype=np.float32)) is None
|
||||
assert emb.embed(None) is None
|
||||
assert fake.inputs == []
|
||||
|
||||
|
||||
def test_a_zero_model_output_is_none_not_a_nan_vector(monkeypatch):
|
||||
_install_fake(monkeypatch, output=np.zeros((1, 1024), dtype=np.float32))
|
||||
|
||||
assert emb.embedding_for_bytes(_png(Image.new("RGB", (8, 8)))) is None
|
||||
|
||||
|
||||
def test_a_missing_model_warns_once_and_then_stays_quiet(monkeypatch, caplog, tmp_path):
|
||||
monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "nope.tflite")
|
||||
|
||||
with caplog.at_level("WARNING", logger=emb.__name__):
|
||||
assert emb.available() is False
|
||||
assert emb.available() is False
|
||||
assert emb.embedding_for_bytes(_png(Image.new("RGB", (8, 8)))) is None
|
||||
|
||||
warnings_ = [r for r in caplog.records if r.levelname == "WARNING"]
|
||||
assert len(warnings_) == 1 and "not found" in warnings_[0].getMessage()
|
||||
assert "disabled" in emb.describe()
|
||||
|
||||
|
||||
def test_a_broken_runtime_import_disables_rather_than_raises(monkeypatch, tmp_path):
|
||||
model = tmp_path / "m.tflite"
|
||||
model.write_bytes(b"not a flatbuffer")
|
||||
monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", model)
|
||||
import builtins
|
||||
real_import = builtins.__import__
|
||||
|
||||
def no_litert(name, *a, **k):
|
||||
if name.startswith("ai_edge_litert"):
|
||||
raise ImportError("no wheel for this platform")
|
||||
return real_import(name, *a, **k)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", no_litert)
|
||||
|
||||
assert emb.available() is False
|
||||
assert "not importable" in emb.describe()
|
||||
|
||||
|
||||
def test_inference_is_serialised_on_the_module_lock(monkeypatch):
|
||||
"""Two threads embedding at once must not interleave set_tensor/invoke."""
|
||||
inside = []
|
||||
overlap = []
|
||||
|
||||
class SlowFake(FakeInterpreter):
|
||||
def invoke(self):
|
||||
inside.append(1)
|
||||
if len(inside) > 1:
|
||||
overlap.append(1)
|
||||
threading.Event().wait(0.02)
|
||||
inside.pop()
|
||||
|
||||
fake = SlowFake()
|
||||
monkeypatch.setattr(emb, "_interpreter", fake)
|
||||
monkeypatch.setattr(emb, "_input_index", 0)
|
||||
monkeypatch.setattr(emb, "_output_index", 0)
|
||||
tensor = np.zeros((1, 224, 224, 3), dtype=np.float32)
|
||||
threads = [threading.Thread(target=emb.embed, args=(tensor,)) for _ in range(6)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert not overlap and len(fake.inputs) == 6
|
||||
|
||||
|
||||
def test_to_pg_is_the_same_text_form_the_upsert_uses_for_embedding():
|
||||
assert iv.to_pg([0, 128, 255]) == "[0,128,255]"
|
||||
assert iv.to_pg([0, 0.5, 1]) == "[0.0,0.5,1.0]"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -297,7 +434,7 @@ class MigrationCursor:
|
||||
text = " ".join(str(sql).split())
|
||||
self.statements.append(text)
|
||||
if self._refuse and self._refuse in text and "ADD COLUMN" in text:
|
||||
raise RuntimeError('type "vector(3072)" does not exist')
|
||||
raise RuntimeError('type "vector(1024)" does not exist')
|
||||
|
||||
def fetchall(self):
|
||||
return []
|
||||
@@ -308,23 +445,26 @@ def test_the_migration_adds_both_columns_to_an_existing_table():
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_cadbury")
|
||||
|
||||
assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector vector(3072)" in cur.statements
|
||||
assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector vector(1024)" in cur.statements
|
||||
assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector_src TEXT" in cur.statements
|
||||
|
||||
|
||||
def test_a_server_that_refuses_the_vector_type_still_gets_every_other_column():
|
||||
cur = MigrationCursor(refuse="img_vector vector(3072)")
|
||||
cur = MigrationCursor(refuse="img_vector vector(1024)")
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_cadbury") # must not raise
|
||||
|
||||
assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector_src TEXT" in cur.statements
|
||||
assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS embedding vector(384)" in cur.statements
|
||||
# No column -> nothing to retype and nothing to index.
|
||||
assert not any("pg_attribute" in s or "CREATE INDEX IF NOT EXISTS idx_brand_cadbury_img_vector" in s
|
||||
for s in cur.statements)
|
||||
|
||||
|
||||
def test_the_create_table_ddl_carries_them_and_does_not_index_the_vector():
|
||||
ddl = vector_store.get_brand_table_ddl("Cadbury")
|
||||
|
||||
assert "img_vector vector(3072)" in ddl
|
||||
assert "img_vector vector(1024)" in ddl
|
||||
assert "img_vector_src TEXT" in ddl
|
||||
for line in ddl.splitlines():
|
||||
if "CREATE INDEX" in line:
|
||||
@@ -336,8 +476,96 @@ def test_the_migration_script_can_still_read_col_defs_as_literals():
|
||||
|
||||
declared = _declared_types()
|
||||
|
||||
assert declared["img_vector"] == "vector(3072)"
|
||||
assert declared["img_vector"] == f"vector({vector_store.IMG_VECTOR_DIMS})" == "vector(1024)"
|
||||
assert declared["img_vector_src"] == "TEXT"
|
||||
assert vector_store.IMG_VECTOR_DIMS == emb.EMBED_DIM
|
||||
|
||||
|
||||
class RetypeCursor(MigrationCursor):
|
||||
"""A table that already HAS img_vector, of a given width."""
|
||||
|
||||
def __init__(self, dims, indexed=False):
|
||||
super().__init__()
|
||||
self._dims = dims
|
||||
self._indexed = indexed
|
||||
|
||||
def fetchall(self):
|
||||
last = self.statements[-1] if self.statements else ""
|
||||
if "pg_attribute" in last:
|
||||
return [(self._dims,)]
|
||||
if "pg_indexes" in last:
|
||||
return [(1,)] if self._indexed else []
|
||||
if "information_schema.columns" in last:
|
||||
return [("img_vector", "YES", None), ("img_vector_src", "YES", None)]
|
||||
return []
|
||||
|
||||
|
||||
def test_a_column_of_the_old_width_is_dropped_readded_and_indexed():
|
||||
cur = RetypeCursor(3072)
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_x")
|
||||
|
||||
wanted = [
|
||||
"ALTER TABLE brand_x DROP COLUMN img_vector",
|
||||
"ALTER TABLE brand_x ADD COLUMN IF NOT EXISTS img_vector vector(1024)",
|
||||
"UPDATE brand_x SET img_vector_src = NULL",
|
||||
"CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector ON brand_x USING hnsw (img_vector vector_cosine_ops)",
|
||||
]
|
||||
positions = [cur.statements.index(w) for w in wanted]
|
||||
assert positions == sorted(positions), cur.statements
|
||||
|
||||
|
||||
def test_a_column_of_the_right_width_only_gets_its_index():
|
||||
cur = RetypeCursor(1024)
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_x")
|
||||
|
||||
assert not any("DROP COLUMN" in s or s.startswith("UPDATE") for s in cur.statements)
|
||||
assert "CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector ON brand_x USING hnsw (img_vector vector_cosine_ops)" in cur.statements
|
||||
|
||||
|
||||
def test_a_table_that_is_already_current_emits_no_statement_at_all():
|
||||
"""What the migrate dry run sees after --apply: nothing to report."""
|
||||
cur = RetypeCursor(1024, indexed=True)
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_x")
|
||||
|
||||
assert not any("img_vector" in s and not s.startswith("SELECT") for s in cur.statements), cur.statements
|
||||
|
||||
|
||||
def test_a_cursor_that_cannot_answer_the_width_probe_changes_nothing():
|
||||
"""The migrate dry run and the other schema tests use recorders whose
|
||||
fetchall() is []. That must read as 'nothing to do', never as 'retype'."""
|
||||
cur = MigrationCursor()
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_x")
|
||||
|
||||
assert not any("DROP COLUMN" in s or "CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector" in s
|
||||
for s in cur.statements)
|
||||
|
||||
|
||||
def test_a_refused_index_does_not_fail_the_write():
|
||||
class NoHnsw(RetypeCursor):
|
||||
def execute(self, sql, params=None):
|
||||
super().execute(sql, params)
|
||||
if "USING hnsw" in self.statements[-1]:
|
||||
raise RuntimeError("access method hnsw does not exist")
|
||||
|
||||
cur = NoHnsw(1024)
|
||||
|
||||
vector_store._ensure_columns(cur, "brand_x") # must not raise
|
||||
|
||||
assert any("USING hnsw" in s for s in cur.statements)
|
||||
|
||||
|
||||
def test_the_retype_never_reaches_the_ddl():
|
||||
"""The DDL runs BEFORE _ensure_columns on every write; an index there
|
||||
would hit a still-3072 column and fail the write."""
|
||||
import inspect
|
||||
statements = [line for line in vector_store.get_brand_table_ddl("Cadbury").splitlines()
|
||||
if not line.strip().startswith("--")]
|
||||
assert not any("hnsw" in line for line in statements)
|
||||
assert "DROP COLUMN" not in inspect.getsource(vector_store.get_brand_table_ddl)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -616,7 +844,7 @@ def test_store_vector_does_not_touch_updated_at():
|
||||
|
||||
text, params = cur.statements[-1]
|
||||
assert text == "UPDATE brand_a SET img_vector = %s::vector, img_vector_src = %s WHERE image_id = %s"
|
||||
assert params == ("[1,2,3]", "https://a/x.jpg", "x")
|
||||
assert params == ("[1.0,2.0,3.0]", "https://a/x.jpg", "x")
|
||||
assert "updated_at" not in text
|
||||
|
||||
|
||||
@@ -634,8 +862,9 @@ class _RowsConn:
|
||||
def test_backfill_dry_run_downloads_and_writes_nothing(monkeypatch):
|
||||
cur = RowsCursor(_ROWS)
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur))
|
||||
monkeypatch.setattr(emb, "available", lambda: True)
|
||||
fetched: List[str] = []
|
||||
monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: (fetched.append(url) or [0] * 3072, url))
|
||||
monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: (fetched.append(url) or [0.0] * 1024, url))
|
||||
|
||||
totals = iv.backfill_rows("A", dry_run=True, pause_seconds=0)
|
||||
|
||||
@@ -646,9 +875,10 @@ def test_backfill_dry_run_downloads_and_writes_nothing(monkeypatch):
|
||||
def test_backfill_apply_writes_one_update_per_success_and_counts_failures(monkeypatch):
|
||||
cur = RowsCursor(_ROWS)
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur))
|
||||
monkeypatch.setattr(emb, "available", lambda: True)
|
||||
monkeypatch.setattr(
|
||||
iv, "vector_for_url",
|
||||
lambda url, timeout=None: (None, url) if "4.jpg" in url else ([7] * 3072, url),
|
||||
lambda url, timeout=None: (None, url) if "4.jpg" in url else ([0.03125] * 1024, url),
|
||||
)
|
||||
|
||||
totals = iv.backfill_rows("A", pause_seconds=0)
|
||||
@@ -657,6 +887,29 @@ def test_backfill_apply_writes_one_update_per_success_and_counts_failures(monkey
|
||||
assert totals == {"candidates": 3, "computed": 2, "failed": 1, "skipped": 0}
|
||||
assert [p[2] for _, p in updates] == ["missing", "stale"]
|
||||
assert updates[1][1][1] == "https://a/2-new.jpg"
|
||||
assert updates[0][1][0].startswith("[0.03125,0.03125,")
|
||||
|
||||
|
||||
def test_backfill_downloads_nothing_when_the_embedder_is_unavailable(monkeypatch):
|
||||
cur = RowsCursor(_ROWS)
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur))
|
||||
monkeypatch.setattr(emb, "available", lambda: False)
|
||||
monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: pytest.fail("must not download"))
|
||||
|
||||
totals = iv.backfill_rows("A", pause_seconds=0)
|
||||
|
||||
assert totals["skipped"] == 1 and totals["candidates"] == 0
|
||||
assert not any(t.startswith("SELECT image_id") for t, _ in cur.statements)
|
||||
|
||||
|
||||
def test_force_takes_every_row_with_a_primary_even_the_current_ones():
|
||||
cur = RowsCursor(_ROWS)
|
||||
|
||||
got = iv.rows_needing_vectors(cur, "brand_a", force=True)
|
||||
|
||||
assert [r["image_id"] for r in got] == ["missing", "stale", "fresh", "gallery-only"]
|
||||
text, _ = cur.statements[-1]
|
||||
assert "WHERE" not in text
|
||||
|
||||
|
||||
def test_backfill_skips_a_table_the_migration_has_not_reached(monkeypatch):
|
||||
|
||||
Reference in New Issue
Block a user