KBaba7/llama.cpp
0
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 