Usando IA para encontrar tu estilista de celebridades (Parte II)
En mi publicación de blog anterior, "Usar IA para encontrar a tu estilista de celebridades," expliqué cómo aprovechar tecnologías de inteligencia artificial (IA), como Milvus, una base de datos vectorial de código abierto nativa de IA, y modelos de Hugging Face, para encontrar elecciones de estilo de celebridades que coincidan con las tuyas. En esta publicación de seguimiento, iremos un paso más allá y demostraremos cómo obtener resultados más detallados y precisos haciendo algunos cambios de código al proyecto anterior. Además, proporcionaré sugerencias sobre cómo puedes ampliar este proyecto por tu cuenta.
Si deseas probar este proyecto directamente, descarga las fotos y el notebook completado. Si te interesa el proyecto analizado en mi blog anterior, puedes consultar las fotos que usa y su tutorial.
Recapitulando el tutorial de mi proyecto anterior de IA para moda
Antes de profundizar en este proyecto, permíteme recapitular brevemente el tutorial que analizamos en mi publicación anterior. Así, no tienes que salir de esta página para conocer el contexto.
Importar todas las bibliotecas necesarias para la manipulación de imágenes
Comenzamos el código importando todas las bibliotecas necesarias para la manipulación de imágenes, incluido torch para la extracción de características, el objeto segformer de transformers, matplotlib y algunas importaciones de torchvision, como Resize, masks_to_boxes y 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
Preprocesar las imágenes de celebridades
Después de importar todos los paquetes necesarios para la manipulación de imágenes, puedes comenzar a procesar tus imágenes. Las siguientes tres funciones (get_segmentation, get_masks y crop_images) se usan para segmentar prendas de vestir y recortarlas para una ingesta posterior.
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
Almacenar los datos de imagen en una base de datos vectorial
Usamos Milvus, una base de datos vectorial de código abierto nativa de IA, para almacenar datos de imágenes. Para comenzar, descomprime el archivo zip photos para este proyecto e incluye la carpeta en el mismo directorio raíz que el notebook. Una vez completado este paso, puedes ejecutar el código siguiente para procesar las imágenes y almacenar los datos en 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()
A continuación, puedes ejecutar el código siguiente para generar embeddings usando el modelo Nvidia ResNet 50 de 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 función siguiente define cómo incrustar e insertar datos. Después de eso, el código recorre todas las imágenes, las incrusta y las inserta en Milvus.
Nota: Muchos de los componentes siguientes cambiarán o se eliminarán al utilizar la nueva función de esquema dinámico 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()
Consultar la base de datos vectorial
El siguiente código demuestra cómo consultar Milvus con imágenes de entrada y recuperar los tres mejores resultados para cada prenda de ropa.
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)
Más coincidencia de patrones: seleccionar objetos de cada imagen
Siguiendo el tutorial resumido arriba, puedes descubrir las tres principales coincidencias de estilo de celebridades para cada prenda que busques. También puedes crear una imagen como la de abajo sin los cuadros delimitadores de los artículos coincidentes. En esta sección, explicaré cómo encontrar estilos de moda con patrones más cercanos a los tuyos, con unos pocos cambios de código respecto al utilizado en el tutorial anterior.
imagen
Importar todas las bibliotecas necesarias para la manipulación de imágenes
Para empezar, importa en tu código todas las bibliotecas necesarias para la manipulación de imágenes. Si ya lo has hecho, puedes omitir este paso.
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
Preprocesar tus imágenes
Cuando hayas importado todos los paquetes necesarios para la manipulación de imágenes, continúa con el proceso de segmentación de imágenes, que implica tres funciones: get_segmentation, get_masks y crop_images.
No tenemos que hacer ningún cambio de código en la función get_segmentation.
Para la función get_masks, solo necesitamos tomar las segmentaciones que corresponden a los ID de segmentación en la lista wanted. Es una nueva adición que incluye ID de segmentación para prendas de vestir, como se indica en la tarjeta del modelo en Hugging Face.
Haremos la mayoría de los cambios de código en la función crop_image. En mi tutorial anterior, esta función devolvía una lista de imágenes recortadas. Después de cambiar algo de código, ahora devuelve tres objetos: embeddings de las imágenes recortadas, una lista de las coordenadas de los cuadros en la imagen original y una lista de los ID de segmentación. Esta nueva configuración mueve el embedding de la inserción por lotes al paso de transformación.
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
devuelve dos listas masks (tensor) y obj_ids (int)
"mattmdjaga/segformer_b2_clothes" de 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
Ahora que tenemos las imágenes, es hora de cargarlas. Este paso implica una inserción por lotes, que cubrimos en mi tutorial anterior. En este tutorial, insertaremos todos nuestros datos de una sola vez como una lista de diccionarios en lugar de una lista de listas. Me parece que este método de inserción es mucho más limpio, y nos permite añadir un nuevo campo al esquema en el momento de la inserción. En este caso, añadiremos una lista de esquinas de recorte.
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()
### Consultar la base de datos vectorial
Ahora, es hora de hacer consultas en Milvus, nuestra [base de datos vectorial](https://zilliz.com/learn/what-is-vector-database). En comparación con los pasos que usamos en el tutorial anterior, aquí hay algunas diferencias:
- Primero, limitamos a cinco el número de "coincidencias" que nos interesan en una imagen.
- Segundo, mostramos las tres imágenes coincidentes más cercanas.
- Tercero, añadimos una función para obtener un mapa de colores para dibujar cuadros delimitadores en diferentes colores.
Ahora, configuramos la figura y los ejes de `matplotlib`. Luego, recorremos todas nuestras imágenes y aplicamos las tres funciones de procesamiento mencionadas anteriormente para obtener las segmentaciones y los cuadros delimitadores.
Después de haber preprocesado las imágenes, podemos buscarlas en Milvus. Obtenemos las tres mejores respuestas para cada imagen según el número de artículos "coincidentes" que contienen. Finalmente, imprimimos los resultados junto con los cuadros delimitadores que devolvieron coincidencias.
from pprint import pprint from PIL import ImageDraw from collections import Counter import matplotlib.patches as patches
LIMIT = 5 # Cuántas coincidencias cercanas por prenda de vestir analizar CLOSEST = 3 # Cuántas imágenes más cercanas mostrar. CLOSEST <= Limit
search_paths = ["./photos/Taylor_Swift/Taylor_Swift_2.jpg", "./photos/Jenna_Ortega/Jenna_Ortega_6.jpg"] # Imágenes que buscar
def get_cmap(n, name='hsv'): '''Devuelve una función que asigna cada índice en 0, 1, ..., n-1 a un color RGB distinto; el argumento de palabra clave name debe ser un nombre estándar de mapa de colores de mpl. Obtenido de https://stackoverflow.com/questions/14720331/how-to-generate-random-colors-in-matplotlib''' return plt.cm.get_cmap(name, n)
Crear los subgráficos de resultados
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)
Una vez que hayas completado los pasos anteriores, deberías obtener un resultado similar al que aparece a continuación o al del comienzo de esta sesión.

