mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge pull request #2888 from rick-github/ollama-image-handling
Update ollama.py for image handling
This commit is contained in:
commit
77e9008eb0
2 changed files with 103 additions and 1 deletions
|
|
@ -121,6 +121,31 @@ class OllamaConfig:
|
|||
and v is not None
|
||||
}
|
||||
|
||||
# ollama wants plain base64 jpeg/png files as images. strip any leading dataURI
|
||||
# and convert to jpeg if necessary.
|
||||
def _convert_image(image):
|
||||
import base64, io
|
||||
try:
|
||||
from PIL import Image
|
||||
except:
|
||||
raise Exception(
|
||||
"ollama image conversion failed please run `pip install Pillow`"
|
||||
)
|
||||
|
||||
orig = image
|
||||
if image.startswith("data:"):
|
||||
image = image.split(",")[-1]
|
||||
try:
|
||||
image_data = Image.open(io.BytesIO(base64.b64decode(image)))
|
||||
if image_data.format in ["JPEG", "PNG"]:
|
||||
return image
|
||||
except:
|
||||
return orig
|
||||
jpeg_image = io.BytesIO()
|
||||
image_data.convert("RGB").save(jpeg_image, "JPEG")
|
||||
jpeg_image.seek(0)
|
||||
return base64.b64encode(jpeg_image.getvalue()).decode("utf-8")
|
||||
|
||||
|
||||
# ollama implementation
|
||||
def get_ollama_response(
|
||||
|
|
@ -158,7 +183,7 @@ def get_ollama_response(
|
|||
if format is not None:
|
||||
data["format"] = format
|
||||
if images is not None:
|
||||
data["images"] = images
|
||||
data["images"] = [_convert_image(image) for image in images]
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
|
|||
|
|
@ -1397,6 +1397,83 @@ def test_hf_classifier_task():
|
|||
pytest.fail(f"Error occurred: {str(e)}")
|
||||
|
||||
|
||||
def test_ollama_image():
|
||||
"""
|
||||
Test that datauri prefixes are removed, JPEG/PNG images are passed
|
||||
through, and other image formats are converted to JPEG. Non-image
|
||||
data is untouched.
|
||||
"""
|
||||
|
||||
import io, base64
|
||||
from PIL import Image
|
||||
|
||||
def mock_post(url, **kwargs):
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"Content-Type": "application/json"}
|
||||
mock_response.json.return_value = {
|
||||
# return the image in the response so that it can be tested
|
||||
# against the original
|
||||
"response": kwargs["json"]["images"]
|
||||
}
|
||||
return mock_response
|
||||
|
||||
def make_b64image(format):
|
||||
image = Image.new(mode='RGB', size=(1, 1))
|
||||
image_buffer = io.BytesIO()
|
||||
image.save(image_buffer, format)
|
||||
return base64.b64encode(image_buffer.getvalue()).decode("utf-8")
|
||||
|
||||
jpeg_image = make_b64image("JPEG")
|
||||
webp_image = make_b64image("WEBP")
|
||||
png_image = make_b64image("PNG")
|
||||
|
||||
base64_data = base64.b64encode(b"some random data")
|
||||
datauri_base64_data = f"data:text/plain;base64,{base64_data}"
|
||||
|
||||
tests = [
|
||||
# input expected
|
||||
[ jpeg_image, jpeg_image ],
|
||||
[ webp_image, None ],
|
||||
[ png_image, png_image ],
|
||||
[ f"data:image/jpeg;base64,{jpeg_image}", jpeg_image ],
|
||||
[ f"data:image/webp;base64,{webp_image}", None ],
|
||||
[ f"data:image/png;base64,{png_image}", png_image ],
|
||||
[ datauri_base64_data, datauri_base64_data ]
|
||||
]
|
||||
|
||||
for test in tests:
|
||||
try:
|
||||
with patch("requests.post", side_effect=mock_post):
|
||||
response = completion(
|
||||
model="ollama/llava",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Whats in this image?"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": test[0]
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
],
|
||||
)
|
||||
if not test[1]:
|
||||
# the conversion process may not always generate the same image,
|
||||
# so just check for a JPEG image when a conversion was done.
|
||||
image_data = response["choices"][0]["message"]["content"][0]
|
||||
image = Image.open(io.BytesIO(base64.b64decode(image_data)))
|
||||
assert image.format == "JPEG"
|
||||
else:
|
||||
assert response["choices"][0]["message"]["content"][0] == test[1]
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
||||
########################### End of Hugging Face Tests ##############################################
|
||||
# def test_completion_hf_api():
|
||||
# # failing on circle-ci commenting out
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue