Skip to content

Commit 9ca5ca9

Browse files
committed
feat: swap embedding backend to model2vec and fix search arg-swap
1 parent a6a94f7 commit 9ca5ca9

7 files changed

Lines changed: 53 additions & 544 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,6 @@ dependencies = [
1414
"litellm==1.83.7",
1515
"pydantic==2.12.5",
1616
"qdrant-client>=1.16.1",
17-
"sentence-transformers>=5.1.2",
17+
"model2vec>=0.3.0",
1818
"tucana==0.0.72",
1919
]

src/model.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,8 @@
1-
from sentence_transformers import SentenceTransformer
1+
from model2vec import StaticModel
22

3-
MODEL_NAME = 'all-MiniLM-L6-v2'
3+
MODEL_NAME = 'minishlab/potion-base-8M'
4+
EMBEDDING_DIM = 256
45

56

6-
def load_vector_model() -> SentenceTransformer:
7-
return SentenceTransformer(MODEL_NAME)
7+
def load_vector_model() -> StaticModel:
8+
return StaticModel.from_pretrained(MODEL_NAME)

src/store/few_shots_store.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,8 +2,8 @@
22
from pathlib import Path
33
from typing import Any, List
44

5+
from model2vec import StaticModel
56
from qdrant_client import QdrantClient
6-
from sentence_transformers import SentenceTransformer
77

88
from src.logger import get_logger
99
from src.schema.few_shot_schema import FewShot
@@ -13,7 +13,7 @@
1313

1414

1515
class FewShotsStore(Store):
16-
def __init__(self, memory_client: QdrantClient, vector_model: SentenceTransformer):
16+
def __init__(self, memory_client: QdrantClient, vector_model: StaticModel):
1717
super().__init__(
1818
memory_client,
1919
vector_model,
@@ -50,7 +50,7 @@ def validate(self, payload: Any) -> FewShot:
5050
return FewShot.model_validate(payload)
5151

5252
def search(self, group_identifier: str, prompt: str, limit=5) -> List[FewShot]:
53-
return super().search(prompt, group_identifier, limit)
53+
return super().search(group_identifier, prompt, limit)
5454

5555
def find(self, group_identifier: str, identifier: str) -> FewShot | None:
5656
return super().find(group_identifier, identifier)

src/store/flow_type_store.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,8 @@
11
import re
22
from typing import List, Any
33

4+
from model2vec import StaticModel
45
from qdrant_client import QdrantClient
5-
from sentence_transformers import SentenceTransformer
66

77
from src.schema.data_type_schema import DataType
88
from src.schema.flow_type_schema import FlowType
@@ -81,7 +81,7 @@ def replacement(match):
8181

8282

8383
class FlowTypeStore(Store):
84-
def __init__(self, memory_client: QdrantClient, vector_model: SentenceTransformer):
84+
def __init__(self, memory_client: QdrantClient, vector_model: StaticModel):
8585
super().__init__(
8686
memory_client,
8787
vector_model,
@@ -105,7 +105,7 @@ def validate(self, payload: Any) -> FlowType:
105105
return FlowType.model_validate(payload)
106106

107107
def search(self, group_identifier: str, prompt: str, limit=5) -> List[FlowType]:
108-
return super().search(prompt, group_identifier, limit)
108+
return super().search(group_identifier, prompt, limit)
109109

110110
def find(self, group_identifier: str, identifier: str) -> FlowType | None:
111111
return super().find(group_identifier, identifier)

src/store/function_store.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,8 @@
11
import re
22
from typing import List, Any
33

4+
from model2vec import StaticModel
45
from qdrant_client import QdrantClient
5-
from sentence_transformers import SentenceTransformer
66

77
from src.schema.data_type_schema import DataType
88
from src.schema.function_schema import FunctionDefinition
@@ -81,7 +81,7 @@ def replacement(match):
8181

8282

8383
class FunctionStore(Store):
84-
def __init__(self, memory_client: QdrantClient, vector_model: SentenceTransformer):
84+
def __init__(self, memory_client: QdrantClient, vector_model: StaticModel):
8585
super().__init__(
8686
memory_client,
8787
vector_model,
@@ -105,7 +105,7 @@ def validate(self, payload: Any) -> FunctionDefinition:
105105
return FunctionDefinition.model_validate(payload)
106106

107107
def search(self, group_identifier: str, prompt: str, limit=5) -> List[FunctionDefinition]:
108-
return super().search(prompt, group_identifier, limit)
108+
return super().search(group_identifier, prompt, limit)
109109

110110
def find(self, group_identifier: str, identifier: str) -> FunctionDefinition | None:
111111
return super().find(group_identifier, identifier)

src/store/vector_store.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,9 @@
55
from apscheduler.schedulers.background import BackgroundScheduler
66
from qdrant_client import QdrantClient
77
from qdrant_client.models import Distance, VectorParams, PointStruct, MatchAny, Filter, FieldCondition, MatchValue
8-
from sentence_transformers import SentenceTransformer
8+
from model2vec import StaticModel
9+
10+
from src.model import EMBEDDING_DIM
911

1012
from src.logger import get_logger
1113

@@ -16,7 +18,7 @@ class Store(ABC):
1618
def __init__(
1719
self,
1820
client: QdrantClient,
19-
model: SentenceTransformer,
21+
model: StaticModel,
2022
collection_name: str,
2123
payload_identifier: str,
2224
group_identifier: str,
@@ -39,7 +41,7 @@ def _setup_collection(self):
3941
self.client.create_collection(
4042
collection_name=self.collection_name,
4143
vectors_config=VectorParams(
42-
size=384,
44+
size=EMBEDDING_DIM,
4345
distance=Distance.COSINE
4446
),
4547
)

0 commit comments

Comments
 (0)