Utiliser l’IA pour trouver votre styliste de célébrités (Partie II)
Dans mon précédent article de blog, "Utiliser l’IA pour trouver votre styliste de célébrité", j’ai expliqué comment tirer parti des technologies d’intelligence artificielle (IA), telles que Milvus, une base de données vectorielle open source native pour l’IA, et les modèles Hugging Face, pour trouver des choix de style de célébrités qui correspondent aux vôtres. Dans cet article de suivi, nous irons un peu plus loin et montrerons comment obtenir des résultats plus détaillés et précis en apportant quelques modifications au code du projet précédent. De plus, je fournirai des suggestions sur la façon dont vous pouvez étendre ce projet par vous-même.
Si vous souhaitez essayer ce projet directement, téléchargez les photos et le notebook terminé. Si vous vous intéressez au projet abordé dans mon précédent blog, vous pouvez consulter les photos qu’il utilise et son tutoriel.
Récapitulatif du tutoriel de mon précédent projet d’IA pour la mode
Avant d’examiner ce projet en profondeur, laissez-moi récapituler brièvement le tutoriel dont nous avons parlé dans mon précédent article. Ainsi, vous n’avez pas besoin de quitter cette page pour comprendre le contexte.
Importation de toutes les bibliothèques nécessaires à la manipulation d’images
Nous commençons le code en important toutes les bibliothèques nécessaires à la manipulation d’images, notamment torch pour l’extraction de caractéristiques, l’objet segformer de transformers, matplotlib, ainsi que quelques imports de torchvision tels que Resize, masks_to_boxes et 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
Prétraitement des images de célébrités
Après avoir importé tous les packages nécessaires à la manipulation d’images, vous pouvez commencer à traiter vos images. Les trois fonctions suivantes (get_segmentation, get_masks et crop_images) sont utilisées pour segmenter les articles vestimentaires et les recadrer en vue d’une ingestion ultérieure.
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
Stocker les données d’image dans une base de données vectorielle
Nous utilisons Milvus, une base de données vectorielle open source native pour l’IA, pour stocker les données d’images. Pour commencer, décompressez le fichier zip photos pour ce projet et incluez le dossier dans le même répertoire racine que le notebook. Une fois cette étape terminée, vous pouvez exécuter le code ci-dessous pour traiter les images et stocker les données dans 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()
Ensuite, vous pouvez exécuter le code ci-dessous pour générer des embeddings à l’aide du modèle Nvidia ResNet 50 depuis 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 fonction ci-dessous définit comment intégrer et insérer les données. Ensuite, le code parcourt toutes les images, puis les intègre et les insère dans Milvus.
Remarque : De nombreux composants ci-dessous changeront ou seront supprimés lors de l’utilisation de la nouvelle fonctionnalité de schéma dynamique de 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()
Interroger la base de données vectorielle
Le code suivant montre comment interroger Milvus avec des images d’entrée et récupérer les trois meilleurs résultats pour chaque article vestimentaire.
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)
Plus de correspondance de motifs : sélectionner des objets dans chaque image
En suivant le tutoriel récapitulé ci-dessus, vous pouvez découvrir les trois meilleures correspondances de style de célébrités pour chaque vêtement que vous recherchez. Vous pouvez également créer une image comme celle ci-dessous sans les boîtes englobantes des éléments correspondants. Dans cette section, j’expliquerai comment trouver des styles de mode avec des motifs plus proches des vôtres, avec quelques modifications de code par rapport à celui utilisé dans le tutoriel précédent.
image
Importer toutes les bibliothèques nécessaires à la manipulation d’images
Pour commencer, importez dans votre code toutes les bibliothèques nécessaires à la manipulation d’images. Si vous l’avez déjà fait, vous pouvez ignorer cette étape.
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
Prétraiter vos images
Lorsque vous avez importé tous les packages nécessaires à la manipulation d’images, poursuivez avec le processus de segmentation d’image, qui implique trois fonctions : get_segmentation, get_masks et crop_images.
Nous n’avons pas besoin d’apporter de modifications au code de la fonction get_segmentation.
Pour la fonction get_masks, il nous suffit de récupérer les segmentations qui correspondent aux ID de segmentation dans la liste wanted. Il s’agit d’un nouvel ajout qui inclut les ID de segmentation pour les vêtements, comme indiqué dans la fiche du modèle sur Hugging Face.
Nous apporterons la plupart des modifications de code à la fonction crop_image. Dans mon tutoriel précédent, cette fonction renvoyait une liste d’images recadrées. Après avoir modifié une partie du code, elle renvoie désormais trois objets : les embeddings des images recadrées, une liste des coordonnées des boîtes sur l’image originale et une liste des ID de segmentation. Cette nouvelle configuration déplace l’embedding de l’insertion par lot vers l’étape de transformation.
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
# 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:]
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
Maintenant que nous avons les images, il est temps de les charger. Cette étape implique une insertion par lots, que nous avons abordée dans mon tutoriel précédent. Dans ce tutoriel, nous insérerons toutes nos données en une seule fois sous forme de liste de dictionnaires plutôt que de liste de listes. Je trouve cette méthode d’insertion beaucoup plus propre, et elle nous permet d’ajouter un nouveau champ au schéma au moment de l’insertion. Dans ce cas, nous ajouterons une liste de coins de recadrage.
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()
Interroger la base de données vectorielle
Maintenant, il est temps d’effectuer des requêtes dans Milvus, notre base de données vectorielle. Par rapport aux étapes que nous avons utilisées dans le tutoriel précédent, il y a ici quelques différences :
- Premièrement, nous limitons à cinq le nombre de « correspondances » qui nous intéressent dans une image.
- Deuxièmement, nous affichons les trois images correspondantes les plus proches.
- Troisièmement, nous ajoutons une fonction pour obtenir une carte de couleurs permettant de dessiner des boîtes englobantes de différentes couleurs.
Maintenant, nous configurons la figure et les axes matplotlib. Ensuite, nous parcourons toutes nos images et appliquons les trois fonctions de traitement mentionnées ci-dessus afin d’obtenir les segmentations et les boîtes englobantes.
Après avoir prétraité les images, nous pouvons les rechercher dans Milvus. Nous obtenons les trois meilleures réponses pour chaque image en fonction du nombre d’articles « correspondants » qu’elles contiennent. Enfin, nous imprimons les résultats ainsi que les boîtes englobantes qui ont renvoyé des correspondances.
from pprint import pprint
from PIL import ImageDraw
from collections import Counter
import matplotlib.patches as patches
LIMIT = 5 # How many closes matches per article of clothing to analyze
CLOSEST = 3 # How many closest images to display. CLOSEST <= Limit
search_paths = ["./photos/Taylor_Swift/Taylor_Swift_2.jpg", "./photos/Jenna_Ortega/Jenna_Ortega_6.jpg"] # Images to search for
def get_cmap(n, name='hsv'):
'''Returns a function that maps each index in 0, 1, ..., n-1 to a distinct
RGB color; the keyword argument name must be a standard mpl colormap name.
Sourced from https://stackoverflow.com/questions/14720331/how-to-generate-random-colors-in-matplotlib'''
return plt.cm.get_cmap(name, n)
# Create the result subplots
f, axarr = plt.subplots(max(len(search_paths), 2), CLOSEST + 1)
for search_i, path in enumerate(search_paths):
# Generate crops and embeddings for all items found
image = Image.open(path)
segmentation = get_segmentation(image)
masks, ids = get_masks(segmentation)
embeddings, crop_corners, _ = crop_images(masks, ids, image)
# Generate color map
cmap = get_cmap(len(crop_corners))
# Display the first box with image being searched for
axarr[search_i][0].imshow(image)
axarr[search_i][0].set_title('Search Image')
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)
# Search the database for all the crops
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("Total Search Time: ", finish - start)
# Summarize the top unique results and weight them based on position in results
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)])
# Find the most commonly ranked result image
counts = Counter(filepaths)
most_common = [path for path, _ in counts.most_common(CLOSEST)]
# For each image, extract the corresponding item found that correlates to search images
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"))
# Display the most common images in results
for res_i, res_path in enumerate(most_common):
# Display each of the images next to search image
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')
# Add boudning boxes for all matched items
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)
Une fois que vous avez terminé les étapes ci-dessus, vous devriez obtenir un résultat similaire à celui ci-dessous ou à celui du début de cette session.
image
Quelle est la suite ? extensions possibles du projet
Je mets ce projet en pause pour travailler sur d’autres projets pour le moment, mais vous êtes libre de l’étendre si vous le souhaitez ! Voici trois extensions possibles.
Tout d’abord, vous pouvez étoffer un peu plus le jeu de comparaison. Par exemple, vous pouvez regrouper les articles séparés, comme marquer les deux chaussures comme un seul article. Vous pouvez également ajouter davantage de photos de célébrités ou d’amis pour plus de comparaisons.
Deuxièmement, vous pouvez transformer ce projet en identificateur de mode ou en système de recommandation. Au lieu d’utiliser des images de célébrités, vous pouvez utiliser des photos de vêtements qui peuvent être achetés en ligne. Lorsqu’un utilisateur téléverse une photo, vous pouvez la comparer aux images de votre base de données vectorielle et suggérer à l’utilisateur les articles vestimentaires les plus proches.
Troisièmement, vous pouvez créer un générateur de style, ce qui peut être plus difficile. Il existe différentes façons de le faire, mais une idée consiste à prendre plusieurs photos d’un utilisateur et à générer des suggestions à partir de celles-ci. Cette approche implique d’utiliser un modèle d’image générative pour fournir des suggestions de style et de les comparer aux photos les plus proches des utilisateurs à titre de référence. Nous pouvons ensuite suggérer quelque chose de pertinent sur la base de cette comparaison.
Ces trois extensions ne sont que quelques exemples de façons d’améliorer mon projet simple en utilisant un modèle d’image et une base de données vectorielle comme Milvus. L’utilisation d’une base de données vectorielle permet une variété de tâches de recherche de similarité, ce qui est particulièrement précieux pour comparer des photos.
Résumé
Dans ce tutoriel, nous avons étendu notre premier projet de style de célébrité en utilisant le nouveau schéma dynamique de Milvus, en filtrant certains ID de segmentation et en suivant les boîtes englobantes de nos correspondances. Nous avons également trié nos résultats de recherche afin de renvoyer les trois meilleurs résultats en fonction du nombre de correspondances.
Le nouveau schéma dynamique de Milvus nous permet d’ajouter des champs supplémentaires lorsque nous téléversons des données à l’aide d’un format dictionnaire, modifiant ainsi la manière dont nous téléversions initialement par lots une liste de listes. Il a également facilité l’ajout de coordonnées de recadrage sans modifier le schéma.
Comme nouvelle étape de prétraitement, nous avons filtré certains ID qui ne sont pas liés aux vêtements, sur la base de la carte du modèle dans Hugging Face. Nous filtrons ces ID dans la fonction get_masks. Fait amusant, l’objet obj_ids dans cette fonction est en réalité un tenseur.
Nous avons également suivi les boîtes englobantes. Nous avons déplacé l’étape d’embedding vers la fonction de recadrage d’image et renvoyé les embeddings avec les boîtes englobantes et les ID de segmentation. Ensuite, nous avons enregistré ces embeddings dans Milvus à l’aide d’un schéma dynamique.
Au moment de la requête, nous avons agrégé toutes les images renvoyées selon le nombre de boîtes englobantes qu’elles contenaient, ce qui nous permet de trouver l’image de célébrité la plus proche via différents articles vestimentaires. Maintenant, c’est à vous. Vous pouvez prendre mes suggestions et en faire autre chose, comme un système de recommandation de mode, un meilleur système de comparaison de style pour vous et vos amis, ou une application d’IA générative de mode.
Continuer à lire

Zilliz Cloud Update: Smarter Autoscaling for Cost Savings, Stronger Compliance with Audit Logs, and More
What's new in Zilliz Cloud? Smarter autoscaling with scale-down, audit logs GA, enhanced SSO, and Milvus 2.6 in Private Preview.

Why I’m Against Claude Code’s Grep-Only Retrieval? It Just Burns Too Many Tokens
Learn how vector-based code retrieval cuts Claude Code token consumption by 40%. Open-source solution with easy MCP integration. Try claude-context today.

Milvus WebUI: A Visual Management Tool for Your Vector Database
Explore Milvus WebUI to monitor, manage, and optimize your vector database with real-time insights, performance tracking, and system health monitoring.



