| 78 | |
| 79 | |
| 80 | class SentenceTransformersEmbedding: |
| 81 | name = "sentence_transformers" |
| 82 | |
| 83 | def __init__(self, model: str, dimensions: int) -> None: |
| 84 | from sentence_transformers import SentenceTransformer # local import |
| 85 | |
| 86 | self.model = model |
| 87 | self.dimensions = dimensions |
| 88 | kwargs: dict[str, object] = {} |
| 89 | if model.lower().startswith(("qwen/", "baai/")): |
| 90 | kwargs["trust_remote_code"] = True |
| 91 | self._model = SentenceTransformer(model, **kwargs) |
| 92 | |
| 93 | def embed(self, text: str) -> List[float]: |
| 94 | vector = self._model.encode(text or " ", normalize_embeddings=True) |
| 95 | return _resize(list(vector.tolist() if hasattr(vector, "tolist") else vector), self.dimensions) |
| 96 | |
| 97 | def embed_batch(self, texts: Iterable[str]) -> List[List[float]]: |
| 98 | items = [text or " " for text in texts] |
| 99 | vectors = self._model.encode(items, normalize_embeddings=True, convert_to_numpy=True) |
| 100 | return [_resize(list(row.tolist()), self.dimensions) for row in vectors] |
| 101 | |
| 102 | |
| 103 | class OllamaEmbedding: |
no outgoing calls
no test coverage detected