Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_multimodal_ef.py170 linesDownload Raw Back to ef
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 
codekingpro/portable-devtools · Team Ai