Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
test_slot_save.py549 linesDownload Raw Back to unit
1import pytest2from utils import *3import base644import requests5import struct6 7# sequence state file: magic(4) version(4) payload_size(4), then payload_size llama_token words8STATE_FILE_HEADER_SIZE = 129 10server = ServerPreset.tinyllama2()11 12@pytest.fixture(autouse=True)13def create_server(tmp_path):14    global server15    server = ServerPreset.tinyllama2()16    server.slot_save_path = str(tmp_path)17    server.temperature = 0.018 19 20def test_slot_save_restore():21    global server22    server.start()23 24    # First prompt in slot 1 should be fully processed25    res = server.make_request("POST", "/completion", data={26        "prompt": "What is the capital of France?",27        "id_slot": 1,28        "cache_prompt": True,29    })30    assert res.status_code == 20031    assert match_regex("(Whiskers|Flana)+", res.body["content"])32    assert res.body["timings"]["prompt_n"] == 21  # all tokens are processed33 34    # Save state of slot 135    res = server.make_request("POST", "/slots/1?action=save", data={36        "filename": "slot1.bin",37    })38    assert res.status_code == 20039    assert res.body["n_saved"] == 8440 41    # Since we have cache, this should only process the last tokens42    res = server.make_request("POST", "/completion", data={43        "prompt": "What is the capital of Germany?",44        "id_slot": 1,45        "cache_prompt": True,46    })47    assert res.status_code == 20048    assert match_regex("(Jack|said)+", res.body["content"])49    assert res.body["timings"]["prompt_n"] == 6  # only different part is processed50 51    # Loading the saved cache into slot 052    res = server.make_request("POST", "/slots/0?action=restore", data={53        "filename": "slot1.bin",54    })55    assert res.status_code == 20056    assert res.body["n_restored"] == 8457 58    # Since we have cache, slot 0 should only process the last tokens59    res = server.make_request("POST", "/completion", data={60        "prompt": "What is the capital of Germany?",61        "id_slot": 0,62        "cache_prompt": True,63    })64    assert res.status_code == 20065    assert match_regex("(Jack|said)+", res.body["content"])66    assert res.body["timings"]["prompt_n"] == 6  # only different part is processed67 68    # For verification that slot 1 was not corrupted during slot 0 load, same thing should work69    res = server.make_request("POST", "/completion", data={70        "prompt": "What is the capital of Germany?",71        "id_slot": 1,72        "cache_prompt": True,73    })74    assert res.status_code == 20075    assert match_regex("(Jack|said)+", res.body["content"])76    assert res.body["timings"]["prompt_n"] == 177 78 79def test_slot_restore_legacy_token_list():80    global server81    server.start()82 83    res = server.make_request("POST", "/completion", data={84        "prompt": "What is the capital of France?",85        "id_slot": 1,86        "cache_prompt": True,87    })88    assert res.status_code == 20089 90    res = server.make_request("POST", "/slots/1?action=save", data={91        "filename": "slot_legacy.bin",92    })93    assert res.status_code == 20094    assert res.body["n_saved"] == 8495 96    # rewrite the token payload into a plain token list, as written by servers that predate the packed server_tokens format97    path = os.path.join(server.slot_save_path, "slot_legacy.bin")98    with open(path, "rb") as f:99        data = bytearray(f.read())100 101    # the payload written by this server starts with a packed header: LLAMA_TOKEN_NULL(4) version(4) n_tokens(4)102    packed_header_size = 12103 104    payload_size = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE - 4)[0]105    payload_end = STATE_FILE_HEADER_SIZE + payload_size * 4106    n_tokens = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE + 8)[0]107    assert n_tokens == 84108 109    tokens_start = STATE_FILE_HEADER_SIZE + packed_header_size110    data = data[:STATE_FILE_HEADER_SIZE] + data[tokens_start:tokens_start + n_tokens * 4] + data[payload_end:]111    struct.pack_into("=I", data, STATE_FILE_HEADER_SIZE - 4, n_tokens)112 113    with open(path, "wb") as f:114        f.write(data)115 116    # the plain token list must restore, and the restored KV must be reusable117    res = server.make_request("POST", "/slots/0?action=restore", data={118        "filename": "slot_legacy.bin",119    })120    assert res.status_code == 200121    assert res.body["n_restored"] == 84122 123    res = server.make_request("POST", "/completion", data={124        "prompt": "What is the capital of Germany?",125        "id_slot": 0,126        "cache_prompt": True,127    })128    assert res.status_code == 200129    assert res.body["timings"]["prompt_n"] == 6  # only the different part is processed130 131 132 133def test_slot_erase():134    global server135    server.start()136 137    res = server.make_request("POST", "/completion", data={138        "prompt": "What is the capital of France?",139        "id_slot": 1,140        "cache_prompt": True,141    })142    assert res.status_code == 200143    assert match_regex("(Whiskers|Flana)+", res.body["content"])144    assert res.body["timings"]["prompt_n"] == 21  # all tokens are processed145 146    # erase slot 1147    res = server.make_request("POST", "/slots/1?action=erase")148    assert res.status_code == 200149 150    # re-run the same prompt, it should process all tokens again151    res = server.make_request("POST", "/completion", data={152        "prompt": "What is the capital of France?",153        "id_slot": 1,154        "cache_prompt": True,155    })156    assert res.status_code == 200157    assert match_regex("(Whiskers|Flana)+", res.body["content"])158    assert res.body["timings"]["prompt_n"] == 21  # all tokens are processed159 160 161#162# Multimodal server (mmproj loaded) slot save/restore.163#164# A pure-text slot on a multimodal server and a slot containing images must both support save/restore.165# Erase remains gated on the slot's content.166#167 168IMG_URL_CAT = "https://huggingface.co/ggml-org/tinygemma3-GGUF/resolve/main/test/91_cat.png"169IMG_URL_TRUCK = "https://huggingface.co/ggml-org/tinygemma3-GGUF/resolve/main/test/11_truck.png"170 171 172def _get_img_base64(url: str) -> str:173    response = requests.get(url)174    response.raise_for_status()  # Raise an exception for bad status codes175    return base64.b64encode(response.content).decode("utf-8")176 177 178@pytest.fixture179def mmproj_server():180    # tinygemma3 is a small multimodal model: the mmproj is provided by the HF registry API and auto-downloaded on first run.181    os.environ['LLAMA_MEDIA_MARKER'] = '<__media__>'182    mm_server = ServerPreset.tinygemma3()183    mm_server.slot_save_path = "./tmp"184    mm_server.temperature = 0.0185    return mm_server186 187 188def test_slot_save_restore_text_only_on_multimodal(mmproj_server):189    server = mmproj_server190    server.start()191 192    # A pure-text prompt processed on slot 1 of a multimodal server.193    res = server.make_request("POST", "/completion", data={194        "prompt": "The quick brown fox jumps over the lazy dog.",195        "id_slot": 1,196        "cache_prompt": True,197    })198    assert res.status_code == 200199    prompt_n = res.body["timings"]["prompt_n"]200    assert prompt_n > 0  # all tokens are processed201 202    # Saving a pure-text slot must succeed even though an mmproj is loaded.203    res = server.make_request("POST", "/slots/1?action=save", data={204        "filename": "mm_slot1.bin",205    })206    assert res.status_code == 200207    n_saved = res.body["n_saved"]208    assert n_saved > 0  # the slot KV (prompt + generated tokens) was written209 210    # Restore the saved state into slot 0; it must round-trip exactly.211    res = server.make_request("POST", "/slots/0?action=restore", data={212        "filename": "mm_slot1.bin",213    })214    assert res.status_code == 200215    assert res.body["n_restored"] == n_saved216 217    # Prefix reuse is not checked with the default SWA cache.218    res = server.make_request("POST", "/completion", data={219        "prompt": "The quick brown fox jumps over the lazy dog.",220        "id_slot": 0,221        "cache_prompt": True,222    })223    assert res.status_code == 200224 225 226def test_slot_save_restore_with_image(mmproj_server):227    server = mmproj_server228    # Use the full SWA cache so the restored image prefix can be reused.229    server.swa_full = True230    server.start()231 232    prompt_cat = {233        "prompt_string": "What is this: <__media__>\n",234        "multimodal_data": [_get_img_base64(IMG_URL_CAT)],235    }236    res = server.make_request("POST", "/completions", data={237        "temperature": 0.0,238        "top_k": 1,239        "id_slot": 1,240        "cache_prompt": True,241        "prompt": prompt_cat,242    })243    assert res.status_code == 200244    content_cat = res.body["content"]245    prompt_n_full = res.body["timings"]["prompt_n"]246    assert res.body["timings"]["cache_n"] == 0247    assert prompt_n_full > 32  # text plus image tokens are all processed248 249    res = server.make_request("POST", "/slots/1?action=save", data={250        "filename": "mm_slot_image.bin",251    })252    assert res.status_code == 200253    n_saved = res.body["n_saved"]254    n_written = res.body["n_written"]255    assert n_saved > 0256    assert n_written > 0257 258    res = server.make_request("POST", "/slots/1?action=erase")259    assert res.status_code == 200260 261    res = server.make_request("POST", "/slots/0?action=restore", data={262        "filename": "mm_slot_image.bin",263    })264    assert res.status_code == 200265    assert res.body["n_restored"] == n_saved266    assert res.body["n_read"] == n_written267 268    # a different image must not reuse the restored image tokens; only the text prefix before the image is common269    res = server.make_request("POST", "/completions", data={270        "temperature": 0.0,271        "top_k": 1,272        "id_slot": 0,273        "cache_prompt": True,274        "prompt": {275            "prompt_string": "What is this: <__media__>\n",276            "multimodal_data": [_get_img_base64(IMG_URL_TRUCK)],277        },278    })279    assert res.status_code == 200280    cache_n = res.body["timings"]["cache_n"]281    assert cache_n < 16282    assert res.body["timings"]["prompt_n"] == prompt_n_full - cache_n283 284    # restore again and resend the same image: the image tokens must be reused and greedy sampling must reproduce the original content285    res = server.make_request("POST", "/slots/0?action=restore", data={286        "filename": "mm_slot_image.bin",287    })288    assert res.status_code == 200289    assert res.body["n_restored"] == n_saved290 291    res = server.make_request("POST", "/completions", data={292        "temperature": 0.0,293        "top_k": 1,294        "id_slot": 0,295        "cache_prompt": True,296        "prompt": prompt_cat,297    })298    assert res.status_code == 200299    assert res.body["timings"]["cache_n"] == prompt_n_full - 1300    assert res.body["timings"]["prompt_n"] == 1301    assert res.body["content"] == content_cat302 303 304def test_slot_save_restore_with_two_images(mmproj_server):305    server = mmproj_server306    server.swa_full = True307    server.n_ctx = 2048  # two images need more than the default 512 per slot308    server.start()309 310    prompt = {311        "prompt_string": "A: <__media__> B: <__media__>\n",312        "multimodal_data": [_get_img_base64(IMG_URL_CAT), _get_img_base64(IMG_URL_TRUCK)],313    }314    res = server.make_request("POST", "/completions", data={315        "temperature": 0.0,316        "top_k": 1,317        "id_slot": 1,318        "cache_prompt": True,319        "prompt": prompt,320    })321    assert res.status_code == 200322    prompt_n_full = res.body["timings"]["prompt_n"]323    assert prompt_n_full > 64324 325    res = server.make_request("POST", "/slots/1?action=save", data={326        "filename": "mm_slot_two_images.bin",327    })328    assert res.status_code == 200329    n_saved = res.body["n_saved"]330 331    res = server.make_request("POST", "/slots/0?action=restore", data={332        "filename": "mm_slot_two_images.bin",333    })334    assert res.status_code == 200335    assert res.body["n_restored"] == n_saved336 337    res = server.make_request("POST", "/completions", data={338        "temperature": 0.0,339        "top_k": 1,340        "id_slot": 0,341        "cache_prompt": True,342        "prompt": prompt,343    })344    assert res.status_code == 200345    assert res.body["timings"]["cache_n"] == prompt_n_full - 1346    assert res.body["timings"]["prompt_n"] == 1347    content = res.body["content"]348 349    res = server.make_request("POST", "/slots/1?action=restore", data={350        "filename": "mm_slot_two_images.bin",351    })352    assert res.status_code == 200353    assert res.body["n_restored"] == n_saved354 355    res = server.make_request("POST", "/completions", data={356        "temperature": 0.0,357        "top_k": 1,358        "id_slot": 0,359        "cache_prompt": True,360        "prompt": prompt,361    })362    assert res.status_code == 200363    assert res.body["timings"]["cache_n"] == prompt_n_full - 1364    assert res.body["timings"]["prompt_n"] == 1365    content = res.body["content"]366 367    assert res.body["content"] == content368 369 370def test_slot_save_restore_with_image_across_restart(mmproj_server):371    server = mmproj_server372    server.swa_full = True373    server.start()374 375    prompt_cat = {376        "prompt_string": "What is this: <__media__>\n",377        "multimodal_data": [_get_img_base64(IMG_URL_CAT)],378    }379    res = server.make_request("POST", "/completions", data={380        "temperature": 0.0,381        "top_k": 1,382        "id_slot": 0,383        "cache_prompt": True,384        "prompt": prompt_cat,385    })386    assert res.status_code == 200387    content = res.body["content"]388    prompt_n_full = res.body["timings"]["prompt_n"]389 390    res = server.make_request("POST", "/slots/0?action=save", data={391        "filename": "mm_slot_restart.bin",392    })393    assert res.status_code == 200394    n_saved = res.body["n_saved"]395 396    # restart the server with the same model and mmproj: the saved file must restore in the new process and the image KV must be reused397    server.stop()398    server.start()399 400    res = server.make_request("POST", "/slots/0?action=restore", data={401        "filename": "mm_slot_restart.bin",402    })403    assert res.status_code == 200404    assert res.body["n_restored"] == n_saved405 406    res = server.make_request("POST", "/completions", data={407        "temperature": 0.0,408        "top_k": 1,409        "id_slot": 0,410        "cache_prompt": True,411        "prompt": prompt_cat,412    })413    assert res.status_code == 200414    assert res.body["timings"]["cache_n"] == prompt_n_full - 1415    assert res.body["timings"]["prompt_n"] == 1416    assert res.body["content"] == content417 418 419def test_slot_save_restore_image_payload_larger_than_context(mmproj_server):420    server = mmproj_server421    server.swa_full = True422    server.start()423 424    # the slot context, as the server computed it (n_ctx split across the slots)425    res = server.make_request("GET", "/props")426    assert res.status_code == 200427    n_ctx_slot = res.body["default_generation_settings"]["n_ctx"]428 429    # a filler token, used to grow the prompt up to the slot context430    res = server.make_request("POST", "/tokenize", data={"content": " hello" * 8})431    assert res.status_code == 200432    assert len(res.body["tokens"]) == 8433 434    res = server.make_request("POST", "/completions", data={435        "temperature": 0.0,436        "top_k": 1,437        "id_slot": 0,438        "cache_prompt": True,439        "prompt": {440            "prompt_string": "What is this: <__media__>\n",441            "multimodal_data": [_get_img_base64(IMG_URL_CAT)],442        },443    })444    assert res.status_code == 200445 446    prompt_cat = {447        "prompt_string": "What is this: <__media__>\n" + " hello" * (n_ctx_slot - res.body["timings"]["prompt_n"] - 8),448        "multimodal_data": [_get_img_base64(IMG_URL_CAT)],449    }450    res = server.make_request("POST", "/completions", data={451        "temperature": 0.0,452        "top_k": 1,453        "id_slot": 0,454        "cache_prompt": True,455        "prompt": prompt_cat,456    })457    assert res.status_code == 200458    prompt_n_full = res.body["timings"]["cache_n"] + res.body["timings"]["prompt_n"]459 460    res = server.make_request("POST", "/slots/0?action=save", data={461        "filename": "mm_slot_large_payload.bin",462    })463    assert res.status_code == 200464 465    path = os.path.join(server.slot_save_path, "mm_slot_large_payload.bin")466    with open(path, "rb") as f:467        data = bytearray(f.read())468    payload_size = struct.unpack_from("=I", data, STATE_FILE_HEADER_SIZE - 4)[0]469    assert payload_size > n_ctx_slot  # the scenario under test: the payload does not fit in n_ctx470 471    # drop the image from the slot, then restore it from the file472    res = server.make_request("POST", "/completion", data={473        "prompt": "The quick brown fox",474        "id_slot": 0,475        "cache_prompt": True,476    })477    assert res.status_code == 200478 479    res = server.make_request("POST", "/slots/0?action=restore", data={480        "filename": "mm_slot_large_payload.bin",481    })482    assert res.status_code == 200483 484    res = server.make_request("POST", "/completions", data={485        "temperature": 0.0,486        "top_k": 1,487        "id_slot": 0,488        "cache_prompt": True,489        "prompt": prompt_cat,490    })491    assert res.status_code == 200492    assert res.body["timings"]["cache_n"] == prompt_n_full - 1493    assert res.body["timings"]["prompt_n"] == 1494 495 496def test_slot_restore_media_file_without_mmproj(mmproj_server):497    server = mmproj_server498    server.start()499 500    res = server.make_request("POST", "/completions", data={501        "temperature": 0.0,502        "top_k": 1,503        "id_slot": 0,504        "cache_prompt": True,505        "prompt": {506            "prompt_string": "What is this: <__media__>\n",507            "multimodal_data": [_get_img_base64(IMG_URL_CAT)],508        },509    })510    assert res.status_code == 200511 512    res = server.make_request("POST", "/slots/0?action=save", data={513        "filename": "mm_slot_no_mmproj.bin",514    })515    assert res.status_code == 200516 517    # restart the same model without the mmproj: restoring the media file must fail gracefully and leave the slot usable518    server.stop()519    server.no_mmproj = True520    server.start()521 522    res = server.make_request("POST", "/slots/0?action=restore", data={523        "filename": "mm_slot_no_mmproj.bin",524    })525    assert res.status_code == 400526    assert "Cannot restore media tokens without an mmproj" in res.body["error"]["message"]527 528    # A failed restore must leave the slot empty and usable.529    res = server.make_request("POST", "/completions", data={530        "temperature": 0.0,531        "top_k": 1,532        "id_slot": 1,533        "cache_prompt": True,534        "prompt": "The quick brown fox",535    })536    assert res.status_code == 200537    content = res.body["content"]538 539    res = server.make_request("POST", "/completions", data={540        "temperature": 0.0,541        "top_k": 1,542        "id_slot": 0,543        "cache_prompt": True,544        "prompt": "The quick brown fox",545    })546    assert res.status_code == 200547    assert res.body["timings"]["cache_n"] == 0548    assert res.body["content"] == content549