Image vector embedding

This commit is contained in:
sriram
2026-09-16 16:33:39 +05:30
parent ce4fa70dee
commit 9c8dbf1759
10 changed files with 1493 additions and 17 deletions

View File

@@ -0,0 +1,218 @@
#!/usr/bin/env python3
"""
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`
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.
WHAT IT DOES, PER TABLE
-----------------------
1. Skips the table if the migration has not reached it yet
(`python -m scripts.migrate_brand_schema --apply` first).
2. Selects rows with a usable primary image and no vector - or, with
`--recompute-stale`, a vector computed from a URL that is no longer the
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.
USAGE
-----
python -m scripts.backfill_image_vectors --brands Amul # dry run
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
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.
"""
from __future__ import annotations
import argparse
import json
import logging
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from urllib.parse import urlparse
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from app.infrastructure.settings import ( # noqa: E402
DB_HOST,
DB_NAME,
IMAGE_VECTOR_HOST_PAUSE_SECONDS,
)
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
logging.basicConfig(level=logging.INFO, format="%(message)s")
logger = logging.getLogger("backfill_image_vectors")
MAX_WORKERS = 6
class _HostPacer:
"""At most one request per host every `pause` seconds, across threads."""
def __init__(self, pause: float):
self._pause = pause
self._lock = threading.Lock()
self._next_ok: Dict[str, float] = {}
self._busy: Dict[str, threading.Lock] = {}
def _host_lock(self, host: str) -> threading.Lock:
with self._lock:
return self._busy.setdefault(host, threading.Lock())
def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[int]], str]:
host = urlparse(url).netloc.lower()
with self._host_lock(host):
with self._lock:
wait = self._next_ok.get(host, 0.0) - time.monotonic()
if wait > 0:
time.sleep(wait)
try:
return iv.vector_for_url(url, timeout=timeout)
finally:
with self._lock:
self._next_ok[host] = time.monotonic() + self._pause
def backfill_table(
cur,
conn,
table: str,
*,
apply: bool,
recompute_stale: bool,
limit: Optional[int],
workers: int,
timeout: Optional[float],
pacer: _HostPacer,
) -> Dict[str, Any]:
entry: Dict[str, Any] = {
"table": table, "candidates": 0, "computed": 0, "failed": 0, "skipped": False,
}
if not iv.table_has_columns(cur, table):
logger.info("%s: no img_vector columns yet - run migrate_brand_schema --apply first", table)
entry["skipped"] = True
return entry
rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit)
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)
if not apply:
for r in rows[:5]:
logger.info(" %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:90])
if len(rows) > 5:
logger.info(" ... and %d more", len(rows) - 5)
return entry
results: Dict[str, Tuple[Optional[List[int]], str]] = {}
with ThreadPoolExecutor(max_workers=workers) as pool:
futures = {
pool.submit(pacer.fetch, r["primary_url"], timeout): r["image_id"] for r in rows
}
for fut in as_completed(futures):
image_id = futures[fut]
try:
results[image_id] = fut.result()
except Exception as exc: # noqa: BLE001 - one bad URL must not stop the table
logger.warning(" %s: %s", image_id, exc)
results[image_id] = (None, "")
# Writes happen here, on the caller's connection, in row order.
for r in rows:
vector, src = results.get(r["image_id"], (None, ""))
if vector is None:
entry["failed"] += 1
logger.info(" failed %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:80])
continue
iv.store_vector(cur, table, r["image_id"], vector, src)
entry["computed"] += 1
conn.commit()
logger.info("%s: %d written, %d failed", table, entry["computed"], entry["failed"])
return entry
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
mode = parser.add_mutually_exclusive_group(required=True)
mode.add_argument("--brands", help="comma-separated brands to fill")
mode.add_argument("--all", action="store_true", help="every brand table, including inactive brands")
parser.add_argument("--apply", action="store_true", help="download and write (default: dry run)")
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("--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,
help="seconds between two requests to the same host")
parser.add_argument("--json", action="store_true", help="machine-readable summary")
args = parser.parse_args()
workers = max(1, min(MAX_WORKERS, args.workers))
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")
conn = _connect()
if conn is None:
logger.error("No database connection.")
return 1
report: List[Dict[str, Any]] = []
pacer = _HostPacer(args.pause)
try:
with conn.cursor() as cur:
brands = None if args.all else [b for b in args.brands.split(",") if b.strip()]
for _suffix, table in _brand_tables(cur, brands):
report.append(backfill_table(
cur, conn, table,
apply=args.apply, recompute_stale=args.recompute_stale,
limit=args.limit, workers=workers, timeout=args.timeout, pacer=pacer,
))
finally:
conn.close()
candidates = sum(e["candidates"] for e in report)
computed = sum(e["computed"] for e in report)
failed = sum(e["failed"] for e in report)
skipped = sum(1 for e in report if e["skipped"])
if args.json:
print(json.dumps({"apply": args.apply, "tables": report,
"candidates": candidates, "computed": computed,
"failed": failed, "skipped_tables": skipped}, indent=2))
else:
logger.info("")
if args.apply:
logger.info("Applied: %d candidate(s), %d written, %d failed, %d table(s) skipped",
candidates, computed, failed, skipped)
else:
logger.info("Dry run: %d row(s) would be attempted across %d table(s), %d table(s) skipped",
candidates, len(report) - skipped, skipped)
if candidates:
logger.info("Re-run with --apply to download and write.")
return 0
if __name__ == "__main__":
sys.exit(main())

View File

@@ -108,8 +108,33 @@ _TYPE_ALIASES = {
"JSONB": {"jsonb"},
"DOUBLE PRECISION": {"double precision"},
"vector(384)": {"USER-DEFINED"},
"vector(3072)": {"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)
def _pgvector_version(cur) -> Optional[str]:
try:
cur.execute("SELECT extversion FROM pg_extension WHERE extname = 'vector'")
row = cur.fetchone()
return str(row[0]) if row and row[0] else None
except Exception: # noqa: BLE001 - informational; a failed probe is reported as unknown
return None
def _version_tuple(text: str) -> tuple:
parts = []
for piece in text.split("."):
digits = "".join(ch for ch in piece if ch.isdigit())
if not digits:
break
parts.append(int(digits))
return tuple(parts)
class RecordingCursor:
"""Wraps a real cursor so a dry run can see the statements without
@@ -155,6 +180,13 @@ def main() -> int:
logger.info("database : %s / %s", DB_HOST, DB_NAME)
logger.info("mode : %s", "APPLY (writing)" if apply else "dry run (no writes)")
with conn.cursor() as cur:
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).",
pgv, ".".join(map(str, _MIN_PGVECTOR_FOR_IMG_VECTOR)))
logger.info("")
declared_types = _declared_types()