Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
test_vision_api.py182 linesDownload Raw Back to unit
1import pytest2from utils import *3import base644import requests5 6server: ServerProcess7 8def get_img_url(id: str) -> str:9    IMG_URL_0 = "https://huggingface.co/ggml-org/tinygemma3-GGUF/resolve/main/test/11_truck.png"10    IMG_URL_1 = "https://huggingface.co/ggml-org/tinygemma3-GGUF/resolve/main/test/91_cat.png"11    if id == "IMG_URL_0":12        return IMG_URL_013    elif id == "IMG_URL_1":14        return IMG_URL_115    elif id == "IMG_BASE64_URI_0":16        response = requests.get(IMG_URL_0)17        response.raise_for_status() # Raise an exception for bad status codes18        return "data:image/png;base64," + base64.b64encode(response.content).decode("utf-8")19    elif id == "IMG_BASE64_0":20        response = requests.get(IMG_URL_0)21        response.raise_for_status() # Raise an exception for bad status codes22        return base64.b64encode(response.content).decode("utf-8")23    elif id == "IMG_BASE64_URI_1":24        response = requests.get(IMG_URL_1)25        response.raise_for_status() # Raise an exception for bad status codes26        return "data:image/png;base64," + base64.b64encode(response.content).decode("utf-8")27    elif id == "IMG_BASE64_1":28        response = requests.get(IMG_URL_1)29        response.raise_for_status() # Raise an exception for bad status codes30        return base64.b64encode(response.content).decode("utf-8")31    else:32        return id33 34JSON_MULTIMODAL_KEY = "multimodal_data"35JSON_PROMPT_STRING_KEY = "prompt_string"36 37@pytest.fixture(autouse=True)38def create_server():39    global server40    os.environ['LLAMA_MEDIA_MARKER'] = '<__media__>'41    server = ServerPreset.tinygemma3()42 43def test_models_supports_multimodal_capability():44    global server45    server.start()46    res = server.make_request("GET", "/models", data={})47    assert res.status_code == 20048    model_info = res.body["models"][0]49    print(model_info)50    assert "completion" in model_info["capabilities"]51    assert "multimodal" in model_info["capabilities"]52 53def test_v1_models_supports_multimodal_capability():54    global server55    server.start()56    res = server.make_request("GET", "/v1/models", data={})57    assert res.status_code == 20058    model_info = res.body["models"][0]59    print(model_info)60    assert "completion" in model_info["capabilities"]61    assert "multimodal" in model_info["capabilities"]62 63@pytest.mark.parametrize(64    "prompt, image_url, success, re_content",65    [66        # test model is trained on CIFAR-10, but it's quite dumb due to small size67        ("What is this:\n", "IMG_URL_0",              True, "(cat)+"),68        ("What is this:\n", "IMG_BASE64_URI_0",       True, "(cat)+"),69        ("What is this:\n", "IMG_URL_1",              True, "(frog)+"),70        ("Test test\n",     "IMG_URL_1",              True, "(frog)+"), # test invalidate cache71        ("What is this:\n", "malformed",              False, None),72        ("What is this:\n", "https://google.com/404", False, None), # non-existent image73        ("What is this:\n", "https://ggml.ai",        False, None), # non-image data74        ("What is this:\n", "data:text/html;base64,aGVsbG8=", False, None), # unsupported data uri mime75        # TODO @ngxson : test with multiple images, no images and with audio76    ]77)78def test_vision_chat_completion(prompt, image_url, success, re_content):79    global server80    server.start()81    res = server.make_request("POST", "/chat/completions", data={82        "temperature": 0.0,83        "top_k": 1,84        "messages": [85            {"role": "user", "content": [86                {"type": "text", "text": prompt},87                {"type": "image_url", "image_url": {88                    "url": get_img_url(image_url),89                }},90            ]},91        ],92    })93    if success:94        assert res.status_code == 20095        choice = res.body["choices"][0]96        assert "assistant" == choice["message"]["role"]97        assert match_regex(re_content, choice["message"]["content"])98    else:99        assert res.status_code != 200100 101 102def test_vision_chat_completion_token_count():103    global server104    server.start()105    res = server.make_request("POST", "/chat/completions/input_tokens", data={106        "temperature": 0.0,107        "top_k": 1,108        "messages": [109            {"role": "user", "content": [110                {"type": "text", "text": "What is this:"},111                {"type": "image_url", "image_url": {112                    "url": get_img_url("IMG_URL_0"),113                }},114            ]},115        ],116    })117    assert res.status_code == 200118    assert res.body["input_tokens"] > 10119 120 121@pytest.mark.parametrize(122    "prompt, image_data, success, re_content",123    [124        # test model is trained on CIFAR-10, but it's quite dumb due to small size125        ("What is this: <__media__>\n", "IMG_BASE64_0",         True, "(cat)+|(automobile)+"),126        ("What is this: <__media__>\n", "IMG_BASE64_1",         True, "(frog)+"),127        ("What is this: <__media__>\n", "malformed",            False, None), # non-image data128        ("What is this:\n",             "",                     False, None), # empty string129    ]130)131def test_vision_completion(prompt, image_data, success, re_content):132    global server133    server.start()134    res = server.make_request("POST", "/completions", data={135        "temperature": 0.0,136        "top_k": 1,137        "prompt": {138            JSON_PROMPT_STRING_KEY: prompt,139            JSON_MULTIMODAL_KEY: [ get_img_url(image_data) ],140        },141    })142    if success:143        assert res.status_code == 200144        content = res.body["content"]145        assert match_regex(re_content, content)146    else:147        assert res.status_code != 200148 149 150@pytest.mark.parametrize(151    "prompt, image_data, success",152    [153        # test model is trained on CIFAR-10, but it's quite dumb due to small size154        ("What is this: <__media__>\n", "IMG_BASE64_0",         True),155        ("What is this: <__media__>\n", "IMG_BASE64_1",         True),156        ("What is this: <__media__>\n", "malformed",            False), # non-image data157        ("What is this:\n",             "base64",               False), # non-image data158    ]159)160def test_vision_embeddings(prompt, image_data, success):161    global server162    server.server_embeddings = True163    server.n_batch = 512164    server.start()165    image_data = get_img_url(image_data)166    res = server.make_request("POST", "/embeddings", data={167        "content": [168            { JSON_PROMPT_STRING_KEY: prompt, JSON_MULTIMODAL_KEY: [ image_data ] },169            { JSON_PROMPT_STRING_KEY: prompt, JSON_MULTIMODAL_KEY: [ image_data ] },170            { JSON_PROMPT_STRING_KEY: prompt, },171        ],172    })173    if success:174        assert res.status_code == 200175        content = res.body176        # Ensure embeddings are stable when multimodal.177        assert content[0]['embedding'] == content[1]['embedding']178        # Ensure embeddings without multimodal but same prompt do not match multimodal embeddings.179        assert content[0]['embedding'] != content[2]['embedding']180    else:181        assert res.status_code != 200182