Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_custom_ef.py96 linesDownload Raw Back to ef
1from chromadb.api.types import EmbeddingFunction, Embeddable, Embeddings
2import numpy as np
3from typing import cast, Any
4from chromadb.utils.embedding_functions import (
5    register_embedding_function,
6    known_embedding_functions,
7)
8
9
10class LegacyCustomEmbeddingFunction(EmbeddingFunction[Embeddable]):
11    def __call__(self, input: Embeddable) -> Embeddings:
12        return cast(Embeddings, np.array([1, 2, 3]).tolist())
13
14
15class CustomEmbeddingFunction(EmbeddingFunction[Embeddable]):
16    def __call__(self, input: Embeddable) -> Embeddings:
17        return cast(Embeddings, np.array([1, 2, 3]).tolist())
18
19    def __init__(self, *args: Any, **kwargs: Any) -> None:
20        pass
21
22    @staticmethod
23    def name() -> str:
24        return "custom_embedding_function"
25
26    @staticmethod
27    def build_from_config(config: dict[str, Any]) -> "CustomEmbeddingFunction":
28        return CustomEmbeddingFunction()
29
30    def get_config(self) -> dict[str, Any]:
31        return {}
32
33
34@register_embedding_function
35class CustomEmbeddingFunctionWithRegistration(EmbeddingFunction[Embeddable]):
36    def __call__(self, input: Embeddable) -> Embeddings:
37        return cast(Embeddings, np.array([1, 2, 3]).tolist())
38
39    def __init__(self, *args: Any, **kwargs: Any) -> None:
40        pass
41
42    @staticmethod
43    def name() -> str:
44        return "custom_embedding_function_with_registration"
45
46    @staticmethod
47    def build_from_config(
48        config: dict[str, Any]
49    ) -> "CustomEmbeddingFunctionWithRegistration":
50        return CustomEmbeddingFunctionWithRegistration()
51
52    def get_config(self) -> dict[str, Any]:
53        return {}
54
55
56def test_legacy_custom_ef() -> None:
57    ef = LegacyCustomEmbeddingFunction()
58    result = ef(["test"])
59
60    # Check the structure: we expect a list with one NumPy array
61    assert isinstance(result, list), "Result should be a list"
62    assert len(result) == 1, "Result should contain exactly one element"
63    assert isinstance(result[0], np.ndarray), "Result element should be a NumPy array"
64
65    # Compare the contents of the array
66    expected = np.array([1, 2, 3], dtype=np.float32)
67    assert np.array_equal(
68        result[0], expected
69    ), f"Arrays not equal: {result[0]} vs {expected}"
70
71
72def test_custom_ef() -> None:
73    ef = CustomEmbeddingFunction()
74    result = ef(["test"])
75
76    # Same checks as above
77    assert isinstance(result, list), "Result should be a list"
78    assert len(result) == 1, "Result should contain exactly one element"
79    assert isinstance(result[0], np.ndarray), "Result element should be a NumPy array"
80
81    expected = np.array([1, 2, 3], dtype=np.float32)
82    assert np.array_equal(
83        result[0], expected
84    ), f"Arrays not equal: {result[0]} vs {expected}"
85
86
87def test_custom_ef_registration() -> None:
88    # check all 4 embedding functions for registration.
89    # LegacyCustomEmbeddingFunction should not be in known_embedding_functions
90    # CustomEmbeddingFunction should not be in known_embedding_functions
91    # CustomEmbeddingFunctionWithRegistration should be in known_embedding_functions
92
93    assert "legacy_custom_embedding_function" not in known_embedding_functions
94    assert "custom_embedding_function" not in known_embedding_functions
95    assert "custom_embedding_function_with_registration" in known_embedding_functions
96 
codekingpro/portable-devtools · Team Ai