Xây dựng hệ thống tìm kiếm hình ảnh thương mại điện tử với Milvus và Python

Xây dựng hệ thống tìm kiếm hình ảnh thương mại điện tử với Milvus và Python

Việc triển khai tìm kiếm bằng hình ảnh trong các nền tảng thương mại điện tử đang trở thành tiêu chuẩn bắt buộc. Người dùng thường chụp ảnh sản phẩm họ quan tâm hoặc tải lên hình ảnh từ nguồn khác để tìm kiếm sản phẩm tương tự. Cơ sở dữ liệu vector Milvus cung cấp giải pháp lưu trữ và truy vấn hiệu quả cho các embedding đa chiều, giúp thực hiện tìm kiếm tương đồng nhanh chóng trên quy mô lớn.

Bài viết này hướng dẫn triển khai hệ thống tìm kiếm hình ảnh hoàn chỉnh, từ thiết lập môi trường đến viết mã triển khai. Hệ thống sử dụng mạng nơ-ron tích chập để trích xuất đặc trưng, Milvus để lập chỉ mục và truy vấn vector, cùng Python làm ngôn ngữ phát triển chính.

1. Thiết lập môi trường và khởi động Milvus

Milvus có thể triển khai qua Docker Compose để có phiên bản standalone hoạt động ngay lập tức. Cấu hình này bao gồm ba thành phần: Milvus server, Etcd cho metadata, và MinIO cho lưu trữ đối tượng.

# Tải cấu hình Docker Compose chính thức
curl -L https://github.com/milvus-io/milvus/releases/download/v2.4.0/milvus-standalone-docker-compose.yml \
     -o docker-compose.yml

# Khởi động các dịch vụ
docker-compose up -d

Sau khi khởi động, kiểm tra trạng thái container:

docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Ports}}"

Dịch vụ Milvus lắng nghe tại cổng 19530. Nếu cổng này bị chiếm dụng, chỉnh sửa ánh xạ cổng trong file docker-compose.yml.

Cài đặt các thư viện Python cần thiết:

# Tạo môi trường ảo
python -m venv image_search_env
source image_search_env/bin/activate  # macOS/Linux
# image_search_env\Scripts\activate  # Windows

# Cài đặt SDK và thư viện xử lý
pip install pymilvus==2.4.0
pip install torch torchvision pillow numpy

2. Trích xuất đặc trưng hình ảnh bằng mạng nơ-ron

Quá trình chuyển đổi hình ảnh thành vector số học sử dụng mô hình ResNet50 đã được huấn luyện trước trên tập ImageNet. Lớp trước lớp phân loại cuối cùng tạo ra vector 2048 chiều, biểu diễn nội dung ngữ nghĩa của hình ảnh.

Tạo file embedding_generator.py:

import torch
from torchvision import models, transforms
from PIL import Image
import numpy as np


class VisualEmbeddingEngine:
    """
    Engine tạo embedding vector từ hình ảnh sử dụng ResNet50.
    """
    
    IMAGENET_MEAN = [0.485, 0.456, 0.406]
    IMAGENET_STD = [0.229, 0.224, 0.225]
    EMBEDDING_DIM = 2048
    
    def __init__(self):
        # Khởi tạo mô hình ResNet50 pretrained
        backbone = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
        
        # Tách phần feature extractor (bỏ lớp fc cuối)
        self.feature_extractor = torch.nn.Sequential(*list(backbone.children())[:-1])
        self.feature_extractor.eval()
        
        # Pipeline tiền xử lý ảnh
        self.transform_pipeline = transforms.Compose([
            transforms.Resize(size=256, interpolation=transforms.InterpolationMode.BILINEAR),
            transforms.CenterCrop(size=224),
            transforms.ToTensor(),
            transforms.Normalize(mean=self.IMAGENET_MEAN, std=self.IMAGENET_STD)
        ])
    
    def generate(self, image_file):
        """
        Tạo embedding vector từ file hình ảnh.
        
        Parameters:
            image_file: Đường dẫn đến file ảnh
            
        Returns:
            numpy.ndarray: Vector 2048 chiều đã chuẩn hóa L2
        """
        # Đọc và chuyển đổi sang RGB
        pil_image = Image.open(image_file).convert(mode='RGB')
        
        # Áp dụng biến đổi
        tensor_input = self.transform_pipeline(pil_image)
        batch_input = torch.unsqueeze(tensor_input, dim=0)  # Thêm chiều batch
        
        # Suy luận không tính gradient
        with torch.inference_mode():
            embedding = self.feature_extractor(batch_input)
        
        # Định dạng lại và chuẩn hóa
        embedding = embedding.squeeze().numpy()
        normalized = embedding / np.linalg.norm(embedding)
        
        return normalized.astype(np.float32)
    
    def batch_generate(self, image_files):
        """
        Tạo embedding cho nhiều hình ảnh cùng lúc.
        """
        return np.stack([self.generate(f) for f in image_files])

