Usare l’IA per trovare il tuo stylist delle celebrità (Parte II)
Nel mio precedente post del blog, "Using AI to Find Your Celebrity Stylist," ho spiegato come sfruttare tecnologie di intelligenza artificiale (AI), come Milvus, un database vettoriale open-source nativo per l'AI, e i modelli Hugging Face, per trovare scelte di stile delle celebrità che corrispondono alle tue. In questo post di follow-up, faremo un passo avanti e dimostreremo come ottenere risultati più dettagliati e accurati apportando alcune modifiche al codice del progetto precedente. Inoltre, fornirò suggerimenti su come puoi estendere questo progetto autonomamente.
Se desideri provare direttamente questo progetto, scarica le photos e il notebook completato. Se sei interessato al progetto discusso nel mio blog precedente, puoi consultare le foto che utilizza e il relativo tutorial.
Riepilogo del tutorial per il mio precedente progetto Fashion AI
Prima di approfondire questo progetto, lascia che riepiloghi brevemente il tutorial di cui abbiamo parlato nel mio post precedente. Così non devi lasciare questa pagina per conoscere il contesto.
Importazione di tutte le librerie necessarie per la manipolazione delle immagini
Iniziamo il codice importando tutte le librerie necessarie per la manipolazione delle immagini, inclusi torch per l'estrazione delle caratteristiche, l'oggetto segformer da transformers, matplotlib e alcune importazioni di torchvision come Resize, masks_to_boxes e crop.
import torch
from torch import nn, tensor
from transformers import AutoFeatureExtractor, SegformerForSemanticSegmentation
import matplotlib.pyplot as plt
from torchvision.transforms import Resize
import torchvision.transforms as T
from torchvision.ops import masks_to_boxes
from torchvision.transforms.functional import crop
Pre-elaborazione delle immagini delle celebrità
Dopo aver importato tutti i pacchetti necessari per la manipolazione delle immagini, puoi iniziare a elaborare le tue immagini. Le seguenti tre funzioni (get_segmentation, get_masks e crop_images) vengono utilizzate per segmentare i capi di abbigliamento e ritagliarli per un'ulteriore ingestione.
def get_segmentation(extractor, model, image):
inputs = extractor(images=image, return_tensors="pt")
outputs = model(**inputs)
logits = outputs.logits.cpu()
upsampled_logits = nn.functional.interpolate(
logits,
size=image.size[::-1],
mode="bilinear",
align_corners=False,
)
pred_seg = upsampled_logits.argmax(dim=1)[0]
return pred_seg
# returns two lists masks (tensor) and obj_ids (int)
# "mattmdjaga/segformer_b2_clothes" from hugging face
def get_masks(segmentation):
obj_ids = torch.unique(segmentation)
obj_ids = obj_ids[1:]
masks = segmentation == obj_ids[:, None, None]
return masks, obj_ids
def crop_images(masks, obj_ids, img):
boxes = masks_to_boxes(masks)
crop_boxes = []
for box in boxes:
crop_box = tensor([box[0], box[1], box[2]-box[0], box[3]-box[1]])
crop_boxes.append(crop_box)
preprocess = T.Compose([
T.Resize(size=(256, 256)),
T.ToTensor()
])
cropped_images = {}
for i in range(len(crop_boxes)):
crop_box = crop_boxes[i]
cropped = crop(img, crop_box[1].item(), crop_box[0].item(), crop_box[3].item(), crop_box[2].item())
cropped_images[obj_ids[i].item()] = preprocess(cropped)
return cropped_images
Archiviare i dati delle immagini in un database vettoriale
Usiamo Milvus, un database vettoriale open-source nativo per l'AI, per archiviare i dati delle immagini. Per iniziare, decomprimi il file zip photos per questo progetto e includi la cartella nella stessa directory root del notebook. Una volta completato questo passaggio, puoi eseguire il codice qui sotto per elaborare le immagini e archiviare i dati in Milvus.
import os
image_paths = []
for celeb in os.listdir("./photos"):
for image in os.listdir(f"./photos/{celeb}/"):
image_paths.append(f"./photos/{celeb}/{image}")
from milvus import default_server
from pymilvus import utility, connections
default_server.start()
connections.connect(host="127.0.0.1", port=default_server.listen_port)
DIMENSION = 2048
BATCH_SIZE = 128
COLLECTION_NAME = "fashion"
TOP_K = 3
from pymilvus import FieldSchema, CollectionSchema, Collection, DataType
fields = [
FieldSchema(name="id", dtype=DataType.INT64, is_primary=True, auto_id=True),
FieldSchema(name='filepath', dtype=DataType.VARCHAR, max_length=200),
FieldSchema(name="name", dtype=DataType.VARCHAR, max_length=200),
FieldSchema(name="seg_id", dtype=DataType.INT64),
FieldSchema(name='embedding', dtype=DataType.FLOAT_VECTOR, dim=DIMENSION)
]
schema = CollectionSchema(fields=fields)
collection = Collection(name=COLLECTION_NAME, schema=schema)
index_params = {
"index_type": "IVF_FLAT",
"metric_type": "L2",
"params": {"nlist": 128},
}
collection.create_index(field_name="embedding", index_params=index_params)
collection.load()
Successivamente, puoi eseguire il codice qui sotto per generare embedding utilizzando il modello Nvidia ResNet 50 di Hugging Face.
# run this before importing th resnet50 model if you run into an SSL certificate URLError
import ssl
ssl._create_default_https_context = ssl._create_unverified_context
# Load the embedding model with the last layer removed
embeddings_model = torch.hub.load('NVIDIA/DeepLearningExamples:torchhub', 'nvidia_resnet50', pretrained=True)
embeddings_model = torch.nn.Sequential(*(list(embeddings_model.children())[:-1]))
embeddings_model.eval()
La funzione qui sotto definisce come incorporare e inserire i dati. Successivamente, il codice scorre tutte le immagini, le incorpora e le inserisce in Milvus.
Nota: molti dei componenti qui sotto cambieranno o verranno rimossi quando si utilizzerà la nuova funzionalità di schema dinamico di Milvus.
def embed_insert(data, collection, model):
with torch.no_grad():
output = model(torch.stack(data[0])).squeeze()
collection.insert([data[1], data[2], data[3], output.tolist()])
from PIL import Image
data_batch = [[], [], [], []]
for path in image_paths:
image = Image.open(path)
path_split = path.split("/")
name = " ".join(path_split[2].split("_"))
segmentation = get_segmentation(extractor, model, image)
masks, ids = get_masks(segmentation)
cropped_images = crop_images(masks, ids, image)
for key, image in cropped_images.items():
data_batch[0].append(image)
data_batch[1].append(path)
data_batch[2].append(name)
data_batch[3].append(key)
if len(data_batch[0]) % BATCH_SIZE == 0:
embed_insert(data_batch, collection, embeddings_model)
data_batch = [[], [], [], []]
if len(data_batch[0]) != 0:
embed_insert(data_batch, collection, embeddings_model)
collection.flush()
Interroga il database vettoriale
Il codice seguente dimostra come interrogare Milvus con immagini di input e recuperare i primi tre risultati per ogni capo di abbigliamento.
def embed_search_images(data, model):
with torch.no_grad():
output = model(torch.stack(data))
if len(output) > 1:
return output.squeeze().tolist()
else:
return torch.flatten(output, start_dim=1).tolist()
# data_batch[0] is a list of tensors
# data_batch[1] is a list of filepaths to the images (string)
# data_batch[2] is a list of the names of the people in the images (string)
# data_batch[3] is a list of segmentation keys (int)
data_batch = [[], [], [], []]
search_paths = ["./photos/Taylor_Swift/Taylor_Swift_3.jpg", "./photos/Taylor_Swift/Taylor_Swift_8.jpg"]
for path in search_paths:
image = Image.open(path)
path_split = path.split("/")
name = " ".join(path_split[2].split("_"))
segmentation = get_segmentation(extractor, model, image)
masks, ids = get_masks(segmentation)
cropped_images = crop_images(masks, ids, image)
for key, image in cropped_images.items():
data_batch[0].append(image)
data_batch[1].append(path)
data_batch[2].append(name)
data_batch[3].append(key)
embeds = embed_search_images(data_batch[0], embeddings_model)
import time
start = time.time()
res = collection.search(embeds,
anns_field='embedding',
param={"metric_type": "L2",
"params": {"nprobe": 10}},
limit=TOP_K,
output_fields=['filepath'])
finish = time.time()
print(finish - start)
for index, result in enumerate(res):
print(index)
print(result)
Più pattern matching: selezionare oggetti da ogni immagine
Seguendo il tutorial riepilogato sopra, puoi scoprire le prime tre corrispondenze di stile delle celebrità per ogni capo di abbigliamento che cerchi. Puoi anche creare un’immagine come quella qui sotto senza i riquadri di delimitazione degli elementi abbinati. In questa sezione, spiegherò come trovare stili di moda con pattern più vicini ai tuoi, con alcune modifiche al codice usato nel tutorial precedente.
image
Importare tutte le librerie necessarie per la manipolazione delle immagini
Per cominciare, importa nel tuo codice tutte le librerie necessarie per la manipolazione delle immagini. Se lo hai già fatto, puoi saltare questo passaggio.
import torch
from torch import nn, tensor
from transformers import AutoFeatureExtractor, SegformerForSemanticSegmentation
import matplotlib.pyplot as plt
from torchvision.transforms import Resize
import torchvision.transforms as T
from torchvision.ops import masks_to_boxes
from torchvision.transforms.functional import crop
Pre-elaborazione delle immagini
Dopo aver importato tutti i pacchetti necessari per la manipolazione delle immagini, procedi con il processo di segmentazione delle immagini, che coinvolge tre funzioni: get_segmentation, get_masks e crop_images.
Non dobbiamo apportare alcuna modifica al codice della funzione get_segmentation.
Per la funzione get_masks, dobbiamo solo prendere le segmentazioni che corrispondono agli ID di segmentazione nell’elenco wanted. È una nuova aggiunta che include gli ID di segmentazione per i capi di abbigliamento, come indicato nella model card su Hugging Face.
Apporteremo la maggior parte delle modifiche al codice alla funzione crop_image. Nel mio tutorial precedente, questa funzione restituiva un elenco di immagini ritagliate. Dopo aver modificato parte del codice, ora restituisce tre oggetti: embedding delle immagini ritagliate, un elenco delle coordinate dei riquadri sull’immagine originale e un elenco degli ID di segmentazione. Questa nuova configurazione sposta l’embedding dall’inserimento batch alla fase di trasformazione.
wanted = [1, 3, 4, 5, 6, 7, 8, 9, 10, 16, 17]
def get_segmentation(image):
inputs = extractor(images=image, return_tensors="pt")
outputs = segmentation_model(**inputs)
logits = outputs.logits.cpu()
upsampled_logits = nn.functional.interpolate(
logits,
size=image.size[::-1],
mode="bilinear",
align_corners=False,
)
pred_seg = upsampled_logits.argmax(dim=1)[0]
return pred_seg
restituisce due liste masks (tensor) e obj_ids (int)
"mattmdjaga/segformer_b2_clothes" da hugging face
def get_masks(segmentation): obj_ids = torch.unique(segmentation) obj_ids = obj_ids[1:] wanted_ids = [x.item() for x in obj_ids if x in wanted] wanted_ids = torch.Tensor(wanted_ids) masks = segmentation == wanted_ids[:, None, None] return masks, obj_ids
def crop_images(masks, obj_ids, img): boxes = masks_to_boxes(masks) crop_boxes = [] for box in boxes: crop_box = tensor([box[0], box[1], box[2]-box[0], box[3]-box[1]]) crop_boxes.append(crop_box) preprocess = T.Compose([ T.Resize(size=(256, 256)), T.ToTensor() ]) cropped_images = [] seg_ids = [] for i in range(len(crop_boxes)): crop_box = crop_boxes[i] cropped = crop(img, crop_box[1].item(), crop_box[0].item(), crop_box[3].item(), crop_box[2].item()) cropped_images.append(preprocess(cropped)) seg_ids.append(obj_ids[i].item()) with torch.no_grad(): embeddings = embeddings_model(torch.stack(cropped_images)).squeeze().tolist() return embeddings, boxes.tolist(), seg_ids
Ora che abbiamo le immagini, è il momento di caricarle. Questo passaggio comporta l'inserimento in batch, che abbiamo trattato nel mio tutorial precedente. In questo tutorial, inseriremo tutti i nostri dati in una sola volta come lista di dizionari invece che come lista di liste. Trovo che questo metodo di inserimento sia molto più pulito e ci consenta di aggiungere un nuovo campo allo schema al momento dell'inserimento. In questo caso, aggiungeremo una lista di angoli di ritaglio.
for path in image_paths: image = Image.open(path) path_split = path.split("/") name = " ".join(path_split[2].split("_")) segmentation = get_segmentation(image) masks, ids = get_masks(segmentation) embeddings, crop_corners, seg_ids = crop_images(masks, ids, image) inserts = [{"embedding": embeddings[x], "seg_id": seg_ids[x], "name": name, "filepath": path, "crop_corner": crop_corners[x]} for x in range(len(embeddings))] collection.insert(inserts) collection.flush()
### Interrogare il database vettoriale
Ora è il momento di effettuare query in Milvus, il nostro [database vettoriale](https://zilliz.com/learn/what-is-vector-database). Rispetto ai passaggi che abbiamo utilizzato nel tutorial precedente, ci sono alcune differenze qui:
- Primo, limitiamo a cinque il numero di "corrispondenze" che ci interessano in un'immagine.
- Secondo, mostriamo le tre immagini con la corrispondenza più vicina.
- Terzo, aggiungiamo una funzione per ottenere una mappa di colori per disegnare riquadri di delimitazione in colori diversi.
Ora impostiamo la figura e gli assi di `matplotlib`. Poi, cicliamo attraverso tutte le nostre immagini e applichiamo le tre funzioni di elaborazione menzionate sopra per ottenere le segmentazioni e i riquadri di delimitazione.
Dopo aver pre-elaborato le immagini, possiamo cercarle in Milvus. Otteniamo le tre risposte migliori per ciascuna immagine in base al numero di articoli "corrispondenti" che contengono. Infine, stampiamo i risultati insieme ai riquadri di delimitazione che hanno restituito corrispondenze.
from pprint import pprint from PIL import ImageDraw from collections import Counter import matplotlib.patches as patches
LIMIT = 5 # Quante corrispondenze più vicine per capo di abbigliamento analizzare CLOSEST = 3 # Quante immagini più vicine visualizzare. CLOSEST <= Limit
search_paths = ["./photos/Taylor_Swift/Taylor_Swift_2.jpg", "./photos/Jenna_Ortega/Jenna_Ortega_6.jpg"] # Immagini da cercare
def get_cmap(n, name='hsv'): '''Restituisce una funzione che mappa ciascun indice in 0, 1, ..., n-1 su un colore RGB distinto; l'argomento parola chiave name deve essere il nome di una colormap mpl standard. Tratto da https://stackoverflow.com/questions/14720331/how-to-generate-random-colors-in-matplotlib''' return plt.cm.get_cmap(name, n)
Crea i subplot dei risultati
f, axarr = plt.subplots(max(len(search_paths), 2), CLOSEST + 1)
for search_i, path in enumerate(search_paths): # Genera ritagli ed embedding per tutti gli elementi trovati image = Image.open(path) segmentation = get_segmentation(image) masks, ids = get_masks(segmentation) embeddings, crop_corners, _ = crop_images(masks, ids, image)
Genera mappa colori
cmap = get_cmap(len(crop_corners))
# Visualizza il primo riquadro con l'immagine che viene cercata
axarr[search_i][0].imshow(image)
axarr[search_i][0].set_title('Immagine di ricerca')
axarr[search_i][0].axis('off')
for i, (x0, y0, x1, y1) in enumerate(crop_corners):
rect = patches.Rectangle((x0, y0), x1-x0, y1-y0, linewidth=1, edgecolor=cmap(i), facecolor='none')
axarr[search_i][0].add_patch(rect)
# Cerca nel database tutti i ritagli
start = time.time()
res = collection.search(embeddings,
anns_field='embedding',
param={"metric_type": "L2",
"params": {"nprobe": 10}, "offset": 0},
limit=LIMIT,
output_fields=['filepath', 'crop_corner'])
finish = time.time()
print("Tempo totale di ricerca: ", finish - start)
# Riassumi i risultati unici principali e ponderali in base alla posizione nei risultati
filepaths = []
for hits in res:
seen = set()
for i, hit in enumerate(hits):
if hit.entity.get("filepath") not in seen:
seen.add(hit.entity.get("filepath"))
filepaths.extend([hit.entity.get("filepath") for _ in range(len(hits) - i)])
# Trova l'immagine del risultato più comunemente classificata
counts = Counter(filepaths)
most_common = [path for path, _ in counts.most_common(CLOSEST)]
# Per ogni immagine, estrai l'elemento corrispondente trovato che correla con le immagini di ricerca
matches = {}
for i, hits in enumerate(res):
matches[i] = {}
tracker = set(most_common)
for hit in hits:
if hit.entity.get("filepath") in tracker:
matches[i][hit.entity.get("filepath")] = hit.entity.get("crop_corner")
tracker.remove( hit.entity.get("filepath"))
# Visualizza le immagini più comuni nei risultati
for res_i, res_path in enumerate(most_common):
# Visualizza ciascuna delle immagini accanto all'immagine di ricerca
image = Image.open(res_path)
axarr[search_i][res_i+1].imshow(image)
axarr[search_i][res_i+1].set_title(" ".join(res_path.split("/")[2].split("_")))
axarr[search_i][res_i+1].axis('off')
Aggiungi riquadri di delimitazione per tutti gli elementi corrispondenti
for key, value in matches.items():
if res_path in value:
x0, y0, x1, y1 = value[res_path]
rect = patches.Rectangle((x0, y0), x1-x0, y1-y0, linewidth=1, edgecolor=cmap(key), facecolor='none')
axarr[search_i][res_i+1].add_patch(rect)
Una volta completati i passaggi sopra, dovresti ottenere un risultato simile a quello qui sotto o all'inizio di questa sessione.

## Cosa c’è dopo? possibili estensioni del progetto
Per ora sto mettendo questo progetto in pausa per lavorare ad altri progetti, ma sei libero di estenderlo se vuoi! Ecco tre possibili estensioni.
Innanzitutto, puoi sviluppare un po' di più il gioco del confronto. Ad esempio, puoi raggruppare insieme gli elementi divisi, come segnare entrambe le scarpe come un singolo elemento. Puoi anche aggiungere più foto di celebrità o amici per ulteriori confronti.
In secondo luogo, puoi trasformare questo progetto in un identificatore di moda o in un sistema di raccomandazione. Invece di usare immagini di celebrità, puoi usare foto di vestiti che possono essere acquistati online. Quando un utente carica una foto, puoi confrontarla con le immagini nel tuo database vettoriale e suggerire all'utente i capi di abbigliamento più simili.
In terzo luogo, puoi creare un generatore di stile, che potrebbe essere più impegnativo. Esistono vari modi per farlo, ma un’idea è prendere più foto di un utente e generare suggerimenti basati su di esse. Questo approccio prevede l’utilizzo di un modello generativo di immagini per fornire suggerimenti di stile e confrontarli con le foto più simili degli utenti come riferimento. Possiamo quindi suggerire qualcosa che abbia senso sulla base di questo confronto.
Queste tre estensioni sono solo alcuni esempi di modi per migliorare il mio semplice progetto utilizzando un modello di immagini e un database vettoriale come [Milvus](https://zilliz.com/what-is-milvus). L’uso di un database vettoriale abilita una varietà di attività di ricerca per similarità, il che è particolarmente utile per confrontare immagini.
## Riepilogo
In questo tutorial, abbiamo esteso il nostro primo progetto in stile celebrity utilizzando il nuovo schema dinamico di Milvus, filtrando determinati ID di segmentazione e tenendo traccia dei bounding box delle nostre corrispondenze. Abbiamo anche ordinato i risultati della ricerca per restituire i primi tre risultati in base al numero di corrispondenze.
Il nuovo schema dinamico di Milvus ci consente di aggiungere campi extra quando carichiamo dati utilizzando un formato dizionario, cambiando il modo in cui inizialmente caricavamo in batch una lista di liste. Ha inoltre facilitato l’aggiunta delle coordinate di ritaglio senza modificare lo schema.
Come nuovo passaggio di pre-elaborazione, abbiamo filtrato determinati ID che non sono correlati all’abbigliamento in base alla model card in Hugging Face. Filtriamo questi ID nella funzione `get_masks`. Curiosità: l’oggetto `obj_ids` in quella funzione è in realtà un tensore.
Abbiamo anche tenuto traccia dei bounding box. Abbiamo spostato il passaggio di embedding nella funzione di ritaglio dell’immagine e restituito gli embedding con i bounding box e gli ID di segmentazione. Poi, abbiamo salvato questi embedding in Milvus utilizzando uno schema dinamico.
Al momento della query, abbiamo aggregato tutte le immagini restituite in base al numero di bounding box che contenevano, permettendoci di trovare l’immagine di celebrity più simile tramite diversi capi di abbigliamento.
Ora tocca a te. Puoi prendere i miei suggerimenti e trasformarli in qualcos’altro, come un sistema di raccomandazione di moda, un sistema migliore di confronto dello stile per te e i tuoi amici, o un’app di IA generativa per la moda.
Continua a leggere

We spent 8 years making vector databases faster. Then we stopped.
Rarely queried embeddings still need to stay searchable. See how Vector Lakebase enables on-demand vector search without always-on compute costs.

Introducing Customer-Managed Encryption Keys (CMEK) on Zilliz Cloud
We're announcing the general availability of Customer-Managed Encryption Keys (CMEK) on Zilliz Cloud.

Vector Databases vs. Document Databases
Use a vector database for similarity search and AI-powered applications; use a document database for flexible schema and JSON-like data storage.



