Usando IA para encontrar seu estilista de celebridades (Parte II)
Na minha postagem anterior no blog, "Usando IA para encontrar seu stylist de celebridades," expliquei como aproveitar tecnologias de inteligência artificial (IA), como Milvus, um banco de dados vetorial open-source nativo de IA, e modelos do Hugging Face, para encontrar escolhas de estilo de celebridades que combinem com as suas. Nesta postagem de acompanhamento, daremos um passo além e demonstraremos como obter resultados mais detalhados e precisos fazendo algumas alterações no código do projeto anterior. Além disso, fornecerei sugestões de como você pode estender este projeto por conta própria.
Se você quiser experimentar este projeto diretamente, baixe as fotos e o notebook concluído. Se você tiver interesse no projeto discutido no meu blog anterior, pode conferir as fotos que ele usa e seu tutorial.
Recapitulando o tutorial do meu projeto anterior de IA de moda
Antes de mergulharmos fundo neste projeto, deixe-me recapitular brevemente o tutorial que discutimos na minha postagem anterior. Assim, você não precisa sair desta página para entender o contexto.
Importando todas as bibliotecas necessárias para manipulação de imagens
Começamos o código importando todas as bibliotecas necessárias para manipulação de imagens, incluindo torch para extração de recursos, o objeto segformer de transformers, matplotlib e algumas importações de torchvision, como 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
Pré-processando as imagens de celebridades
Depois de importar todos os pacotes necessários para manipulação de imagens, você pode começar a processar suas imagens. As três funções a seguir (get_segmentation, get_masks e crop_images) são usadas para segmentar peças de roupa e recortá-las para ingestão 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
Armazenar os dados de imagem em um banco de dados vetorial
Usamos Milvus, um banco de dados vetorial nativo de IA de código aberto, para armazenar dados de imagens. Para começar, descompacte o arquivo zip photos deste projeto e inclua a pasta no mesmo diretório raiz do notebook. Depois que esta etapa for concluída, você poderá executar o código abaixo para processar as imagens e armazenar os dados no 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()
Em seguida, você pode executar o código abaixo para gerar embeddings usando o modelo Nvidia ResNet 50 do 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()
A função abaixo define como incorporar e inserir dados. Em seguida, o código percorre todas as imagens, as incorpora e as insere no Milvus.
Observação: Muitos dos componentes abaixo serão alterados ou removidos ao utilizar o novo recurso de esquema dinâmico do 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()
Consulte o banco de dados vetorial
O código a seguir demonstra como consultar o Milvus com imagens de entrada e recuperar os três principais resultados para cada peça de roupa.
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)
Mais correspondência de padrões: selecionando objetos de cada imagem
Seguindo o tutorial recapitulado acima, você pode descobrir as três principais correspondências de estilo de celebridades para cada peça de roupa que pesquisar. Você também pode criar uma imagem como a abaixo sem as caixas delimitadoras dos itens correspondentes. Nesta seção, explicarei como encontrar estilos de moda com padrões mais próximos dos seus, com algumas alterações de código em relação ao usado no tutorial anterior.
imagem
Importando todas as bibliotecas necessárias para manipulação de imagens
Para começar, importe todas as bibliotecas necessárias para manipulação de imagens no seu código. Se você já tiver feito isso, pode pular esta etapa.
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é-processando suas imagens
Depois de importar todos os pacotes necessários para manipulação de imagens, prossiga com o processo de segmentação de imagens, que envolve três funções: get_segmentation, get_masks e crop_images.
Não precisamos fazer nenhuma alteração de código na função get_segmentation.
Para a função get_masks, precisamos apenas pegar as segmentações que correspondem aos IDs de segmentação na lista wanted. Trata-se de uma nova adição que inclui IDs de segmentação para peças de roupa, conforme declarado no model card no Hugging Face.
Faremos a maior parte das alterações de código na função crop_image. No meu tutorial anterior, essa função retornava uma lista de imagens recortadas. Depois de alterarmos algum código, ela agora retorna três objetos: embeddings das imagens recortadas, uma lista das coordenadas das caixas na imagem original e uma lista dos IDs de segmentação. Essa nova configuração move o embedding da inserção em lote para a etapa de transformação.
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
# retorna duas listas masks (tensor) e obj_ids (int)
# "mattmdjaga/segformer_b2_clothes" do 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
Agora que temos as imagens, é hora de carregá-las. Esta etapa envolve inserção em lote, que abordamos no meu tutorial anterior. Neste tutorial, inseriremos todos os nossos dados de uma só vez como uma lista de dicionários, em vez de uma lista de listas. Acho esse método de inserção muito mais limpo, e ele nos permite adicionar um novo campo ao esquema no momento da inserção. Neste caso, adicionaremos uma lista de cantos 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()
Consulte o banco de dados vetorial
Agora, é hora de fazer consultas no Milvus, nosso banco de dados vetorial. Em comparação com as etapas que usamos no tutorial anterior, há algumas diferenças aqui:
- Primeiro, limitamos a cinco o número de "correspondências" com as quais nos importamos em uma imagem.
- Segundo, mostramos as três imagens correspondentes mais próximas.
- Terceiro, adicionamos uma função para obter um mapa de cores para desenhar caixas delimitadoras em cores diferentes.
Agora, configuramos a figura e os eixos do matplotlib. Em seguida, percorremos todas as nossas imagens e aplicamos as três funções de processamento mencionadas acima para obter as segmentações e as caixas delimitadoras.
Depois de pré-processarmos as imagens, podemos pesquisá-las no Milvus. Obtemos as três principais respostas para cada imagem com base no número de artigos "correspondentes" que elas contêm. Por fim, imprimimos os resultados junto com as caixas delimitadoras que retornaram correspondências.
from pprint import pprint
from PIL import ImageDraw
from collections import Counter
import matplotlib.patches as patches
LIMIT = 5 # Quantas correspondências mais próximas por peça de roupa analisar
CLOSEST = 3 # Quantas imagens mais próximas exibir. CLOSEST <= Limit
search_paths = ["./photos/Taylor_Swift/Taylor_Swift_2.jpg", "./photos/Jenna_Ortega/Jenna_Ortega_6.jpg"] # Imagens a pesquisar
def get_cmap(n, name='hsv'):
'''Retorna uma função que mapeia cada índice em 0, 1, ..., n-1 para uma cor
RGB distinta; o argumento nomeado name deve ser um nome padrão de mapa de cores do mpl.
Fonte: https://stackoverflow.com/questions/14720331/how-to-generate-random-colors-in-matplotlib'''
return plt.cm.get_cmap(name, n)
# Crie os subplots de resultado
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)
Depois de concluir as etapas acima, você deverá obter um resultado semelhante ao mostrado abaixo ou ao do início desta sessão.
imagem
O que vem a seguir? possíveis extensões do projeto
Estou colocando este projeto em pausa para trabalhar em alguns outros projetos por enquanto, mas você pode estendê-lo se quiser! Aqui estão três possíveis extensões.
Primeiro, você pode desenvolver um pouco mais o jogo de comparação. Por exemplo, você pode agrupar itens separados, como marcar ambos os sapatos como um único item. Você também pode adicionar mais fotos de celebridades ou amigos para mais comparações.
Segundo, você pode transformar este projeto em um identificador de moda ou em um sistema de recomendação. Em vez de usar imagens de celebridades, você pode usar fotos de roupas que podem ser compradas online. Quando um usuário carrega uma foto, você pode compará-la com as imagens no seu banco de dados vetorial e sugerir ao usuário as peças de roupa mais parecidas.
Terceiro, você pode criar um gerador de estilo, o que pode ser mais desafiador. Há várias maneiras de fazer isso, mas uma ideia é pegar várias fotos de um usuário e gerar sugestões com base nelas. Essa abordagem envolve usar um modelo generativo de imagem para fornecer sugestões de estilo e compará-las com as fotos mais próximas de usuários para referência. Podemos então sugerir algo que faça sentido com base nessa comparação.
Essas três extensões são apenas alguns exemplos de maneiras de aprimorar meu projeto simples usando um modelo de imagem e um banco de dados vetorial como Milvus. O uso de um banco de dados vetorial permite uma variedade de tarefas de busca por similaridade, o que é especialmente valioso para comparar imagens.
Resumo
Neste tutorial, estendemos nosso primeiro projeto de estilo de celebridades usando o novo esquema dinâmico do Milvus, filtrando certos IDs de segmentação e mantendo o controle das caixas delimitadoras das nossas correspondências. Também ordenamos nossos resultados de busca para retornar os três principais resultados com base no número de correspondências.
O novo esquema dinâmico do Milvus nos permite adicionar campos extras quando carregamos dados usando um formato de dicionário, mudando a forma como inicialmente fazíamos o upload em lote de uma lista de listas. Ele também facilitou a adição de coordenadas de recorte sem alterar o esquema.
Como uma nova etapa de pré-processamento, filtramos certos IDs que não estão relacionados a roupas com base no model card no Hugging Face. Filtramos esses IDs na função get_masks. Curiosidade: o objeto obj_ids nessa função é, na verdade, um tensor.
Também mantivemos o controle das caixas delimitadoras. Movemos a etapa de embedding para a função de recorte de imagem e retornamos os embeddings com as caixas delimitadoras e os IDs de segmentação. Em seguida, salvamos esses embeddings no Milvus usando um esquema dinâmico.
No momento da consulta, agregamos todas as imagens retornadas pelo número de caixas delimitadoras que elas continham, permitindo-nos encontrar a imagem de celebridade mais próxima correspondente por meio de diferentes peças de roupa. Agora é com você. Você pode pegar minhas sugestões e criar algo diferente a partir delas, como um sistema de recomendação de moda, um sistema melhor de comparação de estilo para você e seus amigos, ou um aplicativo de IA generativa de moda.
Continue lendo

Context Engineering Strategies for AI Agents: A Developer’s Guide
Learn practical context engineering strategies for AI agents. Explore frameworks, tools, and techniques to improve reliability, efficiency, and cost.

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.

Proactive Monitoring for Vector Database: Zilliz Cloud Integrates with Datadog
we're excited to announce Zilliz Cloud's integration with Datadog, enabling comprehensive monitoring and observability for your vectorDB deployments.



