diff --git a/docs/my-website/docs/image_generation.md b/docs/my-website/docs/image_generation.md index 46cf288f5a4..b16c60f55bd 100644 --- a/docs/my-website/docs/image_generation.md +++ b/docs/my-website/docs/image_generation.md @@ -154,7 +154,7 @@ Any non-openai params, will be treated as provider-specific params, and sent in ## OpenAI Image Generation Models ### Usage -```python +```python showLineNumbers from litellm import image_generation import os os.environ['OPENAI_API_KEY'] = "" @@ -171,7 +171,7 @@ response = image_generation(model='gpt-image-1', prompt="cute baby otter") ### API keys This can be set as env variables or passed as **params to litellm.image_generation()** -```python +```python showLineNumbers import os os.environ['AZURE_API_KEY'] = os.environ['AZURE_API_BASE'] = @@ -179,7 +179,7 @@ os.environ['AZURE_API_VERSION'] = ``` ### Usage -```python +```python showLineNumbers from litellm import embedding response = embedding( model="azure/", @@ -197,6 +197,15 @@ print(response) | dall-e-3 | `image_generation(model="azure/", prompt="cute baby otter")` | | dall-e-2 | `image_generation(model="azure/", prompt="cute baby otter")` | +## Xinference Image Generation Models + +Use this for Stable Diffusion models hosted on Xinference + +#### Usage + +See Xinference usage with LiteLLM [here](./providers/xinference.md#image-generation) + + ## OpenAI Compatible Image Generation Models Use this for calling `/image_generation` endpoints on OpenAI Compatible Servers, example https://github.com/xorbitsai/inference @@ -204,7 +213,7 @@ Use this for calling `/image_generation` endpoints on OpenAI Compatible Servers, **Note add `openai/` prefix to model so litellm knows to route to OpenAI** ### Usage -```python +```python showLineNumbers from litellm import image_generation response = image_generation( model = "openai/", # add `openai/` prefix to model so litellm knows to route to OpenAI @@ -218,7 +227,7 @@ Use this for stable diffusion on bedrock ### Usage -```python +```python showLineNumbers import os from litellm import image_generation @@ -239,7 +248,7 @@ print(f"response: {response}") Use this for image generation models on VertexAI -```python +```python showLineNumbers response = litellm.image_generation( prompt="An olympic size swimming pool", model="vertex_ai/imagegeneration@006", diff --git a/docs/my-website/docs/providers/xinference.md b/docs/my-website/docs/providers/xinference.md index 3686c02098a..9951a1ee3ab 100644 --- a/docs/my-website/docs/providers/xinference.md +++ b/docs/my-website/docs/providers/xinference.md @@ -1,6 +1,17 @@ # Xinference [Xorbits Inference] https://inference.readthedocs.io/en/latest/index.html +## Overview + +| Property | Details | +|-------|-------| +| Description | Xinference is an open-source platform to run inference with any open-source LLMs, image generation models, and more. | +| Provider Route on LiteLLM | `xinference/` | +| Link to Provider Doc | [Xinference ↗](https://inference.readthedocs.io/en/latest/index.html) | +| Supported Operations | [`/embeddings`](#sample-usage---embedding), [`/images/generations`](#image-generation) | + +LiteLLM supports Xinference Embedding + Image Generation calls. + ## API Base, Key ```python # env variable @@ -9,7 +20,7 @@ os.environ['XINFERENCE_API_KEY'] = "anything" #[optional] no api key required ``` ## Sample Usage - Embedding -```python +```python showLineNumbers from litellm import embedding import os @@ -22,7 +33,7 @@ print(response) ``` ## Sample Usage `api_base` param -```python +```python showLineNumbers from litellm import embedding import os @@ -34,6 +45,94 @@ response = embedding( print(response) ``` +## Image Generation + +### Usage - LiteLLM Python SDK + +```python showLineNumbers +from litellm import image_generation +import os + +# xinference image generation call +response = image_generation( + model="xinference/stabilityai/stable-diffusion-3.5-large", + prompt="A beautiful sunset over a calm ocean", + api_base="http://127.0.0.1:9997/v1", +) +print(response) +``` + +### Usage - LiteLLM Proxy Server + +#### 1. Setup config.yaml + +```yaml showLineNumbers +model_list: + - model_name: xinference-sd + litellm_params: + model: xinference/stabilityai/stable-diffusion-3.5-large + api_base: http://127.0.0.1:9997/v1 + api_key: anything + model_info: + mode: image_generation + +general_settings: + master_key: sk-1234 +``` + +#### 2. Start the proxy + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Test it + +```bash showLineNumbers +curl --location 'http://0.0.0.0:4000/v1/images/generations' \ +--header 'Content-Type: application/json' \ +--header 'Authorization: Bearer sk-1234' \ +--data '{ + "model": "xinference-sd", + "prompt": "A beautiful sunset over a calm ocean", + "n": 1, + "size": "1024x1024", + "response_format": "url" +}' +``` + +### Advanced Usage - With Additional Parameters + +```python showLineNumbers +from litellm import image_generation +import os + +os.environ['XINFERENCE_API_BASE'] = "http://127.0.0.1:9997/v1" + +response = image_generation( + model="xinference/stabilityai/stable-diffusion-3.5-large", + prompt="A beautiful sunset over a calm ocean", + n=1, # number of images + size="1024x1024", # image size + response_format="b64_json", # return format +) +print(response) +``` + +### Supported Image Generation Models + +Xinference supports various stable diffusion models. Here are some examples: + +| Model Name | Function Call | +|---------------------------------------------------------|----------------------------------------------------------------------------------------------------| +| stabilityai/stable-diffusion-3.5-large | `image_generation(model="xinference/stabilityai/stable-diffusion-3.5-large", prompt="...")` | +| stabilityai/stable-diffusion-xl-base-1.0 | `image_generation(model="xinference/stabilityai/stable-diffusion-xl-base-1.0", prompt="...")` | +| runwayml/stable-diffusion-v1-5 | `image_generation(model="xinference/runwayml/stable-diffusion-v1-5", prompt="...")` | + +For a complete list of supported image generation models, see: https://inference.readthedocs.io/en/latest/models/builtin/image/index.html + ## Supported Models All models listed here https://inference.readthedocs.io/en/latest/models/builtin/embedding/index.html are supported diff --git a/litellm/llms/xinference/image_generation/__init__.py b/litellm/llms/xinference/image_generation/__init__.py new file mode 100644 index 00000000000..bf2265693c6 --- /dev/null +++ b/litellm/llms/xinference/image_generation/__init__.py @@ -0,0 +1,13 @@ +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import XInferenceImageGenerationConfig + +__all__ = [ + "XInferenceImageGenerationConfig", +] + + +def get_xinference_image_generation_config(model: str) -> BaseImageGenerationConfig: + return XInferenceImageGenerationConfig() diff --git a/litellm/llms/xinference/image_generation/transformation.py b/litellm/llms/xinference/image_generation/transformation.py new file mode 100644 index 00000000000..6ff70d0642d --- /dev/null +++ b/litellm/llms/xinference/image_generation/transformation.py @@ -0,0 +1,40 @@ +from typing import List + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams + + +class XInferenceImageGenerationConfig(BaseImageGenerationConfig): + """ + XInference image generation config + + https://inference.readthedocs.io/en/v1.1.1/reference/generated/xinference.client.handlers.ImageModelHandle.text_to_image.html#xinference.client.handlers.ImageModelHandle.text_to_image + """ + + def get_supported_openai_params( + self, model: str + ) -> List[OpenAIImageGenerationOptionalParams]: + return ["n", "response_format", "size", "response_format"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + supported_params = self.get_supported_openai_params(model) + for k in non_default_params.keys(): + if k not in optional_params.keys(): + if k in supported_params: + optional_params[k] = non_default_params[k] + elif drop_params: + pass + else: + raise ValueError( + f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters." + ) + + return optional_params diff --git a/litellm/utils.py b/litellm/utils.py index 3ac19ca1a82..4117ba6fbb9 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7091,6 +7091,12 @@ class ProviderConfigManager: ) return get_azure_image_generation_config(model) + elif LlmProviders.XINFERENCE == provider: + from litellm.llms.xinference.image_generation import ( + get_xinference_image_generation_config, + ) + + return get_xinference_image_generation_config(model) return None @staticmethod diff --git a/tests/image_gen_tests/test_xinference.py b/tests/image_gen_tests/test_xinference.py new file mode 100644 index 00000000000..812f9c5d8c4 --- /dev/null +++ b/tests/image_gen_tests/test_xinference.py @@ -0,0 +1,153 @@ +import logging +import os +import sys +import traceback +import pytest +import json +from unittest.mock import Mock, patch, AsyncMock + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.types.utils import ImageObject + + +@pytest.mark.asyncio +async def test_xinference_image_generation(): + """Test basic xinference image generation with mocked OpenAI client.""" + + # Mock OpenAI response + mock_openai_response = { + "created": 1699623600, + "data": [ + { + "url": "https://example.com/image.png" + } + ] + } + + # Create a proper mock response object + class MockResponse: + def model_dump(self): + return mock_openai_response + + # Create a mock client with the images.generate method + mock_client = AsyncMock() + mock_client.images.generate = AsyncMock(return_value=MockResponse()) + + # Capture the actual arguments sent to OpenAI client + captured_args = None + captured_kwargs = None + + async def capture_generate_call(*args, **kwargs): + nonlocal captured_args, captured_kwargs + captured_args = args + captured_kwargs = kwargs + return MockResponse() + + mock_client.images.generate.side_effect = capture_generate_call + + # Mock the _get_openai_client method to return our mock client + with patch.object(litellm.main.openai_chat_completions, '_get_openai_client', return_value=mock_client): + response = await litellm.aimage_generation( + model="xinference/stabilityai/stable-diffusion-3.5-large", + prompt="A beautiful sunset over a calm ocean", + api_base="http://mock.image.generation.api", + ) + + # Print the captured arguments for debugging + print("Arguments sent to openai_aclient.images.generate:") + print("args:", json.dumps(captured_args, indent=4, default=str)) + print("kwargs:", json.dumps(captured_kwargs, indent=4, default=str)) + + # Validate the response + assert response is not None + assert response.created == 1699623600 + assert response.data is not None + assert len(response.data) == 1 + assert response.data[0].url == "https://example.com/image.png" + + # Validate that the OpenAI client was called with correct parameters + mock_client.images.generate.assert_called_once() + assert captured_kwargs is not None + assert captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" # xinference/ prefix removed + assert captured_kwargs["prompt"] == "A beautiful sunset over a calm ocean" + + +@pytest.mark.asyncio +async def test_xinference_image_generation_with_response_format(): + """ + Test xinference image generation with additional parameters. + Ensure all documented params are passed in. + + https://inference.readthedocs.io/en/v1.1.1/reference/generated/xinference.client.handlers.ImageModelHandle.text_to_image.html#xinference.client.handlers.ImageModelHandle.text_to_image + """ + + # Mock OpenAI response + mock_openai_response = { + "created": 1699623600, + "data": [ + { + "b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChAI9jU77yQAAAABJRU5ErkJggg==" + } + ] + } + + # Create a proper mock response object + class MockResponse: + def model_dump(self): + return mock_openai_response + + # Create a mock client with the images.generate method + mock_client = AsyncMock() + mock_client.images.generate = AsyncMock(return_value=MockResponse()) + + # Capture the actual arguments sent to OpenAI client + captured_args = None + captured_kwargs = None + + async def capture_generate_call(*args, **kwargs): + nonlocal captured_args, captured_kwargs + captured_args = args + captured_kwargs = kwargs + return MockResponse() + + mock_client.images.generate.side_effect = capture_generate_call + + # Mock the _get_openai_client method to return our mock client + with patch.object(litellm.main.openai_chat_completions, '_get_openai_client', return_value=mock_client): + response = await litellm.aimage_generation( + model="xinference/stabilityai/stable-diffusion-3.5-large", + api_base="http://mock.image.generation.api", + prompt="A beautiful sunset over a calm ocean", + response_format="b64_json", + n=1, + size="1024x1024", + ) + + # Print the captured arguments for debugging + print("Arguments sent to openai_aclient.images.generate:") + print("args:", json.dumps(captured_args, indent=4, default=str)) + print("kwargs:", json.dumps(captured_kwargs, indent=4, default=str)) + + # Validate the response + assert response is not None + assert response.created == 1699623600 + assert response.data is not None + assert len(response.data) == 1 + assert response.data[0].b64_json is not None + + # Validate that the OpenAI client was called with correct parameters + mock_client.images.generate.assert_called_once() + assert captured_kwargs is not None + assert captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large" # xinference/ prefix removed + assert captured_kwargs["prompt"] == "A beautiful sunset over a calm ocean" + assert captured_kwargs["response_format"] == "b64_json" + assert captured_kwargs["n"] == 1 + assert captured_kwargs["size"] == "1024x1024" + expected_args = ["model", "prompt", "response_format", "n", "size"] + # only expected args should be present + assert all(arg in captured_kwargs for arg in expected_args) +