Felipe97/llama-cpp-compiled
01.2k
1import pytest2from utils import *3 4# We use a F16 MOE gguf as main model, and q4_0 as draft model5 6server = ServerPreset.stories15m_moe()7 8MODEL_DRAFT_FILE_URL = "https://huggingface.co/ggml-org/tiny-llamas/resolve/main/stories15M-q4_0.gguf"9 10def create_server():11 global server12 server = ServerPreset.stories15m_moe()13 # set default values14 server.model_draft = download_file(MODEL_DRAFT_FILE_URL)15 server.spec_type = "draft-simple"16 server.spec_draft_n_min = 417 server.spec_draft_n_max = 818 server.fa = "off"19 20 21@pytest.fixture(autouse=True)22def fixture_create_server():23 return create_server()24 25 26def test_with_and_without_draft():27 global server28 request = {29 "prompt": "I believe the meaning of life is",30 "temperature": 0.2,31 "top_k": 5,32 "seed": 4242,33 "n_predict": 16,34 "return_tokens": True,35 }36 37 server.model_draft = None # disable draft model38 server.spec_type = None39 server.start()40 res = server.make_request("POST", "/completion", data=request)41 assert res.status_code == 20042 tokens_no_draft = res.body["tokens"]43 server.stop()44 45 # create new server with draft model46 create_server()47 server.start()48 res = server.make_request("POST", "/completion", data=request)49 assert res.status_code == 20050 assert res.body["timings"]["draft_n"] > 051 tokens_draft = res.body["tokens"]52 53 assert tokens_no_draft == tokens_draft54 55 server.stop()56 create_server()57 assert server.spec_draft_n_max is not None58 server.spec_synth_rates = [0.0] * server.spec_draft_n_max59 server.start()60 res = server.make_request("POST", "/completion", data=request)61 62 assert res.status_code == 20063 assert res.body["timings"]["draft_n"] > 064 assert res.body["timings"]["draft_n_accepted"] == 065 assert res.body["tokens"] == tokens_no_draft66 67 68def test_different_draft_min_draft_max():69 global server70 test_values = [71 (1, 2),72 (1, 4),73 (4, 8),74 (4, 12),75 (8, 16),76 ]77 last_content = None78 for draft_min, draft_max in test_values:79 server.stop()80 server.spec_draft_n_min = draft_min81 server.spec_draft_n_max = draft_max82 server.start()83 res = server.make_request("POST", "/completion", data={84 "prompt": "I believe the meaning of life is",85 "temperature": 0.0,86 "top_k": 1,87 "n_predict": 16,88 })89 assert res.status_code == 20090 if last_content is not None:91 assert last_content == res.body["content"]92 last_content = res.body["content"]93 94 95def test_synth_is_deterministic():96 global server97 assert server.spec_draft_n_max is not None98 server.spec_synth_rates = [0.75 ** (i + 1) for i in range(server.spec_draft_n_max)]99 server.start()100 101 request = {102 "prompt": "I believe the meaning of life is",103 "temperature": 0.2,104 "top_k": 5,105 "seed": 4242,106 "n_predict": 32,107 }108 responses = [server.make_request("POST", "/completion", data=request) for _ in range(2)]109 110 for res in responses:111 assert res.status_code == 200112 assert res.body["timings"]["draft_n"] > 0113 assert responses[0].body["timings"]["draft_n"] == responses[1].body["timings"]["draft_n"]114 assert responses[0].body["timings"]["draft_n_accepted"] == responses[1].body["timings"]["draft_n_accepted"]115 116 117def test_synth_ignores_target_tokens():118 global server119 assert server.spec_draft_n_max is not None120 server.spec_synth_rates = [1.0] * server.spec_draft_n_max121 server.start()122 123 res = server.make_request("POST", "/completion", data={124 "prompt": "I believe the meaning of life is",125 "temperature": 0.0,126 "seed": 4242,127 "n_predict": 32,128 })129 130 assert res.status_code == 200131 assert res.body["timings"]["draft_n"] > 0132 assert res.body["timings"]["draft_n_accepted"] == res.body["timings"]["draft_n"]133 134 res = server.make_request("POST", "/completion", data={135 "prompt": "I believe the meaning of life is",136 "temperature": 0.0,137 "seed": 4242,138 "n_predict": 6,139 "grammar": 'root ::= "a"{5,5}',140 })141 assert res.status_code == 200, res.body142 143 res = server.make_request("POST", "/completion", data={144 "prompt": "Respond with only: OK",145 "temperature": 0.0,146 "seed": 4242,147 "n_predict": 64,148 "ignore_eos": True,149 })150 assert res.status_code == 200, res.body151 assert res.body["tokens_predicted"] == 64152 assert res.body["stop_type"] == "limit"153 154 155def test_slot_ctx_not_exceeded():156 global server157 server.n_ctx = 256158 server.start()159 res = server.make_request("POST", "/completion", data={160 "prompt": "Hello " * 248,161 "temperature": 0.0,162 "top_k": 1,163 "speculative.p_min": 0.0,164 })165 assert res.status_code == 200166 assert len(res.body["content"]) > 0167 168 169def test_with_ctx_shift():170 global server171 server.n_ctx = 256172 server.enable_ctx_shift = True173 server.start()174 res = server.make_request("POST", "/completion", data={175 "prompt": "Hello " * 248,176 "temperature": 0.0,177 "top_k": 1,178 "n_predict": 256,179 "speculative.p_min": 0.0,180 })181 assert res.status_code == 200182 assert len(res.body["content"]) > 0183 assert res.body["tokens_predicted"] == 256184 assert res.body["truncated"] == True185 186 187@pytest.mark.parametrize("n_slots,n_requests", [188 (1, 2),189 (2, 2),190])191def test_multi_requests_parallel(n_slots: int, n_requests: int):192 global server193 server.n_slots = n_slots194 server.start()195 tasks = []196 for _ in range(n_requests):197 tasks.append((server.make_request, ("POST", "/completion", {198 "prompt": "I believe the meaning of life is",199 "temperature": 0.0,200 "top_k": 1,201 })))202 results = parallel_function_calls(tasks)203 for res in results:204 assert res.status_code == 200205 assert match_regex("(wise|kind|owl|answer)+", res.body["content"])206 