3. Thiết kế collection và lập chỉ mục trong Milvus

Milvus lưu trữ vector trong các collection, tương tự như khái niệm table trong cơ sở dữ liệu quan hệ. Mỗi collection cần định nghĩa schema gồm trường khóa chính, trường vector, và các trường metadata tùy chọn.

Tạo file vector_repository.py:

from pymilvus import (
    connections,
    FieldSchema, CollectionSchema, DataType,
    Collection,
    utility
)


class ImageVectorRepository:
    """
    Repository quản lý lưu trữ và truy vấn vector hình ảnh trong Milvus.
    """
    
    COLLECTION_NAME = "product_visual_search"
    VECTOR_DIMENSION = 2048
    
    def __init__(self, host="localhost", port="19530"):
        connections.connect(alias="default", host=host, port=port)
        self.collection = self._ensure_collection()
    
    def _ensure_collection(self):
        """Đảm bảo collection tồn tại với cấu trúc phù hợp."""
        if utility.has_collection(self.COLLECTION_NAME):
            return Collection(self.COLLECTION_NAME)
        
        # Định nghĩa các trường
        fields = [
            # Khóa chính tự tăng
            FieldSchema(
                name="image_id",
                dtype=DataType.INT64,
                is_primary=True,
                auto_id=True
            ),
            # Trường vector đặc trưng
            FieldSchema(
                name="visual_embedding",
                dtype=DataType.FLOAT_VECTOR,
                dim=self.VECTOR_DIMENSION
            ),
            # Metadata: đường dẫn file gốc
            FieldSchema(
                name="file_path",
                dtype=DataType.VARCHAR,
                max_length=512
            ),
            # Metadata: danh mục sản phẩm
            FieldSchema(
                name="category",
                dtype=DataType.VARCHAR,
                max_length=128
            ),
            # Metadata: mã SKU
            FieldSchema(
                name="sku_code",
                dtype=DataType.VARCHAR,
                max_length=64
            )
        ]
        
        schema = CollectionSchema(
            fields=fields,
            description="Vector embeddings cho tìm kiếm hình ảnh sản phẩm",
            enable_dynamic_field=False
        )
        
        collection = Collection(name=self.COLLECTION_NAME, schema=schema)
        
        # Tạo index IVF_FLAT cho tìm kiếm ANN hiệu quả
        index_params = {
            "index_type": "IVF_FLAT",
            "metric_type": "L2",
            "params": {"nlist": 128}
        }
        collection.create_index(
            field_name="visual_embedding",
            index_params=index_params
        )
        
        return collection
    
    def ingest_batch(self, embeddings, metadata_list):
        """
        Nhập dữ liệu hàng loạt vào collection.
        
        Parameters:
            embeddings: numpy.ndarray shape (n, 2048)
            metadata_list: list[dict] chứa file_path, category, sku_code
        """
        entities = [
            embeddings.tolist(),
            [m["file_path"] for m in metadata_list],
            [m["category"] for m in metadata_list],
            [m["sku_code"] for m in metadata_list]
        ]
        
        self.collection.insert(entities)
        self.collection.flush()  # Đảm bảo dữ liệu được ghi ổn định
    
    def search_similar(self, query_vector, top_k=10, category_filter=None):
        """
        Tìm kiếm hình ảnh tương đồng.
        
        Parameters:
            query_vector: numpy.ndarray shape (2048,)
            top_k: Số kết quả trả về
            category_filter: Lọc theo danh mục (tùy chọn)
            
        Returns:
            list[dict]: Các kết quả với khoảng cách và metadata
        """
        search_params = {
            "metric_type": "L2",
            "params": {"nprobe": 16}
        }
        
        # Xây dựng điều kiện lọc nếu có
        expr = f'category == "{category_filter}"' if category_filter else ""
        
        results = self.collection.search(
            data=[query_vector.tolist()],
            anns_field="visual_embedding",
            param=search_params,
            limit=top_k,
            expr=expr,
            output_fields=["file_path", "category", "sku_code"]
        )
        
        # Định dạng kết quả
        formatted = []
        for hits in results:
            for hit in hits:
                formatted.append({
                    "id": hit.id,
                    "distance": hit.distance,
                    "file_path": hit.entity.get("file_path"),
                    "category": hit.entity.get("category"),
                    "sku_code": hit.entity.get("sku_code")
                })
        
        return formatted
    
    def get_stats(self):
        """Thống kê số lượng vector đã lưu."""
        self.collection.flush()
        return self.collection.num_entities

