Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
test_completion.py669 linesDownload Raw Back to unit
1import pytest2import requests3import time4import random5 6from openai import OpenAI7from utils import *8 9server = ServerPreset.tinyllama2()10 11JSON_MULTIMODAL_KEY = "multimodal_data"12JSON_PROMPT_STRING_KEY = "prompt_string"13 14@pytest.fixture(autouse=True)15def create_server():16    global server17    server = ServerPreset.tinyllama2()18 19@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated,return_tokens", [20    ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False, False),21    ("Write a joke about AI from a very long prompt which will not be truncated", 64, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False, True),22])23def test_completion(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool, return_tokens: bool):24    global server25    server.start()26    res = server.make_request("POST", "/completion", data={27        "n_predict": n_predict,28        "prompt": prompt,29        "return_tokens": return_tokens,30    })31    assert res.status_code == 20032    assert res.body["timings"]["prompt_n"] == n_prompt33    assert res.body["timings"]["predicted_n"] == n_predicted34    assert res.body["truncated"] == truncated35    assert type(res.body["has_new_line"]) == bool36    assert match_regex(re_content, res.body["content"])37    if return_tokens:38        assert len(res.body["tokens"]) > 039        assert all(type(tok) == int for tok in res.body["tokens"])40    else:41        assert res.body["tokens"] == []42 43 44@pytest.mark.parametrize("prompt,n_predict,re_content,n_prompt,n_predicted,truncated", [45    ("I believe the meaning of life is", 8, "(going|bed)+", 18, 8, False),46    ("Write a joke about AI from a very long prompt which will not be truncated", 64, "(princesses|everyone|kids|Anna|forest)+", 46, 64, False),47])48def test_completion_stream(prompt: str, n_predict: int, re_content: str, n_prompt: int, n_predicted: int, truncated: bool):49    global server50    server.start()51    res = server.make_stream_request("POST", "/completion", data={52        "n_predict": n_predict,53        "prompt": prompt,54        "stream": True,55    })56    content = ""57    for data in res:58        assert "stop" in data and type(data["stop"]) == bool59        if data["stop"]:60            assert data["timings"]["prompt_n"] == n_prompt61            assert data["timings"]["predicted_n"] == n_predicted62            assert data["truncated"] == truncated63            assert data["stop_type"] == "limit"64            assert type(data["has_new_line"]) == bool65            assert "generation_settings" in data66            assert server.n_predict is not None67            assert data["generation_settings"]["n_predict"] == min(n_predict, server.n_predict)68            assert data["generation_settings"]["seed"] == server.seed69            assert "adaptive_target" in data["generation_settings"]70            assert "adaptive_decay" in data["generation_settings"]71            assert match_regex(re_content, content)72        else:73            assert len(data["tokens"]) > 074            assert all(type(tok) == int for tok in data["tokens"])75            content += data["content"]76 77 78def test_completion_stream_vs_non_stream():79    global server80    server.start()81    res_stream = server.make_stream_request("POST", "/completion", data={82        "n_predict": 8,83        "prompt": "I believe the meaning of life is",84        "stream": True,85    })86    res_non_stream = server.make_request("POST", "/completion", data={87        "n_predict": 8,88        "prompt": "I believe the meaning of life is",89    })90    content_stream = ""91    for data in res_stream:92        content_stream += data["content"]93    assert content_stream == res_non_stream.body["content"]94 95 96def test_completion_with_openai_library():97    global server98    server.start()99    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")100    res = client.completions.create(101        model="davinci-002",102        prompt="I believe the meaning of life is",103        max_tokens=8,104    )105    assert res.system_fingerprint is not None and res.system_fingerprint.startswith("b")106    assert res.choices[0].finish_reason == "length"107    assert res.choices[0].text is not None108    assert match_regex("(going|bed)+", res.choices[0].text)109 110 111def test_completion_stream_with_openai_library():112    global server113    server.start()114    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")115    res = client.completions.create(116        model="davinci-002",117        prompt="I believe the meaning of life is",118        max_tokens=8,119        stream=True,120    )121    output_text = ''122    for data in res:123        choice = data.choices[0]124        if choice.finish_reason is None:125            assert choice.text is not None126            output_text += choice.text127    assert match_regex("(going|bed)+", output_text)128 129 130# Test case from https://github.com/ggml-org/llama.cpp/issues/13780131@pytest.mark.slow132def test_completion_stream_with_openai_library_stops():133    global server134    server.model_hf_repo = "bartowski/Phi-3.5-mini-instruct-GGUF:Q4_K_M"135    server.model_hf_file = None136    server.start()137    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")138    res = client.completions.create(139        model="davinci-002",140        prompt="System: You are helpful assistant.\nAssistant:\nHey! How could I help?\nUser:\nTell me a joke.\nAssistant:\n",141        stop=["User:\n", "Assistant:\n"],142        max_tokens=200,143        stream=True,144    )145    output_text = ''146    for data in res:147        choice = data.choices[0]148        if choice.finish_reason is None:149            assert choice.text is not None150            output_text += choice.text151    assert match_regex("Sure, here's one for[\\s\\S]*", output_text), f'Unexpected output: {output_text}'152 153 154@pytest.mark.parametrize("n_slots", [1, 2])155def test_consistent_result_same_seed(n_slots: int):156    global server157    server.n_slots = n_slots158    server.start()159    last_res = None160    for _ in range(4):161        res = server.make_request("POST", "/completion", data={162            "prompt": "I believe the meaning of life is",163            "seed": 42,164            "temperature": 0.0,165            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed166        })167        if last_res is not None:168            assert res.body["content"] == last_res.body["content"]169        last_res = res170 171 172@pytest.mark.parametrize("n_slots", [1, 2])173def test_different_result_different_seed(n_slots: int):174    global server175    server.n_slots = n_slots176    server.start()177    last_res = None178    for seed in range(4):179        res = server.make_request("POST", "/completion", data={180            "prompt": "I believe the meaning of life is",181            "seed": seed,182            "temperature": 1.0,183            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed184        })185        if last_res is not None:186            assert res.body["content"] != last_res.body["content"]187        last_res = res188 189# TODO figure why it don't work with temperature = 1190# @pytest.mark.parametrize("temperature", [0.0, 1.0])191@pytest.mark.parametrize("n_batch", [16, 32])192@pytest.mark.parametrize("temperature", [0.0])193def test_consistent_result_different_batch_size(n_batch: int, temperature: float):194    global server195    server.n_batch = n_batch196    server.start()197    last_res = None198    for _ in range(4):199        res = server.make_request("POST", "/completion", data={200            "prompt": "I believe the meaning of life is",201            "seed": 42,202            "temperature": temperature,203            "cache_prompt": False,  # TODO: remove this once test_cache_vs_nocache_prompt is fixed204        })205        if last_res is not None:206            assert res.body["content"] == last_res.body["content"]207        last_res = res208 209 210@pytest.mark.skip(reason="This test fails on linux, need to be fixed")211def test_cache_vs_nocache_prompt():212    global server213    server.start()214    res_cache = server.make_request("POST", "/completion", data={215        "prompt": "I believe the meaning of life is",216        "seed": 42,217        "temperature": 1.0,218        "cache_prompt": True,219    })220    res_no_cache = server.make_request("POST", "/completion", data={221        "prompt": "I believe the meaning of life is",222        "seed": 42,223        "temperature": 1.0,224        "cache_prompt": False,225    })226    assert res_cache.body["content"] == res_no_cache.body["content"]227 228 229def test_nocache_long_input_prompt():230    global server231    server.start()232    res = server.make_request("POST", "/completion", data={233        "prompt": "I believe the meaning of life is"*32,234        "seed": 42,235        "temperature": 1.0,236        "cache_prompt": False,237    })238    assert res.status_code == 400239 240def test_json_prompt_no_mtmd():241    global server242    server.start()243    res = server.make_request("POST", "/completion", data={244        "prompt": { JSON_PROMPT_STRING_KEY: "I believe the meaning of life is" },245        "seed": 42,246        "temperature": 1.0,247        "cache_prompt": False,248    })249    assert res.status_code == 200250 251def test_json_prompt_mtm_error_when_not_supported():252    global server253    server.start()254    res = server.make_request("POST", "/completion", data={255        "prompt": { JSON_PROMPT_STRING_KEY: "I believe the meaning of life is <__media__>", JSON_MULTIMODAL_KEY: "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=" },256        "seed": 42,257        "temperature": 1.0,258        "cache_prompt": False,259    })260    # MTMD is disabled on this model, so this should fail.261    assert res.status_code != 200262 263def test_completion_with_tokens_input():264    global server265    server.temperature = 0.0266    server.start()267    prompt_str = "I believe the meaning of life is"268    res = server.make_request("POST", "/tokenize", data={269        "content": prompt_str,270        "add_special": True,271    })272    assert res.status_code == 200273    tokens = res.body["tokens"]274 275    # single completion276    res = server.make_request("POST", "/completion", data={277        "prompt": tokens,278    })279    assert res.status_code == 200280    assert type(res.body["content"]) == str281 282    # batch completion283    res = server.make_request("POST", "/completion", data={284        "prompt": [tokens, tokens],285    })286    assert res.status_code == 200287    assert type(res.body) == list288    assert len(res.body) == 2289    assert res.body[0]["content"] == res.body[1]["content"]290 291    # mixed string and tokens292    res = server.make_request("POST", "/completion", data={293        "prompt": [tokens, prompt_str],294    })295    assert res.status_code == 200296    assert type(res.body) == list297    assert len(res.body) == 2298    assert res.body[0]["content"] == res.body[1]["content"]299 300    # mixed JSON and tokens301    res = server.make_request("POST", "/completion", data={302        "prompt": [303            tokens,304            {305                JSON_PROMPT_STRING_KEY: "I believe the meaning of life is",306            },307        ],308    })309    assert res.status_code == 200310    assert type(res.body) == list311    assert len(res.body) == 2312    assert res.body[0]["content"] == res.body[1]["content"]313 314    # mixed string and tokens in one sequence315    res = server.make_request("POST", "/completion", data={316        "prompt": [1, 2, 3, 4, 5, 6, prompt_str, 7, 8, 9, 10, prompt_str],317    })318    assert res.status_code == 200319    assert type(res.body["content"]) == str320 321 322@pytest.mark.parametrize("n_slots,n_requests", [323    (1, 3),324    (2, 2),325    (2, 4),326    (4, 2), # some slots must be idle327    (4, 6),328])329def test_completion_parallel_slots(n_slots: int, n_requests: int):330    global server331    server.n_slots = n_slots332    server.temperature = 0.0333    server.start()334 335    PROMPTS = [336        ("Write a very long book.", "(very|special|big)+"),337        ("Write another a poem.", "(small|house)+"),338        ("What is LLM?", "(Dad|said)+"),339        ("The sky is blue and I love it.", "(climb|leaf)+"),340        ("Write another very long music lyrics.", "(friends|step|sky)+"),341        ("Write a very long joke.", "(cat|Whiskers)+"),342    ]343    def check_slots_status():344        should_all_slots_busy = n_requests >= n_slots345        time.sleep(0.1)346        res = server.make_request("GET", "/slots")347        n_busy = sum([1 for slot in res.body if slot["is_processing"]])348        if should_all_slots_busy:349            assert n_busy == n_slots350        else:351            assert n_busy <= n_slots352 353    tasks = []354    for i in range(n_requests):355        prompt, re_content = PROMPTS[i % len(PROMPTS)]356        tasks.append((server.make_request, ("POST", "/completion", {357            "prompt": prompt,358            "seed": 42,359            "temperature": 1.0,360        })))361    tasks.append((check_slots_status, ()))362    results = parallel_function_calls(tasks)363 364    # check results365    for i in range(n_requests):366        prompt, re_content = PROMPTS[i % len(PROMPTS)]367        res = results[i]368        assert res.status_code == 200369        assert type(res.body["content"]) == str370        assert len(res.body["content"]) > 10371        # FIXME: the result is not deterministic when using other slot than slot 0372        # assert match_regex(re_content, res.body["content"])373 374 375@pytest.mark.parametrize(376    "n_ctx,n_slots,n_predict_vals,expected_success",377    [378        (256, 4, [80, 40, 80, 80], [True,  True,  True,  True]),379        (256, 4, [70, 70, 70, 70], [False, False, False, False]),380        (256, 4, [90, 90, 40, 90], [False, False, True,  False]),381        (256, 4, [90, 90, 40, 75], [True,  True,  True,  True]),382    ],383)384def test_completion_unified(n_ctx, n_slots, n_predict_vals, expected_success):385    global server386    server.n_slots = n_slots387    server.kv_unified = True388    server.n_ctx = n_ctx389    server.start()390    prompt = "A"391    tasks = []392    for n_predict in n_predict_vals:393        tasks.append((server.make_request, ("POST", "/completion", {"prompt": prompt, "n_predict": n_predict})))394    results = parallel_function_calls(tasks)395    for res, n_predict, expect_ok in zip(results, n_predict_vals, expected_success):396        if expect_ok:397            # the pool is aborted as a whole, so a request that fits on its own398            # is still dropped when the slots overlap, and it says so explicitly399            assert res.status_code == 200 or (400                res.status_code == 500401                and "context size has been exceeded" in res.body["error"]["message"].lower()402            )403 404        # note: https://github.com/ggml-org/llama.cpp/pull/18700#issuecomment-3728695581405        if res.status_code == 200:406            assert "content" in res.body407            if "timings" in res.body:408                assert res.body["timings"]["predicted_n"] == n_predict409 410 411@pytest.mark.parametrize(412    "prompt,n_predict,response_fields",413    [414        ("I believe the meaning of life is", 8, []),415        ("I believe the meaning of life is", 32, ["content", "generation_settings/n_predict", "prompt"]),416    ],417)418def test_completion_response_fields(419    prompt: str, n_predict: int, response_fields: list[str]420):421    global server422    server.start()423    res = server.make_request(424        "POST",425        "/completion",426        data={427            "n_predict": n_predict,428            "prompt": prompt,429            "response_fields": response_fields,430        },431    )432    assert res.status_code == 200433    assert "content" in res.body434    assert len(res.body["content"])435    if len(response_fields):436        assert res.body["generation_settings/n_predict"] == n_predict437        assert res.body["prompt"] == "<s> " + prompt438        assert isinstance(res.body["content"], str)439        assert len(res.body) == len(response_fields)440    else:441        assert len(res.body)442        assert "generation_settings" in res.body443 444 445def test_n_probs():446    global server447    server.start()448    res = server.make_request("POST", "/completion", data={449        "prompt": "I believe the meaning of life is",450        "n_probs": 10,451        "temperature": 0.0,452        "n_predict": 5,453    })454    assert res.status_code == 200455    assert "completion_probabilities" in res.body456    assert len(res.body["completion_probabilities"]) == 5457    for tok in res.body["completion_probabilities"]:458        assert "id" in tok and tok["id"] > 0459        assert "token" in tok and type(tok["token"]) == str460        assert "logprob" in tok and tok["logprob"] <= 0.0461        assert "bytes" in tok and type(tok["bytes"]) == list462        assert len(tok["top_logprobs"]) == 10463        for prob in tok["top_logprobs"]:464            assert "id" in prob and prob["id"] > 0465            assert "token" in prob and type(prob["token"]) == str466            assert "logprob" in prob and prob["logprob"] <= 0.0467            assert "bytes" in prob and type(prob["bytes"]) == list468 469 470def test_n_probs_stream():471    global server472    server.start()473    res = server.make_stream_request("POST", "/completion", data={474        "prompt": "I believe the meaning of life is",475        "n_probs": 10,476        "temperature": 0.0,477        "n_predict": 5,478        "stream": True,479    })480    for data in res:481        if data["stop"] == False:482            assert "completion_probabilities" in data483            assert len(data["completion_probabilities"]) == 1484            for tok in data["completion_probabilities"]:485                assert "id" in tok and tok["id"] > 0486                assert "token" in tok and type(tok["token"]) == str487                assert "logprob" in tok and tok["logprob"] <= 0.0488                assert "bytes" in tok and type(tok["bytes"]) == list489                assert len(tok["top_logprobs"]) == 10490                for prob in tok["top_logprobs"]:491                    assert "id" in prob and prob["id"] > 0492                    assert "token" in prob and type(prob["token"]) == str493                    assert "logprob" in prob and prob["logprob"] <= 0.0494                    assert "bytes" in prob and type(prob["bytes"]) == list495 496 497def test_n_probs_post_sampling():498    global server499    server.start()500    res = server.make_request("POST", "/completion", data={501        "prompt": "Today was the day. Today I would finally become a",502        "n_probs": 10,503        "temperature": 1.0,504        "n_predict": 5,505        "post_sampling_probs": True,506    })507    assert res.status_code == 200508    assert "completion_probabilities" in res.body509    assert len(res.body["completion_probabilities"]) == 5510    for (i, tok) in enumerate(res.body["completion_probabilities"]):511        assert "id" in tok and tok["id"] > 0512        assert "token" in tok and type(tok["token"]) == str513        assert "prob" in tok and 0.0 < tok["prob"] <= 1.0514        assert "bytes" in tok and type(tok["bytes"]) == list515        assert "top_probs" in tok and type(tok["top_probs"]) == list516 517        for prob in tok["top_probs"]:518            assert "id" in prob and prob["id"] > 0519            assert "token" in prob and type(prob["token"]) == str520            # 0.0 probability tokens should never be returned by the server521            assert "prob" in prob and 0.0 < prob["prob"] <= 1.0522            assert "bytes" in prob and type(prob["bytes"]) == list523 524        if i == 0:525            # The prompt is vague enough that we should get at least 10 possibilities526            # for the first token.527            assert len(tok["top_probs"]) == 10528 529        if len(tok["top_probs"]) < 10:530            # Getting less than the requested number of probabilities should only happen531            # if the ones we did get already sum to 1.0.532            assert sum(p["prob"] for p in tok["top_probs"]) == pytest.approx(1.0)533 534def test_n_probs_post_backend_sampling():535    """Verify that the same probabilities are returned with and without backend sampling."""536    global server537    server.backend_sampling = True538    server.start()539 540    def make_request(backend_sampling):541        n_predict = 20542 543        res = server.make_request("POST", "/completion", data={544            "prompt": "The countries of Europe, in random order, are:",545            "n_probs": 10,546            "n_predict": n_predict,547            "post_sampling_probs": True,548            "seed": 4242,549            "backend_sampling": backend_sampling,550        })551        assert res.status_code == 200552 553        total_probs = 0554        completions = res.body["completion_probabilities"]555        assert len(completions) == n_predict556        for tok in completions:557            # Handling of 0.0 probabilities differs between samplers and backend sampling. Filter them to normalize the558            # data.559            tok["top_probs"] = [x for x in tok["top_probs"] if x["prob"] > 0.0]560            total_probs += len(tok["top_probs"])561        # Verify that we got at least two top probs on average, to ensure the effectiveness of the test.562        assert total_probs >= 2 * n_predict563        return completions564 565    def verify_token(a, b):566        assert a["id"] == b["id"]567        assert a["token"] == b["token"]568        assert a["bytes"] == b["bytes"]569        assert a["prob"] == pytest.approx(b["prob"], abs=0.01)570 571    for (a, b) in zip(make_request(True), make_request(False)):572        verify_token(a, b)573        assert len(a["top_probs"]) == len(b["top_probs"])574 575        for (aa, bb) in zip(a["top_probs"], b["top_probs"]):576            verify_token(aa, bb)577 578@pytest.mark.parametrize("tokenize,openai_style", [(False, False), (False, True), (True, False), (True, True)])579def test_logit_bias(tokenize, openai_style):580    global server581    server.start()582 583    exclude = ["i", "I", "the", "The", "to", "a", "an", "be", "is", "was", "but", "But", "and", "And", "so", "So", "you", "You", "he", "He", "she", "She", "we", "We", "they", "They", "it", "It", "his", "His", "her", "Her", "book", "Book"]584 585    logit_bias = []586    if tokenize:587        res = server.make_request("POST", "/tokenize", data={588            "content": " " + " ".join(exclude) + " ",589        })590        assert res.status_code == 200591        tokens = res.body["tokens"]592        logit_bias = [[tok, -100] for tok in tokens]593 594    else:595        logit_bias = [[" " + tok + " ", -100] for tok in exclude]596 597    if openai_style:598        logit_bias = {el[0]: -100 for el in logit_bias}599 600    res = server.make_request("POST", "/completion", data={601        "n_predict": 64,602        "prompt": "What is the best book",603        "logit_bias": logit_bias,604        "temperature": 0.0605    })606    assert res.status_code == 200607    output_text = res.body["content"]608    assert all(output_text.find(" " + tok + " ") == -1 for tok in exclude)609 610 611def test_cancel_request():612    global server613    server.n_ctx = 4096614    server.n_predict = -1615    server.n_slots = 1616    server.server_slots = True617    server.start()618    # send a request that will take a long time, but cancel it before it finishes619    try:620        server.make_request("POST", "/completion", data={621            "prompt": "I believe the meaning of life is",622        }, timeout=0.1)623    except requests.exceptions.ReadTimeout:624        pass # expected625    # make sure the slot is free626    time.sleep(2)627    res = server.make_request("GET", "/slots")628    assert res.body[0]["is_processing"] == False629 630 631# this test exercises the host-memory prompt cache632# ref: https://github.com/ggml-org/llama.cpp/pull/16391633# ref: https://github.com/ggml-org/llama.cpp/pull/17078634def test_completion_prompt_cache():635    global server636    server.n_slots = 2637    server.kv_unified = True638    server.start()639 640    for _ in range(16):641        # generate alternating random prompts with variable lengths in order to get them in and out of the cache642        r = random.randint(0, 4)643        prompt = (" Hello " +  str(r)) * (40 + r)644        n_prompt = (40 + r)*5 + 2645        n_predict = random.randint(1, 8)646 647        res = server.make_request(648            "POST",649            "/completion",650            data={651                "prompt": prompt,652                "n_predict": n_predict,653            },654        )655 656        assert res.status_code == 200657        assert "content" in res.body658        content = res.body["content"]659        assert isinstance(content, str)660        assert len(content) > 0661 662        assert type(res.body["has_new_line"]) == bool663        assert "timings" in res.body664        timings = res.body["timings"]665 666        assert "prompt_n" in timings and timings["prompt_n"] + timings["cache_n"] == n_prompt667        assert "predicted_n" in timings and timings["predicted_n"] == n_predict668        assert "tokens" in res.body and isinstance(res.body["tokens"], list)669