image vector dimensionality reduction

This commit is contained in:
sriram
2026-09-17 14:21:48 +05:30
parent deae694a1f
commit afa0bfa743
11 changed files with 907 additions and 183 deletions

View File

@@ -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
View File

@@ -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.

View File

@@ -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

View 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

View File

@@ -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:

View File

@@ -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

View File

@@ -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.

View File

@@ -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:

View File

@@ -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

View 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

View File

@@ -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):