Updates on Image search using vectors
This commit is contained in:
282
scripts/eval_identify.py
Normal file
282
scripts/eval_identify.py
Normal file
@@ -0,0 +1,282 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
How often does a photo of a product card come back as that product?
|
||||
|
||||
WHY THIS EXISTS
|
||||
---------------
|
||||
Colleagues photograph catalogue cards with the Nearle app, and two came back
|
||||
as other brands' products. The thresholds that decide when an image match is
|
||||
"confirmed" (IMAGE_IDENTIFY_MIN_IMAGE_SCORE, IMAGE_SEARCH_MIN_MARGIN) were
|
||||
starting guesses. This measures them, on the database this process points at.
|
||||
|
||||
WHAT IT MEASURES
|
||||
----------------
|
||||
For N sampled catalogue products it downloads the image the card shows and
|
||||
builds four synthetic "photos" of it:
|
||||
|
||||
clean the catalogue image itself (the ceiling: should be 100%)
|
||||
card the image on a white card with text lines under it
|
||||
screen that card dimmed, desaturated, blurred, blue-shifted, JPEG'd
|
||||
tight the screen shot cropped to the pack (a framing guide)
|
||||
|
||||
and, with --photos DIR, real photos named `<brand_table>__<image_id>.jpg`
|
||||
(e.g. `brand_cadbury__cadbury_dairy_milk_lickables_20g.jpg`).
|
||||
|
||||
Each photo is searched three ways:
|
||||
|
||||
image search_by_vector, no text - what the app gets today
|
||||
image+text the identify ladder with the product name as the label
|
||||
(what the card's title reads as)
|
||||
ladder+ocr (real photos only) the identify ladder with server OCR
|
||||
|
||||
and scored: top-1 exact, top-1 "same photo" (a pack-size sibling sharing
|
||||
the target's picture - right product, maybe wrong size), top-3 exact, and
|
||||
the one that matters most - WRONG BUT CONFIRMED: a different product shown
|
||||
with match_confidence "confirmed". Tune the thresholds until that is zero
|
||||
without making everything "low".
|
||||
|
||||
USAGE
|
||||
-----
|
||||
python -m scripts.eval_identify --n 40
|
||||
python -m scripts.eval_identify --n 40 --brands Cadbury,Aachi,Sakthi
|
||||
python -m scripts.eval_identify --photos ~/card_photos --n 0
|
||||
python -m scripts.eval_identify --n 40 --min-margin 0.08 --min-image-score 0.65
|
||||
|
||||
Read-only. Downloads catalogue images (honours the downloader's pacing).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import io
|
||||
import random
|
||||
import sys
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
Key = Tuple[str, str] # (brand_table, image_id)
|
||||
MODES = ("clean", "card", "screen", "tight")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# catalogue
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def load_catalogue(brands: Optional[List[str]]) -> Dict[Key, Dict[str, Any]]:
|
||||
"""Every row with an image vector: (table, image_id) -> name and source URL."""
|
||||
from app.services import vector_store as vs
|
||||
|
||||
conn = vs._connect()
|
||||
if conn is None:
|
||||
raise SystemExit("no database connection")
|
||||
wanted = {vs._table_name(b) for b in brands} if brands else None
|
||||
out: Dict[Key, Dict[str, Any]] = {}
|
||||
try:
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SET TRANSACTION READ ONLY")
|
||||
for suffix in vs._list_brand_table_suffixes(cur):
|
||||
table = f"brand_{suffix}"
|
||||
if wanted is not None and table not in wanted:
|
||||
continue
|
||||
try:
|
||||
cur.execute(f"SELECT image_id, product_name, img_vector_src FROM {table} "
|
||||
f"WHERE img_vector IS NOT NULL AND img_vector_src IS NOT NULL")
|
||||
except Exception: # noqa: BLE001 - a table without the column
|
||||
conn.rollback()
|
||||
cur.execute("SET TRANSACTION READ ONLY")
|
||||
continue
|
||||
for image_id, name, src in cur.fetchall():
|
||||
out[(table, str(image_id))] = {"name": name, "src": src}
|
||||
finally:
|
||||
conn.rollback()
|
||||
conn.close()
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# synthetic photos
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def synthetic_photo(raw: bytes, mode: str) -> Optional[bytes]:
|
||||
import numpy as np
|
||||
from PIL import Image, ImageDraw, ImageEnhance, ImageFilter
|
||||
|
||||
try:
|
||||
img = Image.open(io.BytesIO(raw)).convert("RGB")
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if mode != "clean":
|
||||
w, h = img.size
|
||||
side = max(w, h)
|
||||
pack_h = int(h * side / w)
|
||||
canvas = Image.new("RGB", (int(side * 1.3), int(side * 0.1) + pack_h + int(side * 0.9)), (255, 255, 255))
|
||||
canvas.paste(img.resize((side, pack_h)), (int(side * .15), int(side * .1)))
|
||||
draw = ImageDraw.Draw(canvas)
|
||||
for k in range(6):
|
||||
y = int(side * .1) + pack_h + int(side * .1) + k * int(side * .1)
|
||||
draw.rectangle([int(side * .15), y, int(side * (0.6 + .1 * (k % 3))), y + int(side * .04)],
|
||||
fill=(40, 40, 40))
|
||||
img = canvas
|
||||
if mode in ("screen", "tight"):
|
||||
img = ImageEnhance.Brightness(img).enhance(0.85)
|
||||
img = ImageEnhance.Color(img).enhance(0.8)
|
||||
img = img.filter(ImageFilter.GaussianBlur(2))
|
||||
px = np.array(img).astype(np.float32)
|
||||
px[..., 2] *= 1.08
|
||||
img = Image.fromarray(np.clip(px, 0, 255).astype(np.uint8))
|
||||
if mode == "tight":
|
||||
img = img.crop((int(side * .15), int(side * .1), int(side * 1.15), int(side * .1) + pack_h))
|
||||
buf = io.BytesIO()
|
||||
img.save(buf, "JPEG", quality=70)
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# scoring
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class Tally:
|
||||
def __init__(self) -> None:
|
||||
self.n = 0
|
||||
self.top1 = 0
|
||||
self.top1_same_photo = 0
|
||||
self.top3 = 0
|
||||
self.confirmed = 0
|
||||
self.wrong_confirmed = 0
|
||||
self.misses: List[str] = []
|
||||
|
||||
def add(self, target: Key, rows: List[Dict[str, Any]], confidence: str,
|
||||
catalogue: Dict[Key, Dict[str, Any]]) -> None:
|
||||
self.n += 1
|
||||
keys = [(str(r.get("brand_table") or ""), str(r.get("image_id") or "")) for r in rows]
|
||||
src = catalogue[target]["src"]
|
||||
hit1 = bool(keys) and keys[0] == target
|
||||
same_photo = bool(keys) and (hit1 or catalogue.get(keys[0], {}).get("src") == src)
|
||||
self.top1 += hit1
|
||||
self.top1_same_photo += same_photo
|
||||
self.top3 += target in keys[:3]
|
||||
if confidence == "confirmed":
|
||||
self.confirmed += 1
|
||||
if not same_photo:
|
||||
self.wrong_confirmed += 1
|
||||
if not same_photo and len(self.misses) < 8:
|
||||
got = rows[0].get("product_name") if rows else None
|
||||
self.misses.append(f"{catalogue[target]['name']} -> {got} ({confidence})")
|
||||
|
||||
def line(self, label: str) -> str:
|
||||
if not self.n:
|
||||
return f" {label:22s} (no photos)"
|
||||
pct = lambda k: f"{100.0 * k / self.n:5.1f}%" # noqa: E731
|
||||
return (f" {label:22s} n={self.n:3d} top1 {pct(self.top1)} same-photo {pct(self.top1_same_photo)} "
|
||||
f"top3 {pct(self.top3)} confirmed {pct(self.confirmed)} "
|
||||
f"WRONG+confirmed {self.wrong_confirmed}")
|
||||
|
||||
|
||||
def _ladder_confidence(result) -> str:
|
||||
from app.services.capture_discovery import is_confirmed
|
||||
|
||||
if not result.search.rows:
|
||||
return "none"
|
||||
return "confirmed" if is_confirmed(result.matched_by, result.fallback_reason) else "low"
|
||||
|
||||
|
||||
def evaluate(photos: List[Tuple[str, Key, bytes, bool]], catalogue: Dict[Key, Dict[str, Any]],
|
||||
top_k: int) -> Dict[Tuple[str, str], Tally]:
|
||||
from app.services import image_embedder
|
||||
from app.services.image_match import search_by_vector
|
||||
from app.services import product_identify
|
||||
from app.services.product_identify import identify_product
|
||||
|
||||
floor = product_identify.IMAGE_IDENTIFY_MIN_IMAGE_SCORE
|
||||
tallies: Dict[Tuple[str, str], Tally] = defaultdict(Tally)
|
||||
for mode, target, data, real in photos:
|
||||
vector = image_embedder.embedding_for_bytes(data)
|
||||
if vector is None:
|
||||
continue
|
||||
image = search_by_vector(vector, top_k=top_k)
|
||||
tallies[(mode, "image")].add(target, image.rows, image.match_confidence, catalogue)
|
||||
|
||||
label = catalogue[target]["name"]
|
||||
labelled = identify_product(vector=vector, image_bytes=None, text=label, top_k=top_k,
|
||||
min_image_score=floor)
|
||||
tallies[(mode, "image+text")].add(target, labelled.search.rows, _ladder_confidence(labelled), catalogue)
|
||||
|
||||
if real:
|
||||
ocr = identify_product(vector=vector, image_bytes=data, top_k=top_k, min_image_score=floor)
|
||||
tallies[(mode, "ladder+ocr")].add(target, ocr.search.rows, _ladder_confidence(ocr), catalogue)
|
||||
return tallies
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
parser.add_argument("--n", type=int, default=30, help="catalogue products to sample (0 = none)")
|
||||
parser.add_argument("--brands", help="comma-separated brands to sample from (default: all)")
|
||||
parser.add_argument("--photos", type=Path, help="folder of real photos named <brand_table>__<image_id>.jpg")
|
||||
parser.add_argument("--modes", default=",".join(MODES), help=f"synthetic modes (default: {','.join(MODES)})")
|
||||
parser.add_argument("--seed", type=int, default=7)
|
||||
parser.add_argument("--top-k", type=int, default=5)
|
||||
parser.add_argument("--min-margin", type=float, help="override IMAGE_SEARCH_MIN_MARGIN for this run")
|
||||
parser.add_argument("--min-image-score", type=float, help="override IMAGE_IDENTIFY_MIN_IMAGE_SCORE")
|
||||
args = parser.parse_args()
|
||||
|
||||
from app.services import image_embedder, image_match, product_identify
|
||||
from app.services.image_vector import download_image_bytes
|
||||
|
||||
# Both modules read these at call time; identify_product's own floor is
|
||||
# passed explicitly in evaluate().
|
||||
if args.min_margin is not None:
|
||||
image_match.IMAGE_SEARCH_MIN_MARGIN = args.min_margin
|
||||
product_identify.IMAGE_SEARCH_MIN_MARGIN = args.min_margin
|
||||
if args.min_image_score is not None:
|
||||
image_match.IMAGE_IDENTIFY_MIN_IMAGE_SCORE = args.min_image_score
|
||||
product_identify.IMAGE_IDENTIFY_MIN_IMAGE_SCORE = args.min_image_score
|
||||
if not image_embedder.available():
|
||||
raise SystemExit("the image embedding model is not available here")
|
||||
|
||||
brands = [b.strip() for b in args.brands.split(",")] if args.brands else None
|
||||
catalogue = load_catalogue(brands)
|
||||
print(f"catalogue: {len(catalogue)} rows with an image vector"
|
||||
f"{' in ' + ', '.join(brands) if brands else ''}")
|
||||
|
||||
photos: List[Tuple[str, Key, bytes, bool]] = []
|
||||
modes = [m for m in args.modes.split(",") if m in MODES]
|
||||
sample = random.Random(args.seed).sample(sorted(catalogue), min(args.n, len(catalogue)))
|
||||
for target in sample:
|
||||
raw = download_image_bytes(catalogue[target]["src"])
|
||||
if not raw:
|
||||
print(f" skip (download failed): {catalogue[target]['name']}")
|
||||
continue
|
||||
for mode in modes:
|
||||
data = synthetic_photo(raw, mode)
|
||||
if data:
|
||||
photos.append((mode, target, data, False))
|
||||
|
||||
if args.photos:
|
||||
for path in sorted(args.photos.iterdir()):
|
||||
if "__" not in path.stem:
|
||||
continue
|
||||
table, image_id = path.stem.split("__", 1)
|
||||
if (table, image_id) not in catalogue:
|
||||
print(f" skip (not in catalogue): {path.name}")
|
||||
continue
|
||||
photos.append(("real", (table, image_id), path.read_bytes(), True))
|
||||
|
||||
tallies = evaluate(photos, catalogue, args.top_k)
|
||||
print(f"\nthresholds: min_image_score={product_identify.IMAGE_IDENTIFY_MIN_IMAGE_SCORE} "
|
||||
f"min_margin={product_identify.IMAGE_SEARCH_MIN_MARGIN}\n")
|
||||
for mode in modes + ["real"]:
|
||||
for way in ("image", "image+text", "ladder+ocr"):
|
||||
tally = tallies.get((mode, way))
|
||||
if tally and tally.n:
|
||||
print(tally.line(f"{mode} / {way}"))
|
||||
print("\nexamples of misses (target -> top-1):")
|
||||
for (mode, way), tally in sorted(tallies.items()):
|
||||
for miss in tally.misses[:3]:
|
||||
print(f" [{mode} / {way}] {miss}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
173
scripts/replay_image_query.py
Normal file
173
scripts/replay_image_query.py
Normal file
@@ -0,0 +1,173 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Replay a saved search-by-photo request, and check the app's model against ours.
|
||||
|
||||
WHY THIS EXISTS
|
||||
---------------
|
||||
A colleague photographs a product card and the app shows a different product
|
||||
("Cadbury Dairy Milk Lickables" came back as "Milk Toned", "Aachi Sambar
|
||||
Powder 100g" as Sakthi's "Sambar powder 50g"). On the server, synthetic
|
||||
photos of those cards rank the right product first, so the question is what
|
||||
the app sent. With IMAGE_SEARCH_CAPTURE_DIR set, every request is saved as
|
||||
JSON (app/services/image_search_log.py); this script replays one.
|
||||
|
||||
WHAT IT PRINTS
|
||||
--------------
|
||||
1. The request: route, text, brand, the vector's fingerprint and norm, and
|
||||
what the server answered at the time.
|
||||
2. The image-only ranking (search_by_vector) and the identify ladder
|
||||
(identify_product) for the same vector and text, against the database
|
||||
this process points at. Both, whichever route was called, so you can see
|
||||
whether sending text would have fixed it.
|
||||
3. With --photo: the photo embedded HERE, with our model and preprocessing,
|
||||
and the cosine between that and the app's vector. The same photo through
|
||||
the same model gives >= 0.99. Well below that, the app's on-device model
|
||||
or preprocessing is not ours, and every image score it sends is in a
|
||||
different space: put the app's .tflite in app/services/models/mobilenet/
|
||||
and run `python -m scripts.backfill_image_vectors --all --force --apply`.
|
||||
A saved /identify or /image request carries its photo; use --photo for
|
||||
the vector route, with the picture the colleague took.
|
||||
|
||||
USAGE
|
||||
-----
|
||||
python -m scripts.replay_image_query data/image_search_requests/20260928T101500_ab12cd.json
|
||||
python -m scripts.replay_image_query saved.json --photo lickables_card.jpg
|
||||
python -m scripts.replay_image_query --list data/image_search_requests
|
||||
|
||||
Read-only: nothing here writes.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Sequence
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.services.image_search_log import vector_fingerprint # noqa: E402
|
||||
|
||||
PARITY_OK = 0.99
|
||||
|
||||
|
||||
def _cosine(a: Sequence[float], b: Sequence[float]) -> float:
|
||||
dot = sum(float(x) * float(y) for x, y in zip(a, b))
|
||||
na = math.sqrt(sum(float(x) ** 2 for x in a))
|
||||
nb = math.sqrt(sum(float(y) ** 2 for y in b))
|
||||
return dot / (na * nb) if na and nb else 0.0
|
||||
|
||||
|
||||
def _print_rows(title: str, rows: List[Dict[str, Any]], limit: int = 5) -> None:
|
||||
print(f" {title}")
|
||||
if not rows:
|
||||
print(" (nothing)")
|
||||
for i, row in enumerate(rows[:limit], start=1):
|
||||
score = row.get("score")
|
||||
score_s = f"{float(score):.4f}" if score is not None else " - "
|
||||
print(f" {i}. {score_s} {str(row.get('brand') or row.get('brand_table') or ''):22.22s} "
|
||||
f"{row.get('product_name') or row.get('title')} (overlap {row.get('text_overlap', 0)})")
|
||||
|
||||
|
||||
def _list(folder: Path) -> int:
|
||||
for path in sorted(folder.glob("*.json")):
|
||||
try:
|
||||
saved = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
req, ans = saved.get("request", {}), saved.get("answer", {})
|
||||
top = (ans.get("top") or [{}])[0]
|
||||
print(f"{path.name} {saved.get('route'):16s} text={str(req.get('text'))[:40]!r:44s} "
|
||||
f"-> {top.get('product_name')} ({top.get('score')}) [{ans.get('match_confidence')}]")
|
||||
return 0
|
||||
|
||||
|
||||
def replay(path: Path, photo: Optional[Path], top_k: int) -> int:
|
||||
from app.services import image_embedder
|
||||
from app.services.image_match import search_by_vector
|
||||
from app.services.product_identify import identify_product
|
||||
|
||||
saved = json.loads(path.read_text(encoding="utf-8"))
|
||||
req = saved.get("request") or {}
|
||||
answer = saved.get("answer") or {}
|
||||
vector = req.get("vector")
|
||||
text = req.get("text")
|
||||
|
||||
print(f"== {path.name} route={saved.get('route')} received={saved.get('received_at')}")
|
||||
print(f" text={text!r} brand={req.get('brand')!r} category={req.get('category')!r} "
|
||||
f"text_fallback={req.get('text_fallback')!r}")
|
||||
if vector is not None:
|
||||
norm = math.sqrt(sum(v * v for v in vector))
|
||||
print(f" vector: {len(vector)} floats, norm {norm:.4f}, fingerprint {vector_fingerprint(vector)}")
|
||||
else:
|
||||
print(" vector: none (the server had no model when this was sent)")
|
||||
print(f" answered then: matched_by={answer.get('matched_by')} reason={answer.get('fallback_reason')} "
|
||||
f"confidence={answer.get('match_confidence')} margin={answer.get('margin')}")
|
||||
_print_rows("top then:", answer.get("top") or [])
|
||||
|
||||
photo_bytes: Optional[bytes] = None
|
||||
if photo is not None:
|
||||
photo_bytes = photo.read_bytes()
|
||||
elif saved.get("photo"):
|
||||
candidate = path.parent / saved["photo"]
|
||||
if candidate.exists():
|
||||
photo_bytes = candidate.read_bytes()
|
||||
|
||||
ours: Optional[List[float]] = None
|
||||
if photo_bytes is not None:
|
||||
if not image_embedder.available():
|
||||
print("\n (no embedding model here - cannot run the parity check)")
|
||||
else:
|
||||
ours = image_embedder.embedding_for_bytes(photo_bytes)
|
||||
if ours is None:
|
||||
print("\n (the photo could not be decoded)")
|
||||
elif vector is not None:
|
||||
cos = _cosine(vector, ours)
|
||||
verdict = "SAME model and preprocessing" if cos >= PARITY_OK else (
|
||||
"DIFFERENT - the app's vectors are not comparable with img_vector")
|
||||
print(f"\n model parity: cosine(app vector, server vector of the same photo) = {cos:.4f}"
|
||||
f" -> {verdict}")
|
||||
|
||||
query = vector if vector is not None else ours
|
||||
if query is None:
|
||||
print("\n nothing to replay: no vector in the request and none computed from a photo")
|
||||
return 1
|
||||
|
||||
print("\n-- replayed now, against this process's database --")
|
||||
image = search_by_vector(query, text=text, brand=req.get("brand"), category=req.get("category"),
|
||||
top_k=top_k)
|
||||
print(f" image only: confidence={image.match_confidence} margin={image.margin} "
|
||||
f"detected_brand={image.detected_brand} scoped={image.scoped_to_brand}")
|
||||
_print_rows("ranking:", image.rows)
|
||||
|
||||
identified = identify_product(vector=query, image_bytes=photo_bytes, text=text, brand=req.get("brand"),
|
||||
category=req.get("category"), top_k=top_k)
|
||||
print(f"\n identify ladder: matched_by={identified.matched_by} reason={identified.fallback_reason} "
|
||||
f"ocr_source={identified.ocr_source} ocr_text={identified.ocr_text!r}")
|
||||
_print_rows("ranking:", identified.search.rows)
|
||||
|
||||
if ours is not None and vector is not None and _cosine(vector, ours) < PARITY_OK:
|
||||
print("\n and with the SERVER's vector of the same photo:")
|
||||
_print_rows("ranking:", search_by_vector(ours, text=text, brand=req.get("brand"),
|
||||
category=req.get("category"), top_k=top_k).rows)
|
||||
return 0
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
parser.add_argument("request", nargs="?", help="a saved request JSON")
|
||||
parser.add_argument("--photo", type=Path, help="the photo the request was made from (parity check)")
|
||||
parser.add_argument("--list", type=Path, metavar="DIR", help="summarise every saved request in DIR")
|
||||
parser.add_argument("--top-k", type=int, default=5)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.list:
|
||||
return _list(args.list)
|
||||
if not args.request:
|
||||
parser.error("give a saved request JSON, or --list DIR")
|
||||
return replay(Path(args.request), args.photo, args.top_k)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user