Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
test_chat_completion.py535 linesDownload Raw Back to unit
1import pytest2from openai import OpenAI3from utils import *4 5server: ServerProcess6 7@pytest.fixture(autouse=True)8def create_server():9    global server10    server = ServerPreset.tinyllama2()11 12 13@pytest.mark.parametrize(14    "model,system_prompt,user_prompt,max_tokens,re_content,n_prompt,n_predicted,finish_reason,jinja,chat_template",15    [16        (None, "Book", "Hey", 8, "But she couldn't", 69, 8, "length", False, None),17        (None, "Book", "Hey", 8, "But she couldn't", 69, 8, "length", True, None),18        (None, "Book", "What is the best book", 8, "(Suddenly)+|\\{ \" Sarax.", 77, 8, "length", False, None),19        (None, "Book", "What is the best book", 8, "(Suddenly)+|\\{ \" Sarax.", 77, 8, "length", True,  None),20        (None, "Book", "What is the best book", 8, "(Suddenly)+|\\{ \" Sarax.", 77, 8, "length", True, 'chatml'),21        (None, "Book", "What is the best book", 8, "^ blue",                    23, 8, "length", True, "This is not a chat template, it is"),22        ("codellama70b", "You are a coding assistant.", "Write the fibonacci function in c++.", 128, "(Aside|she|felter|alonger)+", 104, 128, "length", False, None),23        ("codellama70b", "You are a coding assistant.", "Write the fibonacci function in c++.", 128, "(Aside|she|felter|alonger)+", 104, 128, "length", True, None),24        (None, "Book", [{"type": "text", "text": "What is"}, {"type": "text", "text": "the best book"}], 8, "Whillicter", 79, 8, "length", False, None),25        (None, "Book", [{"type": "text", "text": "What is"}, {"type": "text", "text": "the best book"}], 8, "Whillicter", 79, 8, "length", True, None),26    ]27)28def test_chat_completion(model, system_prompt, user_prompt, max_tokens, re_content, n_prompt, n_predicted, finish_reason, jinja, chat_template):29    global server30    server.jinja = jinja31    server.chat_template = chat_template32    server.start()33    res = server.make_request("POST", "/chat/completions", data={34        "model": model,35        "max_tokens": max_tokens,36        "messages": [37            {"role": "system", "content": system_prompt},38            {"role": "user", "content": user_prompt},39        ],40    })41    assert res.status_code == 20042    assert "cmpl" in res.body["id"] # make sure the completion id has the expected format43    assert res.body["system_fingerprint"].startswith("b")44    # we no longer reflect back the model name, see https://github.com/ggml-org/llama.cpp/pull/1766845    # assert res.body["model"] == model if model is not None else server.model_alias46    assert res.body["usage"]["prompt_tokens"] == n_prompt47    assert res.body["usage"]["completion_tokens"] == n_predicted48    choice = res.body["choices"][0]49    assert "assistant" == choice["message"]["role"]50    assert match_regex(re_content, choice["message"]["content"]), f'Expected {re_content}, got {choice["message"]["content"]}'51    assert choice["finish_reason"] == finish_reason52 53 54def test_chat_completion_cached_tokens():55    global server56    server.n_slots = 157    server.start()58    seq = [59        ("1 2 3 4 5 6", 77, 0),60        ("1 2 3 4 5 6", 77, 76),61        ("1 2 3 4 5 9", 77, 51),62        ("1 2 3 9 9 9", 77, 47),63    ]64    for user_prompt, n_prompt, n_cache in seq:65        res = server.make_request("POST", "/chat/completions", data={66            "max_tokens": 8,67            "messages": [68                {"role": "system", "content": "Test"},69                {"role": "user", "content": user_prompt},70            ],71        })72        assert res.body["usage"]["prompt_tokens"] == n_prompt73        assert res.body["usage"]["prompt_tokens_details"]["cached_tokens"] == n_cache74 75@pytest.mark.parametrize(76    "system_prompt,user_prompt,max_tokens,re_content,n_prompt,n_predicted,finish_reason",77    [78        ("Book", "What is the best book", 8, "(Suddenly)+", 77, 8, "length"),79        ("You are a coding assistant.", "Write the fibonacci function in c++.", 128, "(Aside|she|felter|alonger)+", 104, 128, "length"),80    ]81)82def test_chat_completion_stream(system_prompt, user_prompt, max_tokens, re_content, n_prompt, n_predicted, finish_reason):83    global server84    server.model_alias = "llama-test-model"85    server.start()86    res = server.make_stream_request("POST", "/chat/completions", data={87        "max_tokens": max_tokens,88        "messages": [89            {"role": "system", "content": system_prompt},90            {"role": "user", "content": user_prompt},91        ],92        "stream": True,93    })94    content = ""95    last_cmpl_id = None96    for i, data in enumerate(res):97        if data["choices"]:98            choice = data["choices"][0]99            if i == 0:100                # Check first role message for stream=True101                assert choice["delta"]["content"] is None102                assert choice["delta"]["role"] == "assistant"103            else:104                assert "role" not in choice["delta"]105            assert data["system_fingerprint"].startswith("b")106            assert data["model"] == "llama-test-model"107            if last_cmpl_id is None:108                last_cmpl_id = data["id"]109            assert last_cmpl_id == data["id"] # make sure the completion id is the same for all events in the stream110            if choice["finish_reason"] in ["stop", "length"]:111                assert "content" not in choice["delta"]112                assert match_regex(re_content, content)113                assert choice["finish_reason"] == finish_reason114            else:115                assert choice["finish_reason"] is None116                content += choice["delta"]["content"] or ''117        else:118            assert data["usage"]["prompt_tokens"] == n_prompt119            assert data["usage"]["completion_tokens"] == n_predicted120 121 122def test_chat_completion_with_openai_library():123    global server124    server.start()125    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")126    res = client.chat.completions.create(127        model="gpt-3.5-turbo-instruct",128        messages=[129            {"role": "system", "content": "Book"},130            {"role": "user", "content": "What is the best book"},131        ],132        max_tokens=8,133        seed=42,134        temperature=0.8,135    )136    assert res.system_fingerprint is not None and res.system_fingerprint.startswith("b")137    assert res.choices[0].finish_reason == "length"138    assert res.choices[0].message.content is not None139    assert match_regex("(Suddenly)+", res.choices[0].message.content)140 141 142def test_chat_template():143    global server144    server.chat_template = "llama3"145    server.debug = True  # to get the "__verbose" object in the response146    server.start()147    res = server.make_request("POST", "/chat/completions", data={148        "max_tokens": 8,149        "messages": [150            {"role": "system", "content": "Book"},151            {"role": "user", "content": "What is the best book"},152        ]153    })154    assert res.status_code == 200155    assert "__verbose" in res.body156    assert res.body["__verbose"]["prompt"] == "<s> <|start_header_id|>system<|end_header_id|>\n\nBook<|eot_id|><|start_header_id|>user<|end_header_id|>\n\nWhat is the best book<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"157 158 159@pytest.mark.parametrize("prefill,re_prefill", [160    ("Whill", "Whill"),161    ([{"type": "text", "text": "Wh"}, {"type": "text", "text": "ill"}], "Whill"),162])163def test_chat_template_assistant_prefill(prefill, re_prefill):164    global server165    server.chat_template = "llama3"166    server.debug = True  # to get the "__verbose" object in the response167    server.start()168    res = server.make_request("POST", "/chat/completions", data={169        "max_tokens": 8,170        "messages": [171            {"role": "system", "content": "Book"},172            {"role": "user", "content": "What is the best book"},173            {"role": "assistant", "content": prefill},174        ]175    })176    assert res.status_code == 200177    assert "__verbose" in res.body178    assert res.body["__verbose"]["prompt"] == f"<s> <|start_header_id|>system<|end_header_id|>\n\nBook<|eot_id|><|start_header_id|>user<|end_header_id|>\n\nWhat is the best book<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n{re_prefill}"179 180 181def test_apply_chat_template():182    global server183    server.chat_template = "command-r"184    server.start()185    res = server.make_request("POST", "/apply-template", data={186        "messages": [187            {"role": "system", "content": "You are a test."},188            {"role": "user", "content":"Hi there"},189        ]190    })191    assert res.status_code == 200192    assert "prompt" in res.body193    assert res.body["prompt"] == "<|START_OF_TURN_TOKEN|><|SYSTEM_TOKEN|>You are a test.<|END_OF_TURN_TOKEN|><|START_OF_TURN_TOKEN|><|USER_TOKEN|>Hi there<|END_OF_TURN_TOKEN|><|START_OF_TURN_TOKEN|><|CHATBOT_TOKEN|>"194 195 196@pytest.mark.parametrize("response_format,n_predicted,re_content", [197    ({"type": "json_object", "schema": {"const": "42"}}, 6, "\"42\""),198    ({"type": "json_object", "schema": {"items": [{"type": "integer"}]}}, 10, "[ -3000 ]"),199    ({"type": "json_schema", "json_schema": {"schema": {"const": "foooooo"}}}, 10, "\"foooooo\""),200    ({"type": "json_object"}, 10, "(\\{|John)+"),201    ({"type": "sound"}, 0, None),202    # invalid response format (expected to fail)203    ({"type": "json_object", "schema": 123}, 0, None),204    ({"type": "json_object", "schema": {"type": 123}}, 0, None),205    ({"type": "json_object", "schema": {"type": "hiccup"}}, 0, None),206])207def test_completion_with_response_format(response_format: dict, n_predicted: int, re_content: str | None):208    global server209    server.start()210    res = server.make_request("POST", "/chat/completions", data={211        "max_tokens": n_predicted,212        "messages": [213            {"role": "system", "content": "You are a coding assistant."},214            {"role": "user", "content": "Write an example"},215        ],216        "response_format": response_format,217    })218    if re_content is not None:219        assert res.status_code == 200220        choice = res.body["choices"][0]221        assert match_regex(re_content, choice["message"]["content"])222    else:223        assert res.status_code == 400224        assert "error" in res.body225 226 227@pytest.mark.parametrize("jinja,json_schema,n_predicted,re_content", [228    (False, {"const": "42"}, 6, "\"42\""),229    (True, {"const": "42"}, 6, "\"42\""),230])231def test_completion_with_json_schema(jinja: bool, json_schema: dict, n_predicted: int, re_content: str):232    global server233    server.jinja = jinja234    server.debug = True235    server.start()236    res = server.make_request("POST", "/chat/completions", data={237        "max_tokens": n_predicted,238        "messages": [239            {"role": "system", "content": "You are a coding assistant."},240            {"role": "user", "content": "Write an example"},241        ],242        "json_schema": json_schema,243    })244    assert res.status_code == 200, f'Expected 200, got {res.status_code}'245    choice = res.body["choices"][0]246    assert match_regex(re_content, choice["message"]["content"]), f'Expected {re_content}, got {choice["message"]["content"]}'247 248 249@pytest.mark.parametrize("jinja,grammar,n_predicted,re_content", [250    (False, 'root ::= "a"{5,5}', 6, "a{5,5}"),251    (True, 'root ::= "a"{5,5}', 6, "a{5,5}"),252])253def test_completion_with_grammar(jinja: bool, grammar: str, n_predicted: int, re_content: str):254    global server255    server.jinja = jinja256    server.start()257    res = server.make_request("POST", "/chat/completions", data={258        "max_tokens": n_predicted,259        "messages": [260            {"role": "user", "content": "Does not matter what I say, does it?"},261        ],262        "grammar": grammar,263    })264    assert res.status_code == 200, res.body265    choice = res.body["choices"][0]266    assert match_regex(re_content, choice["message"]["content"]), choice["message"]["content"]267 268 269@pytest.mark.parametrize("messages", [270    None,271    "string",272    [123],273    [{}],274    [{"role": 123}],275    [{"role": "system", "content": 123}],276    # [{"content": "hello"}], # TODO: should not be a valid case277    [{"role": "system", "content": "test"}, {}],278    [{"role": "user", "content": "test"}, {"role": "assistant", "content": "test"}, {"role": "assistant", "content": "test"}],279])280def test_invalid_chat_completion_req(messages):281    global server282    server.start()283    res = server.make_request("POST", "/chat/completions", data={284        "messages": messages,285    })286    assert res.status_code == 400 or res.status_code == 500287    assert "error" in res.body288 289 290def test_chat_completion_with_timings_per_token():291    global server292    server.start()293    res = server.make_stream_request("POST", "/chat/completions", data={294        "max_tokens": 10,295        "messages": [{"role": "user", "content": "test"}],296        "stream": True,297        "stream_options": {"include_usage": True},298        "timings_per_token": True,299    })300    stats_received = False301    for i, data in enumerate(res):302        if i == 0:303            # Check first role message for stream=True304            assert data["choices"][0]["delta"]["content"] is None305            assert data["choices"][0]["delta"]["role"] == "assistant"306            assert "timings" not in data, f'First event should not have timings: {data}'307        else:308            if data["choices"]:309                assert "role" not in data["choices"][0]["delta"]310            else:311                assert "timings" in data312                assert "prompt_per_second" in data["timings"]313                assert "predicted_per_second" in data["timings"]314                assert "predicted_n" in data["timings"]315                assert data["timings"]["predicted_n"] <= 10316                stats_received = True317    assert stats_received318 319 320def test_logprobs():321    global server322    server.start()323    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")324    res = client.chat.completions.create(325        model="gpt-3.5-turbo-instruct",326        temperature=0.0,327        messages=[328            {"role": "system", "content": "Book"},329            {"role": "user", "content": "What is the best book"},330        ],331        max_tokens=5,332        logprobs=True,333        top_logprobs=10,334    )335    output_text = res.choices[0].message.content336    aggregated_text = ''337    assert res.choices[0].logprobs is not None338    assert res.choices[0].logprobs.content is not None339    for token in res.choices[0].logprobs.content:340        aggregated_text += token.token341        assert token.logprob <= 0.0342        assert token.bytes is not None343        assert len(token.top_logprobs) > 0344    assert aggregated_text == output_text345 346 347def test_logprobs_stream():348    global server349    server.start()350    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")351    res = client.chat.completions.create(352        model="gpt-3.5-turbo-instruct",353        temperature=0.0,354        messages=[355            {"role": "system", "content": "Book"},356            {"role": "user", "content": "What is the best book"},357        ],358        max_tokens=5,359        logprobs=True,360        top_logprobs=10,361        stream=True,362    )363    output_text = ''364    aggregated_text = ''365    for i, data in enumerate(res):366        if data.choices:367            choice = data.choices[0]368            if i == 0:369                # Check first role message for stream=True370                assert choice.delta.content is None371                assert choice.delta.role == "assistant"372            else:373                assert choice.delta.role is None374                if choice.finish_reason is None:375                    if choice.delta.content:376                        output_text += choice.delta.content377                    assert choice.logprobs is not None378                    assert choice.logprobs.content is not None379                    for token in choice.logprobs.content:380                        aggregated_text += token.token381                        assert token.logprob <= 0.0382                        assert token.bytes is not None383                        assert token.top_logprobs is not None384                        assert len(token.top_logprobs) > 0385    assert aggregated_text == output_text386 387 388def test_logit_bias():389    global server390    server.start()391 392    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"]393 394    res = server.make_request("POST", "/tokenize", data={395        "content": " " + " ".join(exclude) + " ",396    })397    assert res.status_code == 200398    tokens = res.body["tokens"]399    logit_bias = {tok: -100 for tok in tokens}400 401    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")402    res = client.chat.completions.create(403        model="gpt-3.5-turbo-instruct",404        temperature=0.0,405        messages=[406            {"role": "system", "content": "Book"},407            {"role": "user", "content": "What is the best book"},408        ],409        max_tokens=64,410        logit_bias=logit_bias411    )412    output_text = res.choices[0].message.content413    assert output_text414    assert all(output_text.find(" " + tok + " ") == -1 for tok in exclude)415 416def test_context_size_exceeded():417    global server418    server.start()419    res = server.make_request("POST", "/chat/completions", data={420        "messages": [421            {"role": "system", "content": "Book"},422            {"role": "user", "content": "What is the best book"},423        ] * 100, # make the prompt too long424    })425    assert res.status_code == 400426    assert "error" in res.body427    assert res.body["error"]["type"] == "exceed_context_size_error"428    assert res.body["error"]["n_prompt_tokens"] > 0429    assert server.n_ctx is not None430    assert server.n_slots is not None431    assert res.body["error"]["n_ctx"] == server.n_ctx // server.n_slots432 433 434def test_context_size_exceeded_stream():435    global server436    server.start()437    try:438        for _ in server.make_stream_request("POST", "/chat/completions", data={439            "messages": [440                {"role": "system", "content": "Book"},441                {"role": "user", "content": "What is the best book"},442            ] * 100, # make the prompt too long443            "stream": True}):444                pass445        assert False, "Should have failed"446    except ServerError as e:447        assert e.code == 400448        assert "error" in e.body449        assert e.body["error"]["type"] == "exceed_context_size_error"450        assert e.body["error"]["n_prompt_tokens"] > 0451        assert server.n_ctx is not None452        assert server.n_slots is not None453        assert e.body["error"]["n_ctx"] == server.n_ctx // server.n_slots454 455 456@pytest.mark.parametrize(457    "n_batch,batch_count,reuse_cache",458    [459        (64, 4, False),460        (64, 2, True),461    ]462)463def test_return_progress(n_batch, batch_count, reuse_cache):464    global server465    server.n_batch = n_batch466    server.n_ctx = 256467    server.n_slots = 1468    server.start()469    def make_cmpl_request():470        return server.make_stream_request("POST", "/chat/completions", data={471            "max_tokens": 10,472            "messages": [473                {"role": "user", "content": "This is a test" * 10},474            ],475            "stream": True,476            "return_progress": True,477        })478    if reuse_cache:479        # make a first request to populate the cache480        res0 = make_cmpl_request()481        for _ in res0:482            pass # discard the output483 484    res = make_cmpl_request()485    last_progress = None486    total_batch_count = 0487 488    for data in res:489        cur_progress = data.get("prompt_progress", None)490        if cur_progress is None:491            continue492        if total_batch_count == 0:493            # first progress report must have n_cache == n_processed494            assert cur_progress["total"] > 0495            assert cur_progress["cache"] == cur_progress["processed"]496            if reuse_cache:497                # when reusing cache, we expect some cached tokens498                assert cur_progress["cache"] > 0499        if last_progress is not None:500            assert cur_progress["total"] == last_progress["total"]501            assert cur_progress["cache"] == last_progress["cache"]502            assert cur_progress["processed"] > last_progress["processed"]503        total_batch_count += 1504        last_progress = cur_progress505 506    # last progress should indicate completion (all tokens processed)507    assert last_progress is not None508    assert last_progress["total"] > 0509    assert last_progress["processed"] == last_progress["total"]510    assert total_batch_count == batch_count511 512 513def test_chat_completions_multiple_choices():514    global server515    server.start()516    # make sure cache can be reused across multiple choices and multiple requests517    # ref: https://github.com/ggml-org/llama.cpp/pull/18663518    for _ in range(2):519        res = server.make_request("POST", "/chat/completions", data={520            "max_tokens": 8,521            "n": 2,522            "messages": [523                {"role": "system", "content": "Book"},524                {"role": "user", "content": "What is the best book"},525            ],526            # test forcing the same slot to be used527            # the scheduler should not be locked up in this case528            "id_slot": 0,529        })530        assert res.status_code == 200531        assert len(res.body["choices"]) == 2532        for choice in res.body["choices"]:533            assert "assistant" == choice["message"]["role"]534            assert choice["finish_reason"] == "length"535