Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_result_df_transform.py185 linesDownload Raw Back to utils
1import numpy as np
2from typing import List, Dict, Any, cast, Union
3from chromadb.utils.results import (
4    _transform_embeddings,
5    _add_query_fields,
6    _add_get_fields,
7    query_result_to_dfs,
8    get_result_to_df,
9)
10from chromadb.api.types import (
11    QueryResult,
12    GetResult,
13)
14from numpy.typing import NDArray
15
16
17def test_transform_embeddings() -> None:
18    # Test with None input
19    assert _transform_embeddings(None) is None
20
21    # Test with numpy arrays
22    embeddings = cast(
23        List[NDArray[Union[np.int32, np.float32]]],
24        [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
25    )
26    transformed = _transform_embeddings(embeddings)
27    assert isinstance(transformed, list)
28    assert transformed == [[1.0, 2.0], [3.0, 4.0]]
29
30    # Test with list of lists
31    embeddings = cast(
32        List[NDArray[Union[np.int32, np.float32]]],
33        [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
34    )
35    transformed = _transform_embeddings(embeddings)
36    assert transformed == [[1.0, 2.0], [3.0, 4.0]]
37
38
39def test_add_query_fields() -> None:
40    data_dict: Dict[str, Any] = {}
41    query_result: QueryResult = {
42        "ids": [["id1"], ["id2"]],
43        "embeddings": [[np.array([1.0, 2.0])], [np.array([3.0, 4.0])]],
44        "documents": [["doc1"], ["doc2"]],
45        "metadatas": [[{"key": "value1"}], [{"key": "value2"}]],
46        "distances": [[0.1], [0.2]],
47        "uris": [["uri1", "uri2"]],
48        "data": [
49            [np.array([1, 2, 3]), np.array([4, 5, 6])]
50        ],  # Using numpy arrays as Image type
51        "included": ["embeddings", "documents", "metadatas", "distances"],
52    }
53
54    _add_query_fields(data_dict, query_result, 0)
55    assert np.array_equal(data_dict["embedding"], [np.array([1.0, 2.0])])
56    assert data_dict["document"] == ["doc1"]
57    assert data_dict["metadata"] == [{"key": "value1"}]
58    assert data_dict["distance"] == [0.1]
59
60
61def test_add_get_fields() -> None:
62    data_dict: Dict[str, Any] = {}
63    get_result: GetResult = {
64        "ids": ["id1", "id2"],
65        "embeddings": [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
66        "documents": ["doc1", "doc2"],
67        "metadatas": [{"key": "value1"}, {"key": "value2"}],
68        "uris": ["uri1", "uri2"],
69        "data": [
70            np.array([1, 2, 3]),
71            np.array([4, 5, 6]),
72        ],  # Using numpy arrays as Image type
73        "included": ["embeddings", "documents", "metadatas"],
74    }
75
76    _add_get_fields(data_dict, get_result)
77    assert all(
78        np.array_equal(a, b)
79        for a, b in zip(
80            data_dict["embedding"], [np.array([1.0, 2.0]), np.array([3.0, 4.0])]
81        )
82    )
83    assert data_dict["document"] == ["doc1", "doc2"]
84    assert data_dict["metadata"] == [{"key": "value1"}, {"key": "value2"}]
85
86
87def test_query_result_to_dfs() -> None:
88    query_result: QueryResult = {
89        "ids": [["id1", "id2"]],
90        "embeddings": [[np.array([1.0, 2.0]), np.array([3.0, 4.0])]],
91        "documents": [["doc1", "doc2"]],
92        "metadatas": [[{"key": "value1"}, {"key": "value2"}]],
93        "distances": [[0.1, 0.2]],
94        "uris": [["uri1", "uri2"]],
95        "data": [
96            [np.array([1, 2, 3]), np.array([4, 5, 6])]
97        ],  # Using numpy arrays as Image type
98        "included": ["embeddings", "documents", "metadatas", "distances"],
99    }
100
101    dfs = query_result_to_dfs(query_result)
102    assert len(dfs) == 1  # Only one query
103
104    # Test DataFrame
105    df = dfs[0]
106    assert df.index[0] == "id1"
107    assert df["document"].iloc[0] == "doc1"
108    assert df["metadata"].iloc[0] == {"key": "value1"}
109    assert np.array_equal(df["embedding"].iloc[0], np.array([1.0, 2.0]))
110    assert df["distance"].iloc[0] == 0.1
111
112    # Test column order
113    assert list(df.columns) == ["embedding", "document", "metadata", "distance"]
114
115
116def test_get_result_to_df() -> None:
117    get_result: GetResult = {
118        "ids": ["id1", "id2"],
119        "embeddings": [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
120        "documents": ["doc1", "doc2"],
121        "metadatas": [{"key": "value1"}, {"key": "value2"}],
122        "uris": ["uri1", "uri2"],
123        "data": [
124            np.array([1, 2, 3]),
125            np.array([4, 5, 6]),
126        ],  # Using numpy arrays as Image type
127        "included": ["embeddings", "documents", "metadatas"],
128    }
129
130    df = get_result_to_df(get_result)
131    assert len(df) == 2
132    assert list(df.index) == ["id1", "id2"]
133    assert df["document"].tolist() == ["doc1", "doc2"]
134    assert df["metadata"].tolist() == [{"key": "value1"}, {"key": "value2"}]
135    assert all(
136        np.array_equal(a, b)
137        for a, b in zip(
138            df["embedding"].tolist(), [np.array([1.0, 2.0]), np.array([3.0, 4.0])]
139        )
140    )
141
142    # Test column order
143    assert list(df.columns) == ["embedding", "document", "metadata"]
144
145
146def test_query_result_to_dfs_with_missing_fields() -> None:
147    query_result: QueryResult = {
148        "ids": [["id1"]],
149        "documents": [["doc1"]],
150        "embeddings": [[]],  # type:ignore
151        "metadatas": [[]],
152        "distances": [[]],
153        "uris": [[]],
154        "data": [[]],
155        "included": ["documents"],
156    }
157
158    dfs = query_result_to_dfs(query_result)
159    assert len(dfs) == 1
160    df = dfs[0]
161    assert df.index[0] == "id1"
162    assert df["document"].iloc[0] == "doc1"
163    assert "metadata" not in df.columns
164    assert "embedding" not in df.columns
165    assert "distance" not in df.columns
166
167
168def test_get_result_to_df_with_missing_fields() -> None:
169    get_result: GetResult = {
170        "ids": ["id1", "id2"],
171        "documents": ["doc1", "doc2"],
172        "embeddings": [],
173        "metadatas": [],
174        "uris": [],
175        "data": [],
176        "included": ["documents"],
177    }
178
179    df = get_result_to_df(get_result)
180    assert len(df) == 2
181    assert list(df.index) == ["id1", "id2"]
182    assert df["document"].tolist() == ["doc1", "doc2"]
183    assert "metadata" not in df.columns
184    assert "embedding" not in df.columns
185 
codekingpro/portable-devtools · Team Ai