image vector dimensionality reduction
This commit is contained in:
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
|
||||
|
||||
Reference in New Issue
Block a user