PyTorch와 Milvus로 이미지 검색하기 (Image Search with PyTorch and Milvus)
이 가이드는 PyTorch와 Milvus를 통합해 임베딩으로 이미지 검색을 수행하는 예시를 소개해요. PyTorch는 머신러닝 모델을 만들고 배포하는 데 널리 쓰이는 강력한 오픈소스 딥러닝 프레임워크예요. 이 예시에서는 Torchvision 라이브러리와 사전 훈련된 ResNet50 모델을 활용해 이미지 내용을 나타내는 특징 벡터(임베딩)를 생성해요. 이 임베딩을 고성능 벡터 데이터베이스인 Milvus에 저장해 효율적인 유사도 검색을 가능하게 해요. 사용하는 데이터셋은 Kaggle의 Impressionist-Classifier Dataset이에요. PyTorch의 딥러닝 능력과 Milvus의 확장 가능한 검색 기능을 결합하면, 견고하고 효율적인 이미지 검색 시스템을 어떻게 구축하는지 보여 줘요.
바로 시작해 볼게요!
출처: Milvus 문서
본문
요구 사항 설치하기 (Installing the requirements)
이 예시에서는 Milvus 연결·사용에 pymilvus, 임베딩 모델 실행에 torch, 실제 모델과 전처리에 torchvision, 예시 데이터셋 다운로드에 gdown, 로딩 바에 tqdm을 사용해요.
pip install pymilvus torch gdown torchvision tqdm
데이터 가져오기 (Grabbing the data)
gdown으로 Google Drive에서 zip을 받고, 내장된 zipfile 라이브러리로 압축을 풀어요.
import gdown
import zipfile
url = 'https://drive.google.com/uc?id=1OYDHLEy992qu5C4C8HV5uDIkOWRTAR1_'
output = './paintings.zip'
gdown.download(url, output)
with zipfile.ZipFile("./paintings.zip","r") as zip_ref:
zip_ref.extractall("./paintings")
데이터셋 크기는 2.35GB이고, 다운로드에 걸리는 시간은 네트워크 상태에 따라 달라져요.
전역 인수 (Global Arguments)
추적과 업데이트를 쉽게 하기 위해 사용할 주요 전역 인수들이에요.
# Milvus Setup Arguments
COLLECTION_NAME = 'image_search' # Collection name
DIMENSION = 2048 # Embedding vector size in this example
MILVUS_HOST = "localhost"
MILVUS_PORT = "19530"
# Inference Arguments
BATCH_SIZE = 128
TOP_K = 3
Milvus 설정하기 (Setting up Milvus)
이제 Milvus 설정을 시작할게요. 단계는 다음과 같아요.
제공된 URI로 Milvus 인스턴스에 연결해요.
from pymilvus import connections
# Connect to the instance
connections.connect(host=MILVUS_HOST, port=MILVUS_PORT)
컬렉션이 이미 존재하면 제거해요.
from pymilvus import utility
# Remove any previous collections with the same name
if utility.has_collection(COLLECTION_NAME):
utility.drop_collection(COLLECTION_NAME)
ID, 이미지 파일 경로, 임베딩을 담는 컬렉션을 만들어요.
from pymilvus import FieldSchema, CollectionSchema, DataType, Collection
# Create collection which includes the id, filepath of the image, and image embedding
fields = [
FieldSchema(name='id', dtype=DataType.INT64, is_primary=True, auto_id=True),
FieldSchema(name='filepath', dtype=DataType.VARCHAR, max_length=200), # VARCHARS need a maximum length, so for this example they are set to 200 characters
FieldSchema(name='image_embedding', dtype=DataType.FLOAT_VECTOR, dim=DIMENSION)
]
schema = CollectionSchema(fields=fields)
collection = Collection(name=COLLECTION_NAME, schema=schema)
새로 만든 컬렉션에 인덱스를 만들고 메모리에 로드해요.
# Create an AutoIndex index for collection
index_params = {
'metric_type':'L2',
'index_type':"IVF_FLAT",
'params':{'nlist': 16384}
}
collection.create_index(field_name="image_embedding", index_params=index_params)
collection.load()
이 단계가 끝나면 컬렉션은 데이터 삽입과 검색이 가능한 상태가 돼요. 추가되는 데이터는 자동으로 인덱싱되어 즉시 검색할 수 있어요. 데이터가 아주 새 것이라면, 아직 인덱싱 중인 데이터에 brute force 검색이 사용되므로 검색이 더 느릴 수 있어요.
데이터 삽입하기 (Inserting the data)
이 예시에서는 torch와 그 모델 허브가 제공하는 ResNet50 모델을 사용해요. 임베딩을 얻기 위해 마지막 분류 레이어를 떼어내면, 모델이 2048차원 임베딩을 만들어 줘요. torch에서 찾을 수 있는 모든 비전 모델은 여기 포함된 것과 같은 전처리를 사용해요.
이어지는 몇 단계에서 우리는 신경 쓸 게에요.
데이터를 불러와요.
import glob
# Get the filepaths of the images
paths = glob.glob('./paintings/paintings/**/*.jpg', recursive=True)
len(paths)
데이터를 배치로 전처리해요.
import torch
# Load the embedding model with the last layer removed
model = torch.hub.load('pytorch/vision:v0.10.0', 'resnet50', pretrained=True)
model = torch.nn.Sequential(*(list(model.children())[:-1]))
model.eval()
데이터를 임베딩해요.
from torchvision import transforms
# Preprocessing for images
preprocess = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
데이터를 삽입해요.
from PIL import Image
from tqdm import tqdm
# Embed function that embeds the batch and inserts it
def embed(data):
with torch.no_grad():
output = model(torch.stack(data[0])).squeeze()
collection.insert([data[1], output.tolist()])
data_batch = [[],[]]
# Read the images into batches for embedding and insertion
for path in tqdm(paths):
im = Image.open(path).convert('RGB')
data_batch[0].append(preprocess(im))
data_batch[1].append(path)
if len(data_batch[0]) % BATCH_SIZE == 0:
embed(data_batch)
data_batch = [[],[]]
# Embed and insert the remainder
if len(data_batch[0]) != 0:
embed(data_batch)
# Call a flush to index any unsealed segments.
collection.flush()
- 이 단계는 임베딩에 시간이 걸리므로 비교적 시간이 오래 걸려요. 커피 한 잔 마시며 잠시 쉬어도 좋아요.
- PyTorch는 Python 3.9 이하 버전에서는 잘 동작하지 않을 수 있어요. Python 3.10 이상을 사용하는 걸 고려해 보세요.
검색 수행하기 (Performing the search)
모든 데이터를 Milvus에 삽입했으니 이제 검색을 시작할 수 있어요. 이 예시에서는 두 개의 예시 이미지를 검색해요. 배치 검색이므로 검색 시간은 배치의 이미지들이 공유해요.
import glob
# Get the filepaths of the search images
search_paths = glob.glob('./paintings/test_paintings/**/*.jpg', recursive=True)
len(search_paths)
import time
from matplotlib import pyplot as plt
# Embed the search images
def embed(data):
with torch.no_grad():
ret = model(torch.stack(data))
# If more than one image, use squeeze
if len(ret) > 1:
return ret.squeeze().tolist()
# Squeeze would remove batch for single image, so using flatten
else:
return torch.flatten(ret, start_dim=1).tolist()
data_batch = [[],[]]
for path in search_paths:
im = Image.open(path).convert('RGB')
data_batch[0].append(preprocess(im))
data_batch[1].append(path)
embeds = embed(data_batch[0])
start = time.time()
res = collection.search(embeds, anns_field='image_embedding', param={'nprobe': 128}, limit=TOP_K, output_fields=['filepath'])
finish = time.time()
# Show the image results
f, axarr = plt.subplots(len(data_batch[1]), TOP_K + 1, figsize=(20, 10), squeeze=False)
for hits_i, hits in enumerate(res):
axarr[hits_i][0].imshow(Image.open(data_batch[1][hits_i]))
axarr[hits_i][0].set_axis_off()
axarr[hits_i][0].set_title('Search Time: ' + str(finish - start))
for hit_i, hit in enumerate(hits):
axarr[hits_i][hit_i + 1].imshow(Image.open(hit.entity.get('filepath')))
axarr[hits_i][hit_i + 1].set_axis_off()
axarr[hits_i][hit_i + 1].set_title('Distance: ' + str(hit.distance))
# Save the search result in a separate image file alongside your script.
plt.savefig('search_result.png')
검색 결과 이미지는 다음과 비슷해요.
더 알아보기 (Learn more)
- Milvus 벡터 인덱스 — IVF_FLAT을 포함한 인덱스 유형
- 유사도 메트릭 —
L2등 거리 메트릭 동작