## ¿Qué sigue? posibles extensiones del proyecto
Estoy poniendo este proyecto en pausa para trabajar en algunos otros proyectos por ahora, ¡pero puedes ampliarlo si quieres! Aquí tienes tres posibles extensiones.
Primero, puedes desarrollar un poco más el juego de comparación. Por ejemplo, puedes agrupar artículos divididos, como marcar ambos zapatos como un solo artículo. También puedes añadir más fotos de celebridades o amigos para hacer más comparaciones.
Segundo, puedes convertir este proyecto en un identificador de moda o en un sistema de recomendaciones. En lugar de usar imágenes de celebridades, puedes usar fotos de ropa que se pueda comprar en línea. Cuando un usuario suba una foto, puedes compararla con las imágenes de tu base de datos vectorial y sugerir al usuario las prendas de vestir más parecidas.
Tercero, puedes crear un generador de estilos, lo cual puede ser más desafiante. Hay varias formas de hacerlo, pero una idea es tomar varias fotos de un usuario y generar sugerencias basadas en ellas. Este enfoque implica usar un modelo generativo de imágenes para proporcionar sugerencias de estilo y compararlas con las fotos más cercanas de usuarios como referencia. Luego podemos sugerir algo que tenga sentido en función de esta comparación.
Estas tres extensiones son solo algunos ejemplos de formas de mejorar mi proyecto simple usando un modelo de imágenes y una base de datos vectorial como [Milvus](https://zilliz.com/what-is-milvus). El uso de una base de datos vectorial permite una variedad de tareas de búsqueda por similitud, lo cual es especialmente valioso para comparar imágenes.
## Resumen
En este tutorial, ampliamos nuestro primer proyecto de estilo de celebridades usando el nuevo esquema dinámico de Milvus, filtrando ciertos IDs de segmentación y llevando un registro de los cuadros delimitadores de nuestras coincidencias. También ordenamos nuestros resultados de búsqueda para devolver los tres mejores resultados según el número de coincidencias.
El nuevo esquema dinámico de Milvus nos permite agregar campos adicionales cuando cargamos datos usando un formato de diccionario, cambiando la forma en que inicialmente cargábamos por lotes una lista de listas. También facilitó agregar coordenadas de recorte sin cambiar el esquema.
Como nuevo paso de preprocesamiento, filtramos ciertos IDs que no están relacionados con ropa según la model card en Hugging Face. Filtramos estos IDs en la función `get_masks`. Dato curioso: el objeto `obj_ids` en esa función es en realidad un tensor.
También llevamos un registro de los cuadros delimitadores. Movimos el paso de embeddings a la función de recorte de imágenes y devolvimos los embeddings con los cuadros delimitadores y los IDs de segmentación. Luego, guardamos estos embeddings en Milvus usando un esquema dinámico.
En el momento de la consulta, agregamos todas las imágenes devueltas según el número de cuadros delimitadores que contenían, lo que nos permite encontrar la imagen de celebridad más parecida mediante diferentes prendas de vestir.
Ahora depende de ti. Puedes tomar mis sugerencias y crear algo distinto con ellas, como un sistema de recomendación de moda, un mejor sistema de comparación de estilos para ti y tus amigos, o una app de IA generativa de moda.
Sigue leyendo

Introducing Functions and Model Inference on Zilliz Cloud: Automatic Embedding and Reranking with Hosted Models
Zilliz Cloud Functions auto-generate embeddings via OpenAI, Voyage AI, Cohere, or Zilliz Hosted Models. Built-in reranking — just insert text and search.

Introducing Business Critical Plan: Enterprise-Grade Security and Compliance for Mission-Critical AI Applications
Discover Zilliz Cloud’s Business Critical Plan—offering advanced security, compliance, and uptime for mission-critical AI and vector database workloads.

Zilliz Cloud Update: Tiered Storage, Business Critical Plan, Cross-Region Backup, and Pricing Changes
This release offers a rebuilt tiered storage with lower costs, a new Business Critical plan for enhanced security, and pricing updates, among other features.



