Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
test_embedding.py238 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(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