diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index a88ba02a6b..4596e2b620 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -73,6 +73,25 @@ def ollama_pt(model, messages): # https://github.com/jmorganca/ollama/blob/af4cf final_prompt_value="### Response:", messages=messages ) + elif "llava" in model: + prompt = "" + images = [] + for message in messages: + if isinstance(message["content"], str): + prompt += message["content"] + elif isinstance(message["content"], list): + # see https://docs.litellm.ai/docs/providers/openai#openai-vision-models + for element in message["content"]: + if isinstance(element, dict): + if element["type"] == "text": + prompt += element["text"] + elif element["type"] == "image_url": + image_url = element["image_url"]["url"] + images.append(image_url) + return { + "prompt": prompt, + "images": images + } else: prompt = "".join(m["content"] if isinstance(m['content'], str) is str else "".join(m['content']) for m in messages) return prompt diff --git a/litellm/main.py b/litellm/main.py index 31613a67aa..0ad091bf73 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1308,7 +1308,14 @@ def completion( ) else: prompt = prompt_factory(model=model, messages=messages, custom_llm_provider=custom_llm_provider) - + if isinstance(prompt, dict): + # for multimode models - ollama/llava prompt_factory returns a dict { + # "prompt": prompt, + # "images": images + # } + prompt, images = prompt["prompt"], prompt["images"] + optional_params["images"] = images + ## LOGGING generator = ollama.get_ollama_response_stream(api_base, model, prompt, optional_params, logging_obj=logging, acompletion=acompletion, model_response=model_response, encoding=encoding) if acompletion is True: