← back to Dw Photo Capture
Wire in fine-tuned DW CLIP: sanity-check picks best epoch (ep4, 52.9% held-out recall@1 vs 23.0% base — 2.3x, no overfit); load dw_clip_ft.pt in search_service + embed_catalog (both sides); catalog re-embedding on FT model (311k, snapshot saved for rollback)
ec02bd7dd11c1476da6b3993165c0d621601ccca · 2026-07-07 06:39:29 -0700 · Steve Abrams
Files touched
A visual-search/.gitignoreM visual-search/embed_catalog.pyM visual-search/search_service.pyA visual-search/train/sanity_check.py
Diff
commit ec02bd7dd11c1476da6b3993165c0d621601ccca
Author: Steve Abrams <steve@designerwallcoverings.com>
Date: Tue Jul 7 06:39:29 2026 -0700
Wire in fine-tuned DW CLIP: sanity-check picks best epoch (ep4, 52.9% held-out recall@1 vs 23.0% base — 2.3x, no overfit); load dw_clip_ft.pt in search_service + embed_catalog (both sides); catalog re-embedding on FT model (311k, snapshot saved for rollback)
---
visual-search/.gitignore | 3 ++
visual-search/embed_catalog.py | 7 +++++
visual-search/search_service.py | 7 +++++
visual-search/train/sanity_check.py | 61 +++++++++++++++++++++++++++++++++++++
4 files changed, 78 insertions(+)
diff --git a/visual-search/.gitignore b/visual-search/.gitignore
new file mode 100644
index 0000000..1d61c93
--- /dev/null
+++ b/visual-search/.gitignore
@@ -0,0 +1,3 @@
+dw_clip_ft.pt
+*.pt
+image_embeddings*.sql.gz
diff --git a/visual-search/embed_catalog.py b/visual-search/embed_catalog.py
index 79fc5aa..418afa8 100644
--- a/visual-search/embed_catalog.py
+++ b/visual-search/embed_catalog.py
@@ -21,6 +21,13 @@ args = ap.parse_args()
print("loading CLIP ViT-B-32 (laion2b)…", flush=True)
model, _, preprocess = open_clip.create_model_and_transforms("ViT-B-32", pretrained="laion2b_s34b_b79k")
+
+import torch as _t, os as _o
+_ft = _o.path.join(_o.path.dirname(_o.path.abspath(__file__)), "dw_clip_ft.pt")
+if _o.path.exists(_ft):
+ model.load_state_dict(_t.load(_ft, map_location="cpu"), strict=False)
+ print("loaded fine-tuned DW CLIP weights:", _ft, flush=True)
+
model.eval()
torch.set_num_threads(max(1, os.cpu_count() - 2))
DIM = model.visual.output_dim
diff --git a/visual-search/search_service.py b/visual-search/search_service.py
index 9279439..dd444c7 100644
--- a/visual-search/search_service.py
+++ b/visual-search/search_service.py
@@ -15,6 +15,13 @@ PORT = int(os.environ.get("VS_PORT", "9914"))
print("loading CLIP…", flush=True)
model, _, preprocess = open_clip.create_model_and_transforms("ViT-B-32", pretrained="laion2b_s34b_b79k")
+
+import torch as _t, os as _o
+_ft = _o.path.join(_o.path.dirname(_o.path.abspath(__file__)), "dw_clip_ft.pt")
+if _o.path.exists(_ft):
+ model.load_state_dict(_t.load(_ft, map_location="cpu"), strict=False)
+ print("loaded fine-tuned DW CLIP weights:", _ft, flush=True)
+
model.eval(); torch.set_num_threads(max(1, os.cpu_count() - 2))
LOCK = threading.Lock()
diff --git a/visual-search/train/sanity_check.py b/visual-search/train/sanity_check.py
new file mode 100644
index 0000000..a736c2b
--- /dev/null
+++ b/visual-search/train/sanity_check.py
@@ -0,0 +1,61 @@
+#!/usr/bin/env python3
+# Pick the best fine-tuned checkpoint by HELD-OUT image→caption retrieval accuracy (recall@1).
+# Held-out = catalog rows NOT in the training cache. Compares base CLIP vs each ckpt/dw_clip_ft_ep*.pt.
+# Higher held-out recall@1 = better generalization (guards against the low-loss overfit trap).
+import os, io, glob, urllib.request, urllib.parse, subprocess
+os.environ.setdefault('PYTORCH_ENABLE_MPS_FALLBACK', '1')
+import torch, torch.nn.functional as F, open_clip
+from PIL import Image
+import finetune as T # reuse caption()
+
+HERE = os.path.dirname(os.path.abspath(__file__))
+DB = os.environ.get('DW_UNIFIED_DB', 'postgresql://dw_admin@127.0.0.1:5432/dw_unified')
+N = 300
+trained = set(os.path.splitext(f)[0] for f in os.listdir(os.path.join(HERE, 'cache')))
+
+def small(u):
+ if 'shopify' in u:
+ pr=urllib.parse.urlparse(u); q=urllib.parse.parse_qs(pr.query); q['width']=['512']
+ return urllib.parse.urlunparse(pr._replace(query=urllib.parse.urlencode(q,doseq=True)))
+ return u
+
+# pull held-out rows (id, url, pattern, color, vendor, type) not already trained
+sql = ("select id,image_url,regexp_replace(pattern_name,E'[\\t\\n\\r]',' ','g'),"
+ "coalesce(nullif(regexp_replace(coalesce(color_name,''),E'.*\"Name\":\\s*\"([^\"]+)\".*',E'\\1'),color_name),color_primary,''),"
+ "coalesce(original_vendor_name,vendor_code,''),coalesce(product_type,'Wallcovering') "
+ "from vendor_catalog where image_url like 'http%' and pattern_name is not null and pattern_name<>'' "
+ "order by random() limit 1200")
+rows = [l.split('\t') for l in subprocess.check_output(['psql',DB,'-F','\t','-tA','-c',sql]).decode().splitlines()]
+heldout = [r for r in rows if len(r)>=6 and r[0] not in trained][:N]
+print(f"held-out candidates: {len(heldout)}", flush=True)
+
+imgs, caps = [], []
+for r in heldout:
+ try:
+ data = urllib.request.urlopen(urllib.request.Request(small(r[1]),headers={'User-Agent':'t'}),timeout=20).read()
+ imgs.append(Image.open(io.BytesIO(data)).convert('RGB')); caps.append(T.caption(r[2],r[3],r[4],r[5]))
+ except Exception: pass
+print(f"downloaded {len(imgs)} held-out images", flush=True)
+
+dev='mps' if torch.backends.mps.is_available() else 'cpu'
+model,_,pp = open_clip.create_model_and_transforms('ViT-B-32', pretrained='laion2b_s34b_b79k')
+tok = open_clip.get_tokenizer('ViT-B-32'); model=model.to(dev).float().eval()
+IM = torch.stack([pp(i) for i in imgs]).to(dev)
+TX = tok(caps).to(dev)
+
+def recall1(sd=None):
+ if sd: model.load_state_dict(sd, strict=False)
+ else: model.load_state_dict(open_clip.create_model('ViT-B-32', pretrained='laion2b_s34b_b79k').state_dict(), strict=False)
+ with torch.no_grad():
+ imf=F.normalize(model.encode_image(IM),dim=-1); txf=F.normalize(model.encode_text(TX),dim=-1)
+ sim=imf@txf.t(); top=sim.argmax(dim=1); correct=(top==torch.arange(len(imgs),device=dev)).float().mean().item()
+ return correct*100
+
+print(f"\n{'model':<26} held-out recall@1")
+print(f"{'base laion2b':<26} {recall1(None):.1f}%")
+best=('base',recall1(None))
+for ck in sorted(glob.glob(os.path.join(HERE,'ckpt','dw_clip_ft_ep*.pt'))):
+ acc=recall1(torch.load(ck,map_location='cpu')); name=os.path.basename(ck)
+ print(f"{name:<26} {acc:.1f}%")
+ if acc>best[1]: best=(name,acc)
+print(f"\n★ BEST: {best[0]} @ {best[1]:.1f}% held-out recall@1")
← 5673cad Contrarian gate fix: _updateNeedsConfirm invariant — ANY amb
·
back to Dw Photo Capture
·
Device-test checklist for the iPhone-Safari real-device pass 7d4ccb6 →