Felipe97/llama-cpp-compiled
01.2k
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 