Spaces:
Sleeping
Sleeping
File size: 2,437 Bytes
2db8ee1 671fed1 2db8ee1 de88b8d 671fed1 2db8ee1 671fed1 de88b8d 671fed1 de88b8d 2db8ee1 671fed1 2db8ee1 671fed1 d8c0ee6 2db8ee1 d8c0ee6 2db8ee1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | import faiss
import numpy as np
from PIL import Image
import os
import pickle
class ImageVectorStore:
def __init__(self):
self._model = None
self._preprocess = None
self._device = None
self.index = faiss.IndexFlatL2(512)
self.metadata = []
@property
def device(self):
if self._device is None:
import torch
self._device = "cuda" if torch.cuda.is_available() else "cpu"
return self._device
def _load_clip(self):
if self._model is None:
import clip
self._model, self._preprocess = clip.load("ViT-B/32", device=self.device)
@property
def model(self):
self._load_clip()
return self._model
@property
def preprocess(self):
self._load_clip()
return self._preprocess
def add_images(self, image_paths, metadatas):
import torch
images = [
self.preprocess(Image.open(p)).unsqueeze(0)
for p in image_paths
]
images = torch.cat(images).to(self.device)
with torch.no_grad():
emb = self.model.encode_image(images)
self.index.add(emb.cpu().numpy().astype("float32"))
for path, meta in zip(image_paths, metadatas):
enriched_meta = meta.copy()
enriched_meta["image_path"] = path
self.metadata.append(enriched_meta)
def search(self, query_text, k=5):
import clip
import torch
if self.index.ntotal == 0:
return []
k = min(k, self.index.ntotal)
text_tokens = clip.tokenize([query_text]).to(self.device)
with torch.no_grad():
q_emb = self.model.encode_text(text_tokens)
_, idxs = self.index.search(q_emb.cpu().numpy().astype("float32"), k)
return [self.metadata[i] for i in idxs[0] if 0 <= i < len(self.metadata)]
def save_local(self, folder_path):
os.makedirs(folder_path, exist_ok=True)
faiss.write_index(self.index, os.path.join(folder_path, "index.faiss"))
with open(os.path.join(folder_path, "metadata.pkl"), "wb") as f:
pickle.dump(self.metadata, f)
def load_local(self, folder_path):
self.index = faiss.read_index(os.path.join(folder_path, "index.faiss"))
with open(os.path.join(folder_path, "metadata.pkl"), "rb") as f:
self.metadata = pickle.load(f)
|