ChromaDB 임베딩 함수

ChromaDB 임베딩 함수 (Embedding Functions)

텍스트를 그대로 저장하는 대신, ChromaDB에 넣기 전에 벡터로 바꿔야 해요. 그 "텍스트 → 벡터" 변환을 대신 해주는 게 임베딩 함수(embedding function) 예요. 컬렉션에 연결해 두면 add, update, upsert, query를 호출할 때마다 자동으로 사용돼요.

출처: 공식문서 - Embedding Functions

컬렉션에 임베딩 함수 연결하기

컬렉션을 만들 때 embedding_function 인자로 임베딩 함수를 넘기면 돼요. 예를 들어 OpenAI 임베딩 함수를 쓰려면 이렇게 해요. 이때 OPENAI_API_KEY 환경 변수를 미리 설정해 둬야 해요.

from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction

collection = client.create_collection(
    name="my_collection",
    embedding_function=OpenAIEmbeddingFunction(
        model_name="text-embedding-3-small"
    )
)

# Chroma가 OpenAIEmbeddingFunction으로 문서를 임베딩해줘요
collection.add(
    ids=["id1", "id2"],
    documents=["doc1", "doc2"]
)

임베딩 함수는 디버깅에 편하게 직접 호출할 수도 있어요. 기본 임베딩 함수인 DefaultEmbeddingFunction으로 벡터를 뽑아 바로 query에 넘길 수 있죠.

from chromadb.utils.embedding_functions import DefaultEmbeddingFunction

default_ef = DefaultEmbeddingFunction()
embeddings = default_ef(["foo"])
print(embeddings) # [[0.05035809800028801, 0.0626462921500206, -0.061827320605516434...]]

collection.query(query_embeddings=embeddings)

커스텀 임베딩 함수

ChromaDB와 맞춰 쓸 커스텀 임베딩 함수를 만들려면 EmbeddingFunction 인터페이스를 구현하면 돼요. __call__에서 입력 문서들을 벡터로 바꿔 돌려주면 되고, 저장·복원을 위한 메서드도 함께 정의해야 해요.

from typing import Dict, Any
from chromadb import Documents, EmbeddingFunction, Embeddings
from chromadb.utils.embedding_functions import register_embedding_function

@register_embedding_function
class MyEmbeddingFunction(EmbeddingFunction):

    def __init__(self, model):
        self.model = model

    def __call__(self, input: Documents) -> Embeddings:
        # embed the documents somehow
        return embeddings

    @staticmethod
    def name() -> str:
        return "my-ef"

    def get_config(self) -> Dict[str, Any]:
        return dict(model=self.model)

    @staticmethod
    def build_from_config(config: Dict[str, Any]) -> "EmbeddingFunction":
        return MyEmbeddingFunction(config['model'])

기본값: all-MiniLM-L6-v2

컬렉션을 만들 때 임베딩 함수를 지정하지 않으면 ChromaDB가 DefaultEmbeddingFunction으로 설정해요. 이 함수는 로컬에서 실행되며, 필요하면 모델 파일을 자동으로 내려받아요.

collection = client.create_collection(name="my_collection")

더 알아보기 (Learn more)