Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
test_completion.py429 linesDownload Raw Back to unit
1import pytest2import requests3import time4from openai import OpenAI5from utils import *6 7server = ServerPreset.tinyllama2()8 9 10@pytest.fixture(scope="module", autouse=True)11def create_server():12    global server13    server = ServerPreset.tinyllama2()14 15@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated,return_tokens", [16    ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False, False),17    ("Write a joke about AI from a very long prompt which will not be truncated", 256, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False, True),18])19def test_completion(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool, return_tokens: bool):20    global server21    server.start()22    res = server.make_request("POST", "/completion", data={23        "n_predict": n_predict,24        "prompt": prompt,25        "return_tokens": return_tokens,26    })27    assert res.status_code == 20028    assert res.body["timings"]["prompt_n"] == n_prompt29    assert res.body["timings"]["predicted_n"] == n_predicted30    assert res.body["truncated"] == truncated31    assert type(res.body["has_new_line"]) == bool32    assert match_regex(re_content, res.body["content"])33    if return_tokens:34        assert len(res.body["tokens"]) > 035        assert all(type(tok) == int for tok in res.body["tokens"])36    else:37        assert res.body["tokens"] == []38 39 40@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated", [41    ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False),42    ("Write a joke about AI from a very long prompt which will not be truncated", 256, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False),43])44def test_completion_stream(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool):45    global server46    server.start()47    res = server.make_stream_request("POST", "/completion", data={48        "n_predict": n_predict,49        "prompt": prompt,50        "stream": True,51    })52    content = ""53    for data in res:54        assert "stop" in data and type(data["stop"]) == bool55        if data["stop"]:56            assert data["timings"]["prompt_n"] == n_prompt57            assert data["timings"]["predicted_n"] == n_predicted58            assert data["truncated"] == truncated59            assert data["stop_type"] == "limit"60            assert type(data["has_new_line"]) == bool61            assert "generation_settings" in data62            assert server.n_predict is not None63            assert data["generation_settings"]["n_predict"] == min(n_predict, server.n_predict)64            assert data["generation_settings"]["seed"] == server.seed65            assert match_regex(re_content, content)66        else:67            assert len(data["tokens"]) > 068            assert all(type(tok) == int for tok in data["tokens"])69            content += data["content"]70 71 72def test_completion_stream_vs_non_stream():73    global server74    server.start()75    res_stream = server.make_stream_request("POST", "/completion", data={76        "n_predict": 8,77        "prompt": "I believe the meaning of life is",78        "stream": True,79    })80    res_non_stream = server.make_request("POST", "/completion", data={81        "n_predict": 8,82        "prompt": "I believe the meaning of life is",83    })84    content_stream = ""85    for data in res_stream:86        content_stream += data["content"]87    assert content_stream == res_non_stream.body["content"]88 89 90def test_completion_with_openai_library():91    global server92    server.start()93    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")94    res = client.completions.create(95        model="davinci-002",96        prompt="I believe the meaning of life is",97        max_tokens=8,98    )99    assert res.system_fingerprint is not None and res.system_fingerprint.startswith("b")100    assert res.choices[0].finish_reason == "length"101    assert res.choices[0].text is not None102    assert match_regex("(going|bed)+", res.choices[0].text)103 104 105def test_completion_stream_with_openai_library():106    global server107    server.start()108    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")109    res = client.completions.create(110        model="davinci-002",111        prompt="I believe the meaning of life is",112        max_tokens=8,113        stream=True,114    )115    output_text = ''116    for data in res:117        choice = data.choices[0]118        if choice.finish_reason is None:119            assert choice.text is not None120            output_text += choice.text121    assert match_regex("(going|bed)+", output_text)122 123 124@pytest.mark.parametrize("n_slots", [1, 2])125def test_consistent_result_same_seed(n_slots: int):126    global server127    server.n_slots = n_slots128    server.start()129    last_res = None130    for _ in range(4):131        res = server.make_request("POST", "/completion", data={132            "prompt": "I believe the meaning of life is",133            "seed": 42,134            "temperature": 0.0,135            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed136        })137        if last_res is not None:138            assert res.body["content"] == last_res.body["content"]139        last_res = res140 141 142@pytest.mark.parametrize("n_slots", [1, 2])143def test_different_result_different_seed(n_slots: int):144    global server145    server.n_slots = n_slots146    server.start()147    last_res = None148    for seed in range(4):149        res = server.make_request("POST", "/completion", data={150            "prompt": "I believe the meaning of life is",151            "seed": seed,152            "temperature": 1.0,153            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed154        })155        if last_res is not None:156            assert res.body["content"] != last_res.body["content"]157        last_res = res158 159# TODO figure why it don't work with temperature = 1160# @pytest.mark.parametrize("temperature", [0.0, 1.0])161@pytest.mark.parametrize("n_batch", [16, 32])162@pytest.mark.parametrize("temperature", [0.0])163def test_consistent_result_different_batch_size(n_batch: int, temperature: float):164    global server165    server.n_batch = n_batch166    server.start()167    last_res = None168    for _ in range(4):169        res = server.make_request("POST", "/completion", data={170            "prompt": "I believe the meaning of life is",171            "seed": 42,172            "temperature": temperature,173            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed174        })175        if last_res is not None:176            assert res.body["content"] == last_res.body["content"]177        last_res = res178 179 180@pytest.mark.skip(reason="This test fails on linux, need to be fixed")181def test_cache_vs_nocache_prompt():182    global server183    server.start()184    res_cache = server.make_request("POST", "/completion", data={185        "prompt": "I believe the meaning of life is",186        "seed": 42,187        "temperature": 1.0,188        "cache_prompt": True,189    })190    res_no_cache = server.make_request("POST", "/completion", data={191        "prompt": "I believe the meaning of life is",192        "seed": 42,193        "temperature": 1.0,194        "cache_prompt": False,195    })196    assert res_cache.body["content"] == res_no_cache.body["content"]197 198 199def test_completion_with_tokens_input():200    global server201    server.temperature = 0.0202    server.start()203    prompt_str = "I believe the meaning of life is"204    res = server.make_request("POST", "/tokenize", data={205        "content": prompt_str,206        "add_special": True,207    })208    assert res.status_code == 200209    tokens = res.body["tokens"]210 211    # single completion212    res = server.make_request("POST", "/completion", data={213        "prompt": tokens,214    })215    assert res.status_code == 200216    assert type(res.body["content"]) == str217 218    # batch completion219    res = server.make_request("POST", "/completion", data={220        "prompt": [tokens, tokens],221    })222    assert res.status_code == 200223    assert type(res.body) == list224    assert len(res.body) == 2225    assert res.body[0]["content"] == res.body[1]["content"]226 227    # mixed string and tokens228    res = server.make_request("POST", "/completion", data={229        "prompt": [tokens, prompt_str],230    })231    assert res.status_code == 200232    assert type(res.body) == list233    assert len(res.body) == 2234    assert res.body[0]["content"] == res.body[1]["content"]235 236    # mixed string and tokens in one sequence237    res = server.make_request("POST", "/completion", data={238        "prompt": [1, 2, 3, 4, 5, 6, prompt_str, 7, 8, 9, 10, prompt_str],239    })240    assert res.status_code == 200241    assert type(res.body["content"]) == str242 243 244@pytest.mark.parametrize("n_slots,n_requests", [245    (1, 3),246    (2, 2),247    (2, 4),248    (4, 2), # some slots must be idle249    (4, 6),250])251def test_completion_parallel_slots(n_slots: int, n_requests: int):252    global server253    server.n_slots = n_slots254    server.temperature = 0.0255    server.start()256 257    PROMPTS = [258        ("Write a very long book.", "(very|special|big)+"),259        ("Write another a poem.", "(small|house)+"),260        ("What is LLM?", "(Dad|said)+"),261        ("The sky is blue and I love it.", "(climb|leaf)+"),262        ("Write another very long music lyrics.", "(friends|step|sky)+"),263        ("Write a very long joke.", "(cat|Whiskers)+"),264    ]265    def check_slots_status():266        should_all_slots_busy = n_requests >= n_slots267        time.sleep(0.1)268        res = server.make_request("GET", "/slots")269        n_busy = sum([1 for slot in res.body if slot["is_processing"]])270        if should_all_slots_busy:271            assert n_busy == n_slots272        else:273            assert n_busy <= n_slots274 275    tasks = []276    for i in range(n_requests):277        prompt, re_content = PROMPTS[i % len(PROMPTS)]278        tasks.append((server.make_request, ("POST", "/completion", {279            "prompt": prompt,280            "seed": 42,281            "temperature": 1.0,282        })))283    tasks.append((check_slots_status, ()))284    results = parallel_function_calls(tasks)285 286    # check results287    for i in range(n_requests):288        prompt, re_content = PROMPTS[i % len(PROMPTS)]289        res = results[i]290        assert res.status_code == 200291        assert type(res.body["content"]) == str292        assert len(res.body["content"]) > 10293        # FIXME: the result is not deterministic when using other slot than slot 0294        # assert match_regex(re_content, res.body["content"])295 296 297@pytest.mark.parametrize(298    "prompt,n_predict,response_fields",299    [300        ("I believe the meaning of life is", 8, []),301        ("I believe the meaning of life is", 32, ["content", "generation_settings/n_predict", "prompt"]),302    ],303)304def test_completion_response_fields(305    prompt: str, n_predict: int, response_fields: list[str]306):307    global server308    server.start()309    res = server.make_request(310        "POST",311        "/completion",312        data={313            "n_predict": n_predict,314            "prompt": prompt,315            "response_fields": response_fields,316        },317    )318    assert res.status_code == 200319    assert "content" in res.body320    assert len(res.body["content"])321    if len(response_fields):322        assert res.body["generation_settings/n_predict"] == n_predict323        assert res.body["prompt"] == "<s> " + prompt324        assert isinstance(res.body["content"], str)325        assert len(res.body) == len(response_fields)326    else:327        assert len(res.body)328        assert "generation_settings" in res.body329 330 331def test_n_probs():332    global server333    server.start()334    res = server.make_request("POST", "/completion", data={335        "prompt": "I believe the meaning of life is",336        "n_probs": 10,337        "temperature": 0.0,338        "n_predict": 5,339    })340    assert res.status_code == 200341    assert "completion_probabilities" in res.body342    assert len(res.body["completion_probabilities"]) == 5343    for tok in res.body["completion_probabilities"]:344        assert "id" in tok and tok["id"] > 0345        assert "token" in tok and type(tok["token"]) == str346        assert "logprob" in tok and tok["logprob"] <= 0.0347        assert "bytes" in tok and type(tok["bytes"]) == list348        assert len(tok["top_logprobs"]) == 10349        for prob in tok["top_logprobs"]:350            assert "id" in prob and prob["id"] > 0351            assert "token" in prob and type(prob["token"]) == str352            assert "logprob" in prob and prob["logprob"] <= 0.0353            assert "bytes" in prob and type(prob["bytes"]) == list354 355 356def test_n_probs_stream():357    global server358    server.start()359    res = server.make_stream_request("POST", "/completion", data={360        "prompt": "I believe the meaning of life is",361        "n_probs": 10,362        "temperature": 0.0,363        "n_predict": 5,364        "stream": True,365    })366    for data in res:367        if data["stop"] == False:368            assert "completion_probabilities" in data369            assert len(data["completion_probabilities"]) == 1370            for tok in data["completion_probabilities"]:371                assert "id" in tok and tok["id"] > 0372                assert "token" in tok and type(tok["token"]) == str373                assert "logprob" in tok and tok["logprob"] <= 0.0374                assert "bytes" in tok and type(tok["bytes"]) == list375                assert len(tok["top_logprobs"]) == 10376                for prob in tok["top_logprobs"]:377                    assert "id" in prob and prob["id"] > 0378                    assert "token" in prob and type(prob["token"]) == str379                    assert "logprob" in prob and prob["logprob"] <= 0.0380                    assert "bytes" in prob and type(prob["bytes"]) == list381 382 383def test_n_probs_post_sampling():384    global server385    server.start()386    res = server.make_request("POST", "/completion", data={387        "prompt": "I believe the meaning of life is",388        "n_probs": 10,389        "temperature": 0.0,390        "n_predict": 5,391        "post_sampling_probs": True,392    })393    assert res.status_code == 200394    assert "completion_probabilities" in res.body395    assert len(res.body["completion_probabilities"]) == 5396    for tok in res.body["completion_probabilities"]:397        assert "id" in tok and tok["id"] > 0398        assert "token" in tok and type(tok["token"]) == str399        assert "prob" in tok and 0.0 < tok["prob"] <= 1.0400        assert "bytes" in tok and type(tok["bytes"]) == list401        assert len(tok["top_probs"]) == 10402        for prob in tok["top_probs"]:403            assert "id" in prob and prob["id"] > 0404            assert "token" in prob and type(prob["token"]) == str405            assert "prob" in prob and 0.0 <= prob["prob"] <= 1.0406            assert "bytes" in prob and type(prob["bytes"]) == list407        # because the test model usually output token with either 100% or 0% probability, we need to check all the top_probs408        assert any(prob["prob"] == 1.0 for prob in tok["top_probs"])409 410 411def test_cancel_request():412    global server413    server.n_ctx = 4096414    server.n_predict = -1415    server.n_slots = 1416    server.server_slots = True417    server.start()418    # send a request that will take a long time, but cancel it before it finishes419    try:420        server.make_request("POST", "/completion", data={421            "prompt": "I believe the meaning of life is",422        }, timeout=0.1)423    except requests.exceptions.ReadTimeout:424        pass # expected425    # make sure the slot is free426    time.sleep(1) # wait for HTTP_POLLING_SECONDS427    res = server.make_request("GET", "/slots")428    assert res.body[0]["is_processing"] == False429