Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
test_chat_completion.py268 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, 64, "length", False, None),23        ("codellama70b", "You are a coding assistant.", "Write the fibonacci function in c++.", 128, "(Aside|she|felter|alonger)+", 104, 64, "length", True, None),24    ]25)26def test_chat_completion(model, system_prompt, user_prompt, max_tokens, re_content, n_prompt, n_predicted, finish_reason, jinja, chat_template):27    global server28    server.jinja = jinja29    server.chat_template = chat_template30    server.start()31    res = server.make_request("POST", "/chat/completions", data={32        "model": model,33        "max_tokens": max_tokens,34        "messages": [35            {"role": "system", "content": system_prompt},36            {"role": "user", "content": user_prompt},37        ],38    })39    assert res.status_code == 20040    assert "cmpl" in res.body["id"] # make sure the completion id has the expected format41    assert res.body["system_fingerprint"].startswith("b")42    assert res.body["model"] == model if model is not None else server.model_alias43    assert res.body["usage"]["prompt_tokens"] == n_prompt44    assert res.body["usage"]["completion_tokens"] == n_predicted45    choice = res.body["choices"][0]46    assert "assistant" == choice["message"]["role"]47    assert match_regex(re_content, choice["message"]["content"])48    assert choice["finish_reason"] == finish_reason49 50 51@pytest.mark.parametrize(52    "system_prompt,user_prompt,max_tokens,re_content,n_prompt,n_predicted,finish_reason",53    [54        ("Book", "What is the best book", 8, "(Suddenly)+", 77, 8, "length"),55        ("You are a coding assistant.", "Write the fibonacci function in c++.", 128, "(Aside|she|felter|alonger)+", 104, 64, "length"),56    ]57)58def test_chat_completion_stream(system_prompt, user_prompt, max_tokens, re_content, n_prompt, n_predicted, finish_reason):59    global server60    server.model_alias = None # try using DEFAULT_OAICOMPAT_MODEL61    server.start()62    res = server.make_stream_request("POST", "/chat/completions", data={63        "max_tokens": max_tokens,64        "messages": [65            {"role": "system", "content": system_prompt},66            {"role": "user", "content": user_prompt},67        ],68        "stream": True,69    })70    content = ""71    last_cmpl_id = None72    for data in res:73        choice = data["choices"][0]74        assert data["system_fingerprint"].startswith("b")75        assert "gpt-3.5" in data["model"] # DEFAULT_OAICOMPAT_MODEL, maybe changed in the future76        if last_cmpl_id is None:77            last_cmpl_id = data["id"]78        assert last_cmpl_id == data["id"] # make sure the completion id is the same for all events in the stream79        if choice["finish_reason"] in ["stop", "length"]:80            assert data["usage"]["prompt_tokens"] == n_prompt81            assert data["usage"]["completion_tokens"] == n_predicted82            assert "content" not in choice["delta"]83            assert match_regex(re_content, content)84            assert choice["finish_reason"] == finish_reason85        else:86            assert choice["finish_reason"] is None87            content += choice["delta"]["content"]88 89 90def test_chat_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.chat.completions.create(95        model="gpt-3.5-turbo-instruct",96        messages=[97            {"role": "system", "content": "Book"},98            {"role": "user", "content": "What is the best book"},99        ],100        max_tokens=8,101        seed=42,102        temperature=0.8,103    )104    assert res.system_fingerprint is not None and res.system_fingerprint.startswith("b")105    assert res.choices[0].finish_reason == "length"106    assert res.choices[0].message.content is not None107    assert match_regex("(Suddenly)+", res.choices[0].message.content)108 109 110def test_chat_template():111    global server112    server.chat_template = "llama3"113    server.debug = True  # to get the "__verbose" object in the response114    server.start()115    res = server.make_request("POST", "/chat/completions", data={116        "max_tokens": 8,117        "messages": [118            {"role": "system", "content": "Book"},119            {"role": "user", "content": "What is the best book"},120        ]121    })122    assert res.status_code == 200123    assert "__verbose" in res.body124    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"125 126 127def test_apply_chat_template():128    global server129    server.chat_template = "command-r"130    server.start()131    res = server.make_request("POST", "/apply-template", data={132        "messages": [133            {"role": "system", "content": "You are a test."},134            {"role": "user", "content":"Hi there"},135        ]136    })137    assert res.status_code == 200138    assert "prompt" in res.body139    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|>"140 141 142@pytest.mark.parametrize("response_format,n_predicted,re_content", [143    ({"type": "json_object", "schema": {"const": "42"}}, 6, "\"42\""),144    ({"type": "json_object", "schema": {"items": [{"type": "integer"}]}}, 10, "[ -3000 ]"),145    ({"type": "json_object"}, 10, "(\\{|John)+"),146    ({"type": "sound"}, 0, None),147    # invalid response format (expected to fail)148    ({"type": "json_object", "schema": 123}, 0, None),149    ({"type": "json_object", "schema": {"type": 123}}, 0, None),150    ({"type": "json_object", "schema": {"type": "hiccup"}}, 0, None),151])152def test_completion_with_response_format(response_format: dict, n_predicted: int, re_content: str | None):153    global server154    server.start()155    res = server.make_request("POST", "/chat/completions", data={156        "max_tokens": n_predicted,157        "messages": [158            {"role": "system", "content": "You are a coding assistant."},159            {"role": "user", "content": "Write an example"},160        ],161        "response_format": response_format,162    })163    if re_content is not None:164        assert res.status_code == 200165        choice = res.body["choices"][0]166        assert match_regex(re_content, choice["message"]["content"])167    else:168        assert res.status_code != 200169        assert "error" in res.body170 171 172@pytest.mark.parametrize("messages", [173    None,174    "string",175    [123],176    [{}],177    [{"role": 123}],178    [{"role": "system", "content": 123}],179    # [{"content": "hello"}], # TODO: should not be a valid case180    [{"role": "system", "content": "test"}, {}],181])182def test_invalid_chat_completion_req(messages):183    global server184    server.start()185    res = server.make_request("POST", "/chat/completions", data={186        "messages": messages,187    })188    assert res.status_code == 400 or res.status_code == 500189    assert "error" in res.body190 191 192def test_chat_completion_with_timings_per_token():193    global server194    server.start()195    res = server.make_stream_request("POST", "/chat/completions", data={196        "max_tokens": 10,197        "messages": [{"role": "user", "content": "test"}],198        "stream": True,199        "timings_per_token": True,200    })201    for data in res:202        assert "timings" in data203        assert "prompt_per_second" in data["timings"]204        assert "predicted_per_second" in data["timings"]205        assert "predicted_n" in data["timings"]206        assert data["timings"]["predicted_n"] <= 10207 208 209def test_logprobs():210    global server211    server.start()212    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")213    res = client.chat.completions.create(214        model="gpt-3.5-turbo-instruct",215        temperature=0.0,216        messages=[217            {"role": "system", "content": "Book"},218            {"role": "user", "content": "What is the best book"},219        ],220        max_tokens=5,221        logprobs=True,222        top_logprobs=10,223    )224    output_text = res.choices[0].message.content225    aggregated_text = ''226    assert res.choices[0].logprobs is not None227    assert res.choices[0].logprobs.content is not None228    for token in res.choices[0].logprobs.content:229        aggregated_text += token.token230        assert token.logprob <= 0.0231        assert token.bytes is not None232        assert len(token.top_logprobs) > 0233    assert aggregated_text == output_text234 235 236def test_logprobs_stream():237    global server238    server.start()239    client = OpenAI(api_key="dummy", base_url=f"http://{server.server_host}:{server.server_port}/v1")240    res = client.chat.completions.create(241        model="gpt-3.5-turbo-instruct",242        temperature=0.0,243        messages=[244            {"role": "system", "content": "Book"},245            {"role": "user", "content": "What is the best book"},246        ],247        max_tokens=5,248        logprobs=True,249        top_logprobs=10,250        stream=True,251    )252    output_text = ''253    aggregated_text = ''254    for data in res:255        choice = data.choices[0]256        if choice.finish_reason is None:257            if choice.delta.content:258                output_text += choice.delta.content259            assert choice.logprobs is not None260            assert choice.logprobs.content is not None261            for token in choice.logprobs.content:262                aggregated_text += token.token263                assert token.logprob <= 0.0264                assert token.bytes is not None265                assert token.top_logprobs is not None266                assert len(token.top_logprobs) > 0267    assert aggregated_text == output_text268