codekingpro/portable-devtools
114k
1import os
2from typing import Generator, cast
3import numpy as np
4import pytest
5import chromadb
6from chromadb.api.types import (
7 Embeddable,
8 EmbeddingFunction,
9 Embeddings,
10 Image,
11 Document,
12)
13from chromadb.test.property.strategies import hashing_embedding_function
14from chromadb.test.property.invariants import _exact_distances
15from chromadb.config import Settings
16
17
18# A 'standard' multimodal embedding function, which converts inputs to strings
19# then hashes them to a fixed dimension.
20class hashing_multimodal_ef(EmbeddingFunction[Embeddable]):
21 def __init__(self) -> None:
22 self._hef = hashing_embedding_function(dim=10, dtype=np.float64)
23
24 def __call__(self, input: Embeddable) -> Embeddings:
25 to_texts = [str(i) for i in input]
26 embeddings = np.array(self._hef(to_texts))
27 # Normalize the embeddings
28 # This is so we can generate random unit vectors and have them be close to the embeddings
29 embeddings /= np.linalg.norm(embeddings, axis=1, keepdims=True) # type: ignore[misc]
30 return cast(Embeddings, embeddings.tolist())
31
32
33def random_image() -> Image:
34 return np.random.randint(0, 255, size=(10, 10, 3), dtype=np.int64)
35
36
37def random_document() -> Document:
38 return str(random_image())
39
40
41@pytest.fixture
42def multimodal_collection(
43 default_ef: EmbeddingFunction[Embeddable] = hashing_multimodal_ef(),
44) -> Generator[chromadb.Collection, None, None]:
45 settings = Settings()
46 if os.environ.get("CHROMA_INTEGRATION_TEST_ONLY"):
47 host = os.environ.get("CHROMA_SERVER_HOST", "localhost")
48 port = int(os.environ.get("CHROMA_SERVER_HTTP_PORT", 0))
49 settings.chroma_api_impl = "chromadb.api.fastapi.FastAPI"
50 settings.chroma_server_http_port = port
51 settings.chroma_server_host = host
52
53 client = chromadb.Client(settings=settings)
54 collection = client.create_collection(
55 name="multimodal_collection", embedding_function=default_ef
56 )
57 yield collection
58 client.clear_system_cache()
59
60
61# Test adding and querying of a multimodal collection consisting of images and documents
62def test_multimodal(
63 multimodal_collection: chromadb.Collection,
64 default_ef: EmbeddingFunction[Embeddable] = hashing_multimodal_ef(),
65 n_examples: int = 10,
66 n_query_results: int = 3,
67) -> None:
68 # Fix numpy's random seed for reproducibility
69 random_state = np.random.get_state()
70 np.random.seed(0)
71
72 image_ids = [str(i) for i in range(n_examples)]
73 images = [random_image() for _ in range(n_examples)]
74 image_embeddings = default_ef(images)
75
76 document_ids = [str(i) for i in range(n_examples, 2 * n_examples)]
77 documents = [random_document() for _ in range(n_examples)]
78 document_embeddings = default_ef(documents)
79
80 # Trying to add a document and an image at the same time should fail
81 with pytest.raises(
82 ValueError,
83 # This error string may be in any order
84 match=r"Exactly one of (images|documents|uris)(?:, (images|documents|uris))?(?:, (images|documents|uris))? must be provided in add\.",
85 ):
86 multimodal_collection.add(
87 ids=image_ids[0], documents=documents[0], images=images[0]
88 )
89
90 # Add some documents
91 multimodal_collection.add(ids=document_ids, documents=documents)
92 # Add some images
93 multimodal_collection.add(ids=image_ids, images=images)
94
95 # get() should return all the documents and images
96 # ids corresponding to images should not have documents
97 get_result = multimodal_collection.get(include=["documents"])
98 assert len(get_result["ids"]) == len(document_ids) + len(image_ids)
99 for i, id in enumerate(get_result["ids"]):
100 assert id in document_ids or id in image_ids
101 assert get_result["documents"] is not None
102 if id in document_ids:
103 assert get_result["documents"][i] == documents[document_ids.index(id)]
104 if id in image_ids:
105 assert get_result["documents"][i] is None
106
107 # Generate a random query image
108 query_image = random_image()
109 query_image_embedding = default_ef([query_image])
110
111 image_neighbor_indices, _ = _exact_distances(
112 query_image_embedding, image_embeddings + document_embeddings
113 )
114 # Get the ids of the nearest neighbors
115 nearest_image_neighbor_ids = [
116 image_ids[i] if i < n_examples else document_ids[i % n_examples]
117 for i in image_neighbor_indices[0][:n_query_results]
118 ]
119
120 # Generate a random query document
121 query_document = random_document()
122 query_document_embedding = default_ef([query_document])
123 document_neighbor_indices, _ = _exact_distances(
124 query_document_embedding, image_embeddings + document_embeddings
125 )
126 nearest_document_neighbor_ids = [
127 image_ids[i] if i < n_examples else document_ids[i % n_examples]
128 for i in document_neighbor_indices[0][:n_query_results]
129 ]
130
131 # Querying with both images and documents should fail
132 with pytest.raises(ValueError):
133 multimodal_collection.query(
134 query_images=[query_image], query_texts=[query_document]
135 )
136
137 # Query with images
138 query_result = multimodal_collection.query(
139 query_images=[query_image], n_results=n_query_results, include=["documents"]
140 )
141
142 assert query_result["ids"][0] == nearest_image_neighbor_ids
143
144 # Query with documents
145 query_result = multimodal_collection.query(
146 query_texts=[query_document], n_results=n_query_results, include=["documents"]
147 )
148
149 assert query_result["ids"][0] == nearest_document_neighbor_ids
150 np.random.set_state(random_state)
151
152
153@pytest.mark.xfail
154def test_multimodal_update_with_image(
155 multimodal_collection: chromadb.Collection,
156) -> None:
157 # Updating an entry with an existing document should remove the documentß
158
159 document = random_document()
160 image = random_image()
161 id = "0"
162
163 multimodal_collection.add(ids=id, documents=document)
164
165 multimodal_collection.update(ids=id, images=image)
166
167 get_result = multimodal_collection.get(ids=id, include=["documents"])
168 assert get_result["documents"] is not None
169 assert get_result["documents"][0] is None
170 