283 lines
12 KiB
Python
283 lines
12 KiB
Python
#!/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())
|