diff --git a/reme2/component/embedding/base_embedding_model.py b/reme2/component/embedding/base_embedding_model.py index c3493978..6bf73b99 100644 --- a/reme2/component/embedding/base_embedding_model.py +++ b/reme2/component/embedding/base_embedding_model.py @@ -28,7 +28,7 @@ class BaseEmbeddingModel(BaseComponent): pass_dimensions: bool = False, max_batch_size: int = 10, max_input_length: int = 8192, - max_cache_size: int = 5000, + max_cache_size: int = 10000, enable_cache: bool = True, max_retries: int = 3, **kwargs, @@ -99,7 +99,7 @@ class BaseEmbeddingModel(BaseComponent): for orig_idx, text, emb in zip(indices, texts, embeddings): if emb is None: continue - emb_array = np.asarray(emb, dtype=np.float32) + emb_array = np.asarray(emb, dtype=np.float16) if len(emb_array) != self.dimensions: if len(emb_array) < self.dimensions: emb_array = np.pad(emb_array, (0, self.dimensions - len(emb_array))) @@ -172,7 +172,7 @@ class BaseEmbeddingModel(BaseComponent): continue if len(self._embedding_cache) >= self.max_cache_size: break - self._embedding_cache[str(key)] = emb.astype(np.float32) + self._embedding_cache[str(key)] = emb.astype(np.float16) def _save_cache(self) -> None: if not self.enable_cache or not self._embedding_cache: diff --git a/reme2/component/file_store/local_file_store.py b/reme2/component/file_store/local_file_store.py index 9292e850..f523bc52 100644 --- a/reme2/component/file_store/local_file_store.py +++ b/reme2/component/file_store/local_file_store.py @@ -133,12 +133,12 @@ class LocalFileStore(BaseFileStore): if not query_embedding: return [] - candidates = [c for c in self.file_chunks.values() if c.embedding] + candidates = [c for c in self.file_chunks.values() if c.embedding is not None] if not candidates: return [] - candidate_embeddings = np.array([c.embedding for c in candidates]) - similarities = batch_cosine_similarity(np.array([query_embedding]), candidate_embeddings)[0] + candidate_embeddings = np.stack([c.embedding for c in candidates]) + similarities = batch_cosine_similarity(query_embedding.reshape(1, -1), candidate_embeddings)[0] results = [ c.model_copy(update={"scores": {"vector": float(s), "score": float(s)}}) diff --git a/reme2/schema/emb_node.py b/reme2/schema/emb_node.py index 36011d23..7a261d3a 100644 --- a/reme2/schema/emb_node.py +++ b/reme2/schema/emb_node.py @@ -1,10 +1,26 @@ from uuid import uuid4 -from pydantic import BaseModel, Field +import numpy as np +from pydantic import BaseModel, Field, field_serializer, field_validator, ConfigDict class EmbNode(BaseModel): + model_config = ConfigDict(arbitrary_types_allowed=True) + id: str = Field(default_factory=lambda: uuid4().hex) text: str = Field(default="") - embedding: list[float] | None = Field(default=None) + embedding: np.ndarray | None = Field(default=None) metadata: dict = Field(default_factory=dict) + + @field_validator('embedding', mode='before') + @classmethod + def validate_embedding(cls, v): + if v is None: + return v + return np.array(v, dtype=np.float16) + + @field_serializer('embedding') + def serialize_embedding(self, v: np.ndarray | None, _info): + if v is None: + return None + return v.tolist() diff --git a/reme2/schema/file_chunk.py b/reme2/schema/file_chunk.py index 4f8077cc..b30d803a 100644 --- a/reme2/schema/file_chunk.py +++ b/reme2/schema/file_chunk.py @@ -16,4 +16,4 @@ class FileChunk(EmbNode): def set_hash_id(self): from ..utils import hash_text self.id = hash_text(" ".join([self.path, str(self.start_line), str(self.end_line), self.text])) - return self.id \ No newline at end of file + return self \ No newline at end of file