from numpy import argsort from PIL import Image as PImage from sklearn.metrics.pairwise import euclidean_distances, cosine_distances from torch import no_grad def embed_word(word, processor, model, device): txt_t = processor(text=[word], padding="max_length", max_length=64, return_tensors="pt").to(device) with no_grad(): txt_embedding = model.get_text_features(**txt_t).pooler_output return txt_embedding.squeeze().cpu().numpy() def embed_image(image, processor, model, device): img_t = processor(images=[image], return_tensors="pt", padding=True).to(device) with no_grad(): img_embedding = model.get_image_features(**img_t).pooler_output return img_embedding.squeeze().cpu().numpy() def idxs_by_dist(img_embeddings, txt_embedding, cos=True): if cos: dists = cosine_distances([txt_embedding], img_embeddings) else: dists = euclidean_distances([txt_embedding], img_embeddings) return argsort(dists[0]) def make_image(imgs, order): iw = sum(i.width for i in imgs) ih = min(i.height for i in imgs) oimg = PImage.new("RGB", (iw, ih)) cw = 0 for idx in order: mw = imgs[idx].width oimg.paste(imgs[idx], (cw, 0, cw+mw, ih)) cw += mw return oimg def idxs_along_axes(img_embeddings, txt_embeddings): word_dists = euclidean_distances(txt_embeddings, img_embeddings) emb_dists = word_dists[0] / word_dists[1] return argsort(emb_dists)