4. Tích hợp hoàn chỉnh và ví dụ sử dụng

File search_system.py kết hợp các thành phần thành pipeline hoàn chỉnh:

import os
import glob
from pathlib import Path

from embedding_generator import VisualEmbeddingEngine
from vector_repository import ImageVectorRepository


class VisualProductSearch:
    """
    Hệ thống tìm kiếm sản phẩm bằng hình ảnh tích hợp đầy đủ.
    """
    
    def __init__(self):
        self.embedding_engine = VisualEmbeddingEngine()
        self.vector_repo = ImageVectorRepository()
    
    def index_product_images(self, image_dir, category_map=None):
        """
        Lập chỉ mục toàn bộ hình ảnh sản phẩm từ thư mục.
        
        Parameters:
            image_dir: Thư mục chứa hình ảnh
            category_map: Dict mapping filename -> category (tùy chọn)
        """
        image_patterns = ["*.jpg", "*.jpeg", "*.png", "*.webp"]
        image_files = []
        for pattern in image_patterns:
            image_files.extend(glob.glob(os.path.join(image_dir, "**", pattern), recursive=True))
        
        print(f"Tìm thấy {len(image_files)} hình ảnh để lập chỉ mục")
        
        # Xử lý theo batch để tối ưu bộ nhớ
        batch_size = 32
        for i in range(0, len(image_files), batch_size):
            batch_files = image_files[i:i + batch_size]
            
            # Tạo embeddings
            embeddings = self.embedding_engine.batch_generate(batch_files)
            
            # Chuẩn bị metadata
            metadata = []
            for fp in batch_files:
                filename = Path(fp).name
                category = category_map.get(filename, "unknown") if category_map else "general"
                sku = Path(fp).stem
                
                metadata.append({
                    "file_path": fp,
                    "category": category,
                    "sku_code": sku
                })
            
            # Lưu vào Milvus
            self.vector_repo.ingest_batch(embeddings, metadata)
            print(f"Đã lập chỉ mục {min(i + batch_size, len(image_files))}/{len(image_files)}")
        
        print(f"Hoàn tất. Tổng số vector: {self.vector_repo.get_stats()}")
    
    def find_similar_products(self, query_image_path, result_count=8, filter_category=None):
        """
        Tìm sản phẩm tương tự với hình ảnh truy vấn.
        
        Parameters:
            query_image_path: Đường dẫn ảnh cần tìm kiếm
            result_count: Số lượng kết quả mong muốn
            filter_category: Chỉ tìm trong danh mục cụ thể
            
        Returns:
            list: Các sản phẩm tương tự với độ tương đồng
        """
        # Tạo embedding cho ảnh truy vấn
        query_embedding = self.embedding_engine.generate(query_image_path)
        
        # Tìm kiếm trong Milvus
        matches = self.vector_repo.search_similar(
            query_vector=query_embedding,
            top_k=result_count,
            category_filter=filter_category
        )
        
        return matches


# Ví dụ sử dụng
if __name__ == "__main__":
    search_engine = VisualProductSearch()
    
    # Lập chỉ mục dữ liệu mẫu
    # search_engine.index_product_images("./product_catalog")
    
    # Tìm kiếm với ảnh chụp từ người dùng
    # results = search_engine.find_similar_products("./user_upload/captured_item.jpg")
    # for item in results:
    #     print(f"SKU: {item['sku_code']}, Distance: {item['distance']:.4f}")

5. Tối ưu hóa và mở rộng

Thay đổi chiến lược index:

Thuật toán Độ chính xác Tốc độ Bộ nhớ Use case
FLAT Cao nhất Chậm Cao Dataset < 1 triệu
IVF_FLAT Cao Trung bình Trung bình Dataset 1-10 triệu
IVF_SQ8 Trung bình Nhanh Thấp Dataset > 10 triệu
HNSW Rất cao Rất nhanh Cao Yêu cầu latency thấp

Xử lý song song với GPU:

# Chuyển model lên GPU nếu có
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.feature_extractor = self.feature_extractor.to(device)

# Trong phương thức generate
batch_input = batch_input.to(device)

Triển khai production với Milvus Cluster:

# Chuyển từ standalone sang cluster mode
# Sử dụng Helm chart cho Kubernetes deployment
helm repo add milvus https://milvus-io.github.io/milvus-helm/
helm install milvus-cluster milvus/milvus --set cluster.enabled=true

Hệ thống đã sẵn sàng tích hợp vào API backend (FastAPI/Flask) hoặc xây dựng giao diện người dùng cho phép tải lên và hiển thị kết quả tìm kiếm.

Thẻ: Milvus vector-database image-search ResNet PyTorch

Đăng vào ngày 7 tháng 8 lúc 18:22