Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
test_speculative.py206 linesDownload Raw Back to unit
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