Image vector embedding
This commit is contained in:
218
scripts/backfill_image_vectors.py
Normal file
218
scripts/backfill_image_vectors.py
Normal 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())
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user