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

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