Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
test_embedding.py292 linesDownload Raw Back to unit
1import base642import struct3import pytest4from openai import OpenAI5from utils import *6 7server = ServerPreset.bert_bge_small()8 9EPSILON = 1e-310 11@pytest.fixture(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 52def test_embedding_multiple_with_fa():53    server = ServerPreset.bert_bge_small_with_fa()54    server.pooling = 'last'55    server.start()56    # one of these should trigger the FA branch (i.e. context size % 256 == 0)57    res = server.make_request("POST", "/v1/embeddings", data={58        "input": [59            "a "*253,60            "b "*254,61            "c "*255,62            "d "*256,63        ],64    })65    assert res.status_code == 20066    assert len(res.body['data']) == 467    for d in res.body['data']:68        assert 'embedding' in d69        assert len(d['embedding']) > 170 71 72@pytest.mark.parametrize(73    "input,is_multi_prompt",74    [75        # do not crash on empty input76        ("", False),77        # single prompt78        ("string", False),79        ([12, 34, 56], False),80        ([12, 34, "string", 56, 78], False),81        # multiple prompts82        (["string1", "string2"], True),83        (["string1", [12, 34, 56]], True),84        ([[12, 34, 56], [12, 34, 56]], True),85        ([[12, 34, 56], [12, "string", 34, 56]], True),86    ]87)88def test_embedding_mixed_input(input, is_multi_prompt: bool):89    global server90    server.start()91    res = server.make_request("POST", "/v1/embeddings", data={"input": input})92    assert res.status_code == 20093    data = res.body['data']94    if is_multi_prompt:95        assert len(data) == len(input)96        for d in data:97            assert 'embedding' in d98            assert len(d['embedding']) > 199    else:100        assert 'embedding' in data[0]101        assert len(data[0]['embedding']) > 1102 103 104def test_embedding_pooling_mean():105    global server106    server.pooling = 'mean'107    server.start()108    res = server.make_request("POST", "/v1/embeddings", data={109        "input": "I believe the meaning of life is",110    })111    assert res.status_code == 200112    assert len(res.body['data']) == 1113    assert 'embedding' in res.body['data'][0]114    assert len(res.body['data'][0]['embedding']) > 1115 116    # make sure embedding vector is normalized117    assert abs(sum([x ** 2 for x in res.body['data'][0]['embedding']]) - 1) < EPSILON118 119 120def test_embedding_pooling_mean_multiple():121    global server122    server.pooling = 'mean'123    server.start()124    res = server.make_request("POST", "/v1/embeddings", data={125        "input": [126            "I believe the meaning of life is",127            "Write a joke about AI",128            "This is a test",129        ],130    })131    assert res.status_code == 200132    assert len(res.body['data']) == 3133    for d in res.body['data']:134        assert 'embedding' in d135        assert len(d['embedding']) > 1136 137 138def test_embedding_pooling_none():139    global server140    server.pooling = 'none'141    server.start()142    res = server.make_request("POST", "/embeddings", data={143        "input": "hello hello hello",144    })145    assert res.status_code == 200146    assert 'embedding' in res.body[0]147    assert len(res.body[0]['embedding']) == 5 # 3 text tokens + 2 special148 149    # make sure embedding vector is not normalized150    for x in res.body[0]['embedding']:151        assert abs(sum([x ** 2 for x in x]) - 1) > EPSILON152 153 154def test_embedding_pooling_none_oai():155    global server156    server.pooling = 'none'157    server.start()158    res = server.make_request("POST", "/v1/embeddings", data={159        "input": "hello hello hello",160    })161 162    # /v1/embeddings does not support pooling type 'none'163    assert res.status_code == 400164    assert "error" in res.body165 166 167def test_embedding_openai_library_single():168    global server169    server.pooling = 'last'170    server.start()171    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")172    res = client.embeddings.create(model="text-embedding-3-small", input="I believe the meaning of life is")173    assert len(res.data) == 1174    assert len(res.data[0].embedding) > 1175 176 177def test_embedding_openai_library_multiple():178    global server179    server.pooling = 'last'180    server.start()181    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")182    res = client.embeddings.create(model="text-embedding-3-small", input=[183        "I believe the meaning of life is",184        "Write a joke about AI from a very long prompt which will not be truncated",185        "This is a test",186        "This is another test",187    ])188    assert len(res.data) == 4189    for d in res.data:190        assert len(d.embedding) > 1191 192 193def test_embedding_error_prompt_too_long():194    global server195    server.pooling = 'last'196    server.start()197    res = server.make_request("POST", "/v1/embeddings", data={198        "input": "This is a test " * 512,199    })200    assert res.status_code != 200201    assert "too large" in res.body["error"]["message"]202 203 204def test_same_prompt_give_same_result():205    server.pooling = 'last'206    server.start()207    res = server.make_request("POST", "/v1/embeddings", data={208        "input": [209            "I believe the meaning of life is",210            "I believe the meaning of life is",211            "I believe the meaning of life is",212            "I believe the meaning of life is",213            "I believe the meaning of life is",214        ],215    })216    assert res.status_code == 200217    assert len(res.body['data']) == 5218    for i in range(1, len(res.body['data'])):219        v0 = res.body['data'][0]['embedding']220        vi = res.body['data'][i]['embedding']221        for x, y in zip(v0, vi):222            assert abs(x - y) < EPSILON223 224 225@pytest.mark.parametrize(226    "content,n_tokens",227    [228        ("I believe the meaning of life is", 9),229        ("This is a test", 6),230    ]231)232def test_embedding_usage_single(content, n_tokens):233    global server234    server.start()235    res = server.make_request("POST", "/v1/embeddings", data={"input": content})236    assert res.status_code == 200237    assert res.body['usage']['prompt_tokens'] == res.body['usage']['total_tokens']238    assert res.body['usage']['prompt_tokens'] == n_tokens239 240 241def test_embedding_usage_multiple():242    global server243    server.start()244    res = server.make_request("POST", "/v1/embeddings", data={245        "input": [246            "I believe the meaning of life is",247            "I believe the meaning of life is",248        ],249    })250    assert res.status_code == 200251    assert res.body['usage']['prompt_tokens'] == res.body['usage']['total_tokens']252    assert res.body['usage']['prompt_tokens'] == 2 * 9253 254 255def test_embedding_openai_library_base64():256    server.start()257    test_input = "Test base64 embedding output"258 259    # get embedding in default format260    res = server.make_request("POST", "/v1/embeddings", data={261        "input": test_input262    })263    assert res.status_code == 200264    vec0 = res.body["data"][0]["embedding"]265 266    # get embedding in base64 format267    res = server.make_request("POST", "/v1/embeddings", data={268        "input": test_input,269        "encoding_format": "base64"270    })271 272    assert res.status_code == 200273    assert "data" in res.body274    assert len(res.body["data"]) == 1275 276    embedding_data = res.body["data"][0]277    assert "embedding" in embedding_data278    assert isinstance(embedding_data["embedding"], str)279 280    # Verify embedding is valid base64281    decoded = base64.b64decode(embedding_data["embedding"])282    # Verify decoded data can be converted back to float array283    float_count = len(decoded) // 4  # 4 bytes per float284    floats = struct.unpack(f'{float_count}f', decoded)285    assert len(floats) > 0286    assert all(isinstance(x, float) for x in floats)287    assert len(floats) == len(vec0)288 289    # make sure the decoded data is the same as the original290    for x, y in zip(floats, vec0):291        assert abs(x - y) < EPSILON292