KBaba7/llama.cpp
0
1import base642import struct3import pytest4from openai import OpenAI5from utils import *6 7server = ServerPreset.bert_bge_small()8 9EPSILON = 1e-310 11@pytest.fixture(scope="module", autouse=True)12def create_server():13 global server14 server = ServerPreset.bert_bge_small()15 16 17def test_embedding_single():18 global server19 server.pooling = 'last'20 server.start()21 res = server.make_request("POST", "/v1/embeddings", data={22 "input": "I believe the meaning of life is",23 })24 assert res.status_code == 20025 assert len(res.body['data']) == 126 assert 'embedding' in res.body['data'][0]27 assert len(res.body['data'][0]['embedding']) > 128 29 # make sure embedding vector is normalized30 assert abs(sum([x ** 2 for x in res.body['data'][0]['embedding']]) - 1) < EPSILON31 32 33def test_embedding_multiple():34 global server35 server.pooling = 'last'36 server.start()37 res = server.make_request("POST", "/v1/embeddings", data={38 "input": [39 "I believe the meaning of life is",40 "Write a joke about AI from a very long prompt which will not be truncated",41 "This is a test",42 "This is another test",43 ],44 })45 assert res.status_code == 20046 assert len(res.body['data']) == 447 for d in res.body['data']:48 assert 'embedding' in d49 assert len(d['embedding']) > 150 51 52@pytest.mark.parametrize(53 "input,is_multi_prompt",54 [55 # do not crash on empty input56 ("", False),57 # single prompt58 ("string", False),59 ([12, 34, 56], False),60 ([12, 34, "string", 56, 78], False),61 # multiple prompts62 (["string1", "string2"], True),63 (["string1", [12, 34, 56]], True),64 ([[12, 34, 56], [12, 34, 56]], True),65 ([[12, 34, 56], [12, "string", 34, 56]], True),66 ]67)68def test_embedding_mixed_input(input, is_multi_prompt: bool):69 global server70 server.start()71 res = server.make_request("POST", "/v1/embeddings", data={"input": input})72 assert res.status_code == 20073 data = res.body['data']74 if is_multi_prompt:75 assert len(data) == len(input)76 for d in data:77 assert 'embedding' in d78 assert len(d['embedding']) > 179 else:80 assert 'embedding' in data[0]81 assert len(data[0]['embedding']) > 182 83 84def test_embedding_pooling_none():85 global server86 server.pooling = 'none'87 server.start()88 res = server.make_request("POST", "/embeddings", data={89 "input": "hello hello hello",90 })91 assert res.status_code == 20092 assert 'embedding' in res.body[0]93 assert len(res.body[0]['embedding']) == 5 # 3 text tokens + 2 special94 95 # make sure embedding vector is not normalized96 for x in res.body[0]['embedding']:97 assert abs(sum([x ** 2 for x in x]) - 1) > EPSILON98 99 100def test_embedding_pooling_none_oai():101 global server102 server.pooling = 'none'103 server.start()104 res = server.make_request("POST", "/v1/embeddings", data={105 "input": "hello hello hello",106 })107 108 # /v1/embeddings does not support pooling type 'none'109 assert res.status_code == 400110 assert "error" in res.body111 112 113def test_embedding_openai_library_single():114 global server115 server.pooling = 'last'116 server.start()117 client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")118 res = client.embeddings.create(model="text-embedding-3-small", input="I believe the meaning of life is")119 assert len(res.data) == 1120 assert len(res.data[0].embedding) > 1121 122 123def test_embedding_openai_library_multiple():124 global server125 server.pooling = 'last'126 server.start()127 client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")128 res = client.embeddings.create(model="text-embedding-3-small", input=[129 "I believe the meaning of life is",130 "Write a joke about AI from a very long prompt which will not be truncated",131 "This is a test",132 "This is another test",133 ])134 assert len(res.data) == 4135 for d in res.data:136 assert len(d.embedding) > 1137 138 139def test_embedding_error_prompt_too_long():140 global server141 server.pooling = 'last'142 server.start()143 res = server.make_request("POST", "/v1/embeddings", data={144 "input": "This is a test " * 512,145 })146 assert res.status_code != 200147 assert "too large" in res.body["error"]["message"]148 149 150def test_same_prompt_give_same_result():151 server.pooling = 'last'152 server.start()153 res = server.make_request("POST", "/v1/embeddings", data={154 "input": [155 "I believe the meaning of life is",156 "I believe the meaning of life is",157 "I believe the meaning of life is",158 "I believe the meaning of life is",159 "I believe the meaning of life is",160 ],161 })162 assert res.status_code == 200163 assert len(res.body['data']) == 5164 for i in range(1, len(res.body['data'])):165 v0 = res.body['data'][0]['embedding']166 vi = res.body['data'][i]['embedding']167 for x, y in zip(v0, vi):168 assert abs(x - y) < EPSILON169 170 171@pytest.mark.parametrize(172 "content,n_tokens",173 [174 ("I believe the meaning of life is", 9),175 ("This is a test", 6),176 ]177)178def test_embedding_usage_single(content, n_tokens):179 global server180 server.start()181 res = server.make_request("POST", "/v1/embeddings", data={"input": content})182 assert res.status_code == 200183 assert res.body['usage']['prompt_tokens'] == res.body['usage']['total_tokens']184 assert res.body['usage']['prompt_tokens'] == n_tokens185 186 187def test_embedding_usage_multiple():188 global server189 server.start()190 res = server.make_request("POST", "/v1/embeddings", data={191 "input": [192 "I believe the meaning of life is",193 "I believe the meaning of life is",194 ],195 })196 assert res.status_code == 200197 assert res.body['usage']['prompt_tokens'] == res.body['usage']['total_tokens']198 assert res.body['usage']['prompt_tokens'] == 2 * 9199 200 201def test_embedding_openai_library_base64():202 server.start()203 test_input = "Test base64 embedding output"204 205 # get embedding in default format206 res = server.make_request("POST", "/v1/embeddings", data={207 "input": test_input208 })209 assert res.status_code == 200210 vec0 = res.body["data"][0]["embedding"]211 212 # get embedding in base64 format213 res = server.make_request("POST", "/v1/embeddings", data={214 "input": test_input,215 "encoding_format": "base64"216 })217 218 assert res.status_code == 200219 assert "data" in res.body220 assert len(res.body["data"]) == 1221 222 embedding_data = res.body["data"][0]223 assert "embedding" in embedding_data224 assert isinstance(embedding_data["embedding"], str)225 226 # Verify embedding is valid base64227 decoded = base64.b64decode(embedding_data["embedding"])228 # Verify decoded data can be converted back to float array229 float_count = len(decoded) // 4 # 4 bytes per float230 floats = struct.unpack(f'{float_count}f', decoded)231 assert len(floats) > 0232 assert all(isinstance(x, float) for x in floats)233 assert len(floats) == len(vec0)234 235 # make sure the decoded data is the same as the original236 for x, y in zip(floats, vec0):237 assert abs(x - y) < EPSILON238 