mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[Feat] Add FAL AI Image Generations on LiteLLM (#16067)
* add fal-ai provider * fix image_generation_handler * init FalAIImageGenerationConfig * init cost_calculator * init FAL AI * TestFAL_AI_ImageGeneration * fix load_custom_provider_entrypoints * TestFAL_AI_ImageGeneration * add imagen4 transform FAL AI * add FAL AI imagen 4 transform * BaseImageGenTest * test_fal_ai_image_generation_basic * add BRIA + Recraft img gen * add recraft + BRIA * test_fal_ai_image_generation_basic * tests for flux PRO v11 * Add FAL AI SD * test FAL AI SD * docs FAL AI * docs fal ai * Using Model-Specific Parameters * add fal ai model prices * add fall_ai JPG logo * ui fixes FAL AI * fix linting * fix linting * fix bedrock test_get_request_body_stability3 * test_custom_llm_provider_entrypoint
This commit is contained in:
parent
4939793ade
commit
99feefd614
28 changed files with 2037 additions and 31 deletions
310
docs/my-website/docs/providers/fal_ai.md
Normal file
310
docs/my-website/docs/providers/fal_ai.md
Normal file
|
|
@ -0,0 +1,310 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Fal AI
|
||||
|
||||
Fal AI provides fast, scalable access to state-of-the-art image generation models including FLUX, Stable Diffusion, Imagen, and more.
|
||||
|
||||
## Overview
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Description | Fal AI offers optimized infrastructure for running image generation models at scale with low latency. |
|
||||
| Provider Route on LiteLLM | `fal_ai/` |
|
||||
| Provider Doc | [Fal AI Documentation ↗](https://fal.ai/models) |
|
||||
| Supported Operations | [`/images/generations`](#image-generation) |
|
||||
|
||||
## Setup
|
||||
|
||||
### API Key
|
||||
|
||||
```python showLineNumbers
|
||||
import os
|
||||
|
||||
# Set your Fal AI API key
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
```
|
||||
|
||||
Get your API key from [fal.ai](https://fal.ai/).
|
||||
|
||||
## Supported Models
|
||||
|
||||
| Model Name | Description | Documentation |
|
||||
|------------|-------------|---------------|
|
||||
| `fal_ai/fal-ai/flux-pro/v1.1-ultra` | FLUX Pro v1.1 Ultra - High-quality image generation | [Docs ↗](https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra) |
|
||||
| `fal_ai/fal-ai/imagen4/preview` | Google's Imagen 4 - Highest quality model | [Docs ↗](https://fal.ai/models/fal-ai/imagen4/preview) |
|
||||
| `fal_ai/fal-ai/recraft/v3/text-to-image` | Recraft v3 - Multiple style options | [Docs ↗](https://fal.ai/models/fal-ai/recraft/v3/text-to-image) |
|
||||
| `fal_ai/fal-ai/stable-diffusion-v35-medium` | Stable Diffusion v3.5 Medium | [Docs ↗](https://fal.ai/models/fal-ai/stable-diffusion-v35-medium) |
|
||||
| `fal_ai/bria/text-to-image/3.2` | Bria 3.2 - Commercial-grade generation | [Docs ↗](https://fal.ai/models/bria/text-to-image/3.2) |
|
||||
|
||||
## Image Generation
|
||||
|
||||
### Usage - LiteLLM Python SDK
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="basic" label="Basic Usage">
|
||||
|
||||
```python showLineNumbers title="Basic Image Generation"
|
||||
import litellm
|
||||
import os
|
||||
|
||||
# Set your API key
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
|
||||
# Generate an image
|
||||
response = litellm.image_generation(
|
||||
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
prompt="A serene mountain landscape at sunset with vibrant colors"
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="imagen4" label="Imagen 4">
|
||||
|
||||
```python showLineNumbers title="Google Imagen 4 Generation"
|
||||
import litellm
|
||||
import os
|
||||
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
|
||||
# Generate with Imagen 4
|
||||
response = litellm.image_generation(
|
||||
model="fal_ai/fal-ai/imagen4/preview",
|
||||
prompt="A vintage 1960s kitchen with flour package on countertop",
|
||||
aspect_ratio="16:9",
|
||||
num_images=1
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="recraft" label="Recraft v3">
|
||||
|
||||
```python showLineNumbers title="Recraft v3 with Style"
|
||||
import litellm
|
||||
import os
|
||||
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
|
||||
# Generate with specific style
|
||||
response = litellm.image_generation(
|
||||
model="fal_ai/fal-ai/recraft/v3/text-to-image",
|
||||
prompt="A red panda eating bamboo",
|
||||
style="realistic_image",
|
||||
image_size="landscape_4_3"
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="async" label="Async Usage">
|
||||
|
||||
```python showLineNumbers title="Async Image Generation"
|
||||
import litellm
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
async def generate_image():
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
|
||||
response = await litellm.aimage_generation(
|
||||
model="fal_ai/fal-ai/stable-diffusion-v35-medium",
|
||||
prompt="A cyberpunk cityscape with neon lights",
|
||||
guidance_scale=7.5,
|
||||
num_inference_steps=50
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
return response
|
||||
|
||||
asyncio.run(generate_image())
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="advanced" label="Advanced Parameters">
|
||||
|
||||
```python showLineNumbers title="Advanced FLUX Pro Generation"
|
||||
import litellm
|
||||
import os
|
||||
|
||||
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
|
||||
|
||||
# Generate with advanced parameters
|
||||
response = litellm.image_generation(
|
||||
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
prompt="A majestic dragon soaring over mountains",
|
||||
n=2,
|
||||
size="1792x1024", # Maps to aspect_ratio="16:9"
|
||||
seed=42,
|
||||
safety_tolerance="2",
|
||||
enhance_prompt=True
|
||||
)
|
||||
|
||||
for image in response.data:
|
||||
print(f"Generated image: {image.url}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Usage - LiteLLM Proxy Server
|
||||
|
||||
#### 1. Configure your config.yaml
|
||||
|
||||
```yaml showLineNumbers title="Fal AI Image Generation Configuration"
|
||||
model_list:
|
||||
- model_name: flux-ultra
|
||||
litellm_params:
|
||||
model: fal_ai/fal-ai/flux-pro/v1.1-ultra
|
||||
api_key: os.environ/FAL_AI_API_KEY
|
||||
model_info:
|
||||
mode: image_generation
|
||||
|
||||
- model_name: imagen4
|
||||
litellm_params:
|
||||
model: fal_ai/fal-ai/imagen4/preview
|
||||
api_key: os.environ/FAL_AI_API_KEY
|
||||
model_info:
|
||||
mode: image_generation
|
||||
|
||||
- model_name: stable-diffusion
|
||||
litellm_params:
|
||||
model: fal_ai/fal-ai/stable-diffusion-v35-medium
|
||||
api_key: os.environ/FAL_AI_API_KEY
|
||||
model_info:
|
||||
mode: image_generation
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
```
|
||||
|
||||
#### 2. Start LiteLLM Proxy Server
|
||||
|
||||
```bash showLineNumbers title="Start Proxy Server"
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
#### 3. Make requests
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="openai-sdk" label="OpenAI SDK">
|
||||
|
||||
```python showLineNumbers title="Generate via Proxy - OpenAI SDK"
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
base_url="http://localhost:4000",
|
||||
api_key="sk-1234"
|
||||
)
|
||||
|
||||
response = client.images.generate(
|
||||
model="flux-ultra",
|
||||
prompt="A beautiful sunset over the ocean",
|
||||
n=1,
|
||||
size="1024x1024"
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="litellm-sdk" label="LiteLLM SDK">
|
||||
|
||||
```python showLineNumbers title="Generate via Proxy - LiteLLM SDK"
|
||||
import litellm
|
||||
|
||||
response = litellm.image_generation(
|
||||
model="litellm_proxy/imagen4",
|
||||
prompt="A cozy coffee shop interior",
|
||||
api_base="http://localhost:4000",
|
||||
api_key="sk-1234"
|
||||
)
|
||||
|
||||
print(response.data[0].url)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```bash showLineNumbers title="Generate via Proxy - cURL"
|
||||
curl --location 'http://localhost:4000/v1/images/generations' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--data '{
|
||||
"model": "stable-diffusion",
|
||||
"prompt": "A serene Japanese garden with cherry blossoms",
|
||||
"n": 1,
|
||||
"size": "1024x1024"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
|
||||
## Using Model-Specific Parameters
|
||||
|
||||
LiteLLM forwards any additional parameters directly to the Fal AI API. You can pass model-specific parameters in your request and they will be sent to Fal AI.
|
||||
|
||||
```python showLineNumbers title="Pass Model-Specific Parameters"
|
||||
import litellm
|
||||
|
||||
# Any parameters beyond the standard ones are forwarded to Fal AI
|
||||
response = litellm.image_generation(
|
||||
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
prompt="A beautiful sunset",
|
||||
# Model-specific Fal AI parameters
|
||||
aspect_ratio="16:9",
|
||||
safety_tolerance="2",
|
||||
enhance_prompt=True,
|
||||
seed=42
|
||||
)
|
||||
```
|
||||
|
||||
For the complete list of parameters supported by each model, see:
|
||||
- [FLUX Pro v1.1-ultra Parameters ↗](https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra/api)
|
||||
- [Imagen 4 Parameters ↗](https://fal.ai/models/fal-ai/imagen4/preview/api)
|
||||
- [Recraft v3 Parameters ↗](https://fal.ai/models/fal-ai/recraft/v3/text-to-image/api)
|
||||
- [Stable Diffusion v3.5 Parameters ↗](https://fal.ai/models/fal-ai/stable-diffusion-v35-medium/api)
|
||||
- [Bria 3.2 Parameters ↗](https://fal.ai/models/bria/text-to-image/3.2/api)
|
||||
|
||||
## Supported Parameters
|
||||
|
||||
Standard OpenAI-compatible parameters that work across all models:
|
||||
|
||||
| Parameter | Type | Description | Default |
|
||||
|-----------|------|-------------|---------|
|
||||
| `prompt` | string | Text description of desired image | Required |
|
||||
| `model` | string | Fal AI model to use | Required |
|
||||
| `n` | integer | Number of images to generate (1-4) | `1` |
|
||||
| `size` | string | Image dimensions (maps to model-specific format) | Model default |
|
||||
| `api_key` | string | Your Fal AI API key | Environment variable |
|
||||
|
||||
## Getting Started
|
||||
|
||||
1. Sign up at [fal.ai](https://fal.ai/)
|
||||
2. Get your API key from your account settings
|
||||
3. Set `FAL_AI_API_KEY` environment variable
|
||||
4. Choose a model from the [Fal AI model gallery](https://fal.ai/models)
|
||||
5. Start generating images with LiteLLM
|
||||
|
||||
## Additional Resources
|
||||
|
||||
- [Fal AI Documentation](https://fal.ai/docs)
|
||||
- [Model Gallery](https://fal.ai/models)
|
||||
- [API Reference](https://fal.ai/docs/api-reference)
|
||||
- [Pricing](https://fal.ai/pricing)
|
||||
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# 🆕 Github
|
||||
# Github
|
||||
https://github.com/marketplace/models
|
||||
|
||||
:::tip
|
||||
|
|
|
|||
|
|
@ -537,6 +537,7 @@ const sidebars = {
|
|||
"providers/groq",
|
||||
"providers/deepseek",
|
||||
"providers/elevenlabs",
|
||||
"providers/fal_ai",
|
||||
"providers/fireworks_ai",
|
||||
"providers/clarifai",
|
||||
"providers/compactifai",
|
||||
|
|
|
|||
|
|
@ -342,6 +342,7 @@ def image_generation( # noqa: PLR0915
|
|||
litellm.LlmProviders.RECRAFT,
|
||||
litellm.LlmProviders.AIML,
|
||||
litellm.LlmProviders.GEMINI,
|
||||
litellm.LlmProviders.FAL_AI,
|
||||
):
|
||||
if image_generation_config is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -644,8 +644,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
ctype = content.get("type")
|
||||
return not (isinstance(ctype, str) and ctype != "text")
|
||||
|
||||
def _process_messages(self, messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> List[dict[str, Any]]:
|
||||
filtered_messages: List[dict[str, Any]] = []
|
||||
def _process_messages(self, messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> List[Dict[str, Any]]:
|
||||
filtered_messages: List[Dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@ import litellm.types.llms
|
|||
from litellm import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
|
||||
from litellm.litellm_core_utils.token_counter import get_image_type
|
||||
from litellm.types.files import get_file_extension_from_mime_type
|
||||
from litellm.types.llms.anthropic import *
|
||||
from litellm.types.llms.bedrock import CachePointBlock
|
||||
|
|
|
|||
24
litellm/llms/fal_ai/__init__.py
Normal file
24
litellm/llms/fal_ai/__init__.py
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
from .cost_calculator import cost_calculator
|
||||
from .image_generation import (
|
||||
FalAIBaseConfig,
|
||||
FalAIBriaConfig,
|
||||
FalAIFluxProV11UltraConfig,
|
||||
FalAIImageGenerationConfig,
|
||||
FalAIImagen4Config,
|
||||
FalAIRecraftV3Config,
|
||||
FalAIStableDiffusionConfig,
|
||||
get_fal_ai_image_generation_config,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"cost_calculator",
|
||||
"FalAIBaseConfig",
|
||||
"FalAIImageGenerationConfig",
|
||||
"FalAIImagen4Config",
|
||||
"FalAIRecraftV3Config",
|
||||
"FalAIBriaConfig",
|
||||
"FalAIFluxProV11UltraConfig",
|
||||
"FalAIStableDiffusionConfig",
|
||||
"get_fal_ai_image_generation_config",
|
||||
]
|
||||
|
||||
26
litellm/llms/fal_ai/cost_calculator.py
Normal file
26
litellm/llms/fal_ai/cost_calculator.py
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
def cost_calculator(
|
||||
model: str,
|
||||
image_response: Any,
|
||||
) -> float:
|
||||
"""
|
||||
fal.ai image generation cost calculator
|
||||
"""
|
||||
_model_info = litellm.get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider=litellm.LlmProviders.FAL_AI.value,
|
||||
)
|
||||
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
|
||||
num_images: int = 0
|
||||
if isinstance(image_response, ImageResponse):
|
||||
if image_response.data:
|
||||
num_images = len(image_response.data)
|
||||
return output_cost_per_image * num_images
|
||||
else:
|
||||
raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}")
|
||||
|
||||
49
litellm/llms/fal_ai/image_generation/__init__.py
Normal file
49
litellm/llms/fal_ai/image_generation/__init__.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
|
||||
from .bria_transformation import FalAIBriaConfig
|
||||
from .flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig
|
||||
from .imagen4_transformation import FalAIImagen4Config
|
||||
from .recraft_v3_transformation import FalAIRecraftV3Config
|
||||
from .stable_diffusion_transformation import FalAIStableDiffusionConfig
|
||||
from .transformation import FalAIBaseConfig, FalAIImageGenerationConfig
|
||||
|
||||
__all__ = [
|
||||
"FalAIBaseConfig",
|
||||
"FalAIImageGenerationConfig",
|
||||
"FalAIImagen4Config",
|
||||
"FalAIRecraftV3Config",
|
||||
"FalAIBriaConfig",
|
||||
"FalAIFluxProV11UltraConfig",
|
||||
"FalAIStableDiffusionConfig",
|
||||
]
|
||||
|
||||
|
||||
def get_fal_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
||||
"""
|
||||
Get the appropriate Fal AI image generation configuration based on the model.
|
||||
|
||||
Args:
|
||||
model: The Fal AI model name (e.g., "fal-ai/imagen4/preview", "fal-ai/recraft/v3/text-to-image")
|
||||
|
||||
Returns:
|
||||
The appropriate configuration class for the specified model
|
||||
"""
|
||||
model_lower = model.lower()
|
||||
|
||||
# Map model names to their corresponding configuration classes
|
||||
if "imagen4" in model_lower or "imagen-4" in model_lower:
|
||||
return FalAIImagen4Config()
|
||||
elif "recraft" in model_lower:
|
||||
return FalAIRecraftV3Config()
|
||||
elif "bria" in model_lower:
|
||||
return FalAIBriaConfig()
|
||||
elif "flux-pro" in model_lower and "ultra" in model_lower:
|
||||
return FalAIFluxProV11UltraConfig()
|
||||
elif "stable-diffusion" in model_lower:
|
||||
return FalAIStableDiffusionConfig()
|
||||
|
||||
# Default to generic Fal AI configuration
|
||||
return FalAIImageGenerationConfig()
|
||||
|
||||
231
litellm/llms/fal_ai/image_generation/bria_transformation.py
Normal file
231
litellm/llms/fal_ai/image_generation/bria_transformation.py
Normal file
|
|
@ -0,0 +1,231 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIBriaConfig(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Bria Text-to-Image 3.2 model.
|
||||
|
||||
Bria 3.2 is a commercial-grade text-to-image model with prompt enhancement
|
||||
and multiple aspect ratio options.
|
||||
|
||||
Model endpoint: bria/text-to-image/3.2
|
||||
Documentation: https://fal.ai/models/bria/text-to-image/3.2
|
||||
"""
|
||||
IMAGE_GENERATION_ENDPOINT: str = "bria/text-to-image/3.2"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for Bria 3.2.
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to Bria 3.2 parameters.
|
||||
|
||||
Mappings:
|
||||
- size -> aspect_ratio (1:1, 2:3, 3:2, 3:4, 4:3, 4:5, 5:4, 9:16, 16:9)
|
||||
- response_format -> ignored (Bria returns URLs)
|
||||
- n -> ignored (Bria doesn't support multiple images in one call)
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Map OpenAI params to Bria params
|
||||
param_mapping = {
|
||||
"size": "aspect_ratio",
|
||||
}
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# Use mapped parameter name if exists
|
||||
mapped_key = param_mapping.get(k, k)
|
||||
mapped_value = non_default_params[k]
|
||||
|
||||
# Transform specific parameters
|
||||
if k == "response_format":
|
||||
# Bria always returns URLs, so we can ignore this
|
||||
continue
|
||||
elif k == "n":
|
||||
# Bria doesn't support multiple images, ignore
|
||||
continue
|
||||
elif k == "size":
|
||||
# Map OpenAI size format to Bria aspect ratio
|
||||
mapped_value = self._map_aspect_ratio(mapped_value)
|
||||
|
||||
optional_params[mapped_key] = mapped_value
|
||||
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
|
||||
|
||||
def _map_aspect_ratio(self, size: str) -> str:
|
||||
"""
|
||||
Map OpenAI size format to Bria aspect ratio format.
|
||||
|
||||
OpenAI format: "1024x1024", "1792x1024", etc.
|
||||
Bria format: "1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"
|
||||
"""
|
||||
# Map common OpenAI sizes to Bria aspect ratios
|
||||
size_to_aspect_ratio = {
|
||||
"1024x1024": "1:1",
|
||||
"512x512": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1024x768": "4:3",
|
||||
"768x1024": "3:4",
|
||||
"1280x960": "4:3",
|
||||
"960x1280": "3:4",
|
||||
}
|
||||
|
||||
if size in size_to_aspect_ratio:
|
||||
return size_to_aspect_ratio[size]
|
||||
|
||||
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
|
||||
if "x" in size:
|
||||
try:
|
||||
width_str, height_str = size.split("x")
|
||||
width = int(width_str)
|
||||
height = int(height_str)
|
||||
|
||||
# Calculate aspect ratio and find closest match
|
||||
ratio = width / height
|
||||
|
||||
# Map to closest supported aspect ratio
|
||||
if 0.95 <= ratio <= 1.05: # Close to 1:1
|
||||
return "1:1"
|
||||
elif ratio >= 1.7: # Close to 16:9
|
||||
return "16:9"
|
||||
elif ratio <= 0.6: # Close to 9:16
|
||||
return "9:16"
|
||||
elif 1.3 <= ratio <= 1.4: # Close to 4:3
|
||||
return "4:3"
|
||||
elif 0.7 <= ratio <= 0.8: # Close to 3:4
|
||||
return "3:4"
|
||||
elif 1.45 <= ratio <= 1.55: # Close to 3:2
|
||||
return "3:2"
|
||||
elif 0.65 <= ratio <= 0.7: # Close to 2:3
|
||||
return "2:3"
|
||||
elif 1.2 <= ratio <= 1.3: # Close to 5:4
|
||||
return "5:4"
|
||||
elif 0.75 <= ratio <= 0.85: # Close to 4:5
|
||||
return "4:5"
|
||||
except (ValueError, AttributeError, ZeroDivisionError):
|
||||
pass
|
||||
|
||||
# Default to 1:1
|
||||
return "1:1"
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to Bria 3.2 request body.
|
||||
|
||||
Required parameters:
|
||||
- prompt: Prompt for image generation
|
||||
|
||||
Optional parameters:
|
||||
- aspect_ratio: "1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9" (default: "1:1")
|
||||
- prompt_enhancer: Improve the prompt (default: true)
|
||||
- sync_mode: Return image directly in response (default: false)
|
||||
- truncate_prompt: Truncate the prompt (default: true)
|
||||
- guidance_scale: Guidance scale 1-10 (default: 5)
|
||||
- num_inference_steps: Inference steps 20-50 (default: 30)
|
||||
- seed: Random seed for reproducibility (default: 5555)
|
||||
- negative_prompt: Negative prompt string
|
||||
"""
|
||||
bria_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
return bria_request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the Bria 3.2 response to litellm ImageResponse format.
|
||||
|
||||
Expected response format:
|
||||
{
|
||||
"image": {
|
||||
"url": "https://...",
|
||||
"content_type": "image/png",
|
||||
"file_name": "...",
|
||||
"file_size": 123456,
|
||||
"width": 1024,
|
||||
"height": 1024
|
||||
}
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle Bria response format - uses "image" (singular) not "images"
|
||||
image_data = response_data.get("image")
|
||||
if image_data and isinstance(image_data, dict):
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=None, # Bria returns URLs only
|
||||
)
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -0,0 +1,263 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIFluxProV11UltraConfig(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Fal AI Flux Pro v1.1-ultra model.
|
||||
|
||||
FLUX Pro v1.1-ultra is a high-quality text-to-image model with enhanced detail
|
||||
and support for image prompts.
|
||||
|
||||
Model endpoint: fal-ai/flux-pro/v1.1-ultra
|
||||
Documentation: https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra
|
||||
"""
|
||||
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/flux-pro/v1.1-ultra"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for Flux Pro v1.1-ultra.
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to Flux Pro v1.1-ultra parameters.
|
||||
|
||||
Mappings:
|
||||
- n -> num_images (1-4, default 1)
|
||||
- response_format -> output_format (jpeg or png)
|
||||
- size -> aspect_ratio (21:9, 16:9, 4:3, 3:2, 1:1, 2:3, 3:4, 9:16, 9:21)
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Map OpenAI params to Flux Pro v1.1-ultra params
|
||||
param_mapping = {
|
||||
"n": "num_images",
|
||||
"response_format": "output_format",
|
||||
"size": "aspect_ratio",
|
||||
}
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# Use mapped parameter name if exists
|
||||
mapped_key = param_mapping.get(k, k)
|
||||
mapped_value = non_default_params[k]
|
||||
|
||||
# Transform specific parameters
|
||||
if k == "response_format":
|
||||
# Map OpenAI response formats to image formats
|
||||
if mapped_value in ["b64_json", "url"]:
|
||||
mapped_value = "jpeg"
|
||||
elif k == "size":
|
||||
# Map OpenAI size format to Flux aspect ratio
|
||||
mapped_value = self._map_aspect_ratio(mapped_value)
|
||||
|
||||
optional_params[mapped_key] = mapped_value
|
||||
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
|
||||
|
||||
def _map_aspect_ratio(self, size: str) -> str:
|
||||
"""
|
||||
Map OpenAI size format to Flux Pro aspect ratio format.
|
||||
|
||||
OpenAI format: "1024x1024", "1792x1024", etc.
|
||||
Flux format: "21:9", "16:9", "4:3", "3:2", "1:1", "2:3", "3:4", "9:16", "9:21"
|
||||
|
||||
Default: "16:9"
|
||||
"""
|
||||
# Map common OpenAI sizes to Flux aspect ratios
|
||||
size_to_aspect_ratio = {
|
||||
"1024x1024": "1:1",
|
||||
"512x512": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1024x768": "4:3",
|
||||
"768x1024": "3:4",
|
||||
"1536x1024": "3:2",
|
||||
"1024x1536": "2:3",
|
||||
"2048x876": "21:9",
|
||||
"876x2048": "9:21",
|
||||
}
|
||||
|
||||
if size in size_to_aspect_ratio:
|
||||
return size_to_aspect_ratio[size]
|
||||
|
||||
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
|
||||
if "x" in size:
|
||||
try:
|
||||
width_str, height_str = size.split("x")
|
||||
width = int(width_str)
|
||||
height = int(height_str)
|
||||
|
||||
# Calculate aspect ratio and find closest match
|
||||
ratio = width / height
|
||||
|
||||
# Map to closest supported aspect ratio
|
||||
if 0.95 <= ratio <= 1.05: # Close to 1:1
|
||||
return "1:1"
|
||||
elif ratio >= 2.3: # Close to 21:9
|
||||
return "21:9"
|
||||
elif 1.7 <= ratio < 2.3: # Close to 16:9
|
||||
return "16:9"
|
||||
elif 1.3 <= ratio < 1.7: # Close to 4:3
|
||||
return "4:3"
|
||||
elif 1.4 <= ratio < 1.6: # Close to 3:2
|
||||
return "3:2"
|
||||
elif 0.6 <= ratio < 0.7: # Close to 3:4
|
||||
return "3:4"
|
||||
elif 0.65 <= ratio < 0.75: # Close to 2:3
|
||||
return "2:3"
|
||||
elif 0.5 <= ratio < 0.6: # Close to 9:16
|
||||
return "9:16"
|
||||
elif ratio < 0.5: # Close to 9:21
|
||||
return "9:21"
|
||||
except (ValueError, AttributeError, ZeroDivisionError):
|
||||
pass
|
||||
|
||||
# Default to 16:9
|
||||
return "16:9"
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to Flux Pro v1.1-ultra request body.
|
||||
|
||||
Required parameters:
|
||||
- prompt: The prompt to generate an image from
|
||||
|
||||
Optional parameters:
|
||||
- num_images: Number of images (1-4, default: 1)
|
||||
- aspect_ratio: Aspect ratio (default: "16:9")
|
||||
- raw: Generate less processed images (default: false)
|
||||
- output_format: "jpeg" or "png" (default: "jpeg")
|
||||
- image_url: Image URL for image-to-image generation
|
||||
- sync_mode: Return data URI (default: false)
|
||||
- safety_tolerance: Safety level "1"-"6" (default: "2")
|
||||
- enable_safety_checker: Enable safety checker (default: true)
|
||||
- seed: Random seed for reproducibility
|
||||
- image_prompt_strength: Strength of image prompt 0-1 (default: 0.1)
|
||||
- enhance_prompt: Enhance prompt for better results (default: false)
|
||||
"""
|
||||
flux_pro_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
return flux_pro_request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the Flux Pro v1.1-ultra response to litellm ImageResponse format.
|
||||
|
||||
Expected response format:
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"url": "https://...",
|
||||
"width": 1024,
|
||||
"height": 768,
|
||||
"content_type": "image/jpeg"
|
||||
}
|
||||
],
|
||||
"timings": {"inference": 2.5, ...},
|
||||
"seed": 42,
|
||||
"has_nsfw_concepts": [false],
|
||||
"prompt": "original prompt"
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle Flux Pro v1.1-ultra response format
|
||||
images = response_data.get("images", [])
|
||||
if isinstance(images, list):
|
||||
for image_data in images:
|
||||
if isinstance(image_data, dict):
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=None, # Flux Pro returns URLs only
|
||||
)
|
||||
)
|
||||
elif isinstance(image_data, str):
|
||||
# If images is just a list of URLs
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data,
|
||||
b64_json=None,
|
||||
)
|
||||
)
|
||||
|
||||
# Add additional metadata from Flux Pro response
|
||||
if hasattr(model_response, "_hidden_params"):
|
||||
if "seed" in response_data:
|
||||
model_response._hidden_params["seed"] = response_data["seed"]
|
||||
if "timings" in response_data:
|
||||
model_response._hidden_params["timings"] = response_data["timings"]
|
||||
if "has_nsfw_concepts" in response_data:
|
||||
model_response._hidden_params["has_nsfw_concepts"] = response_data[
|
||||
"has_nsfw_concepts"
|
||||
]
|
||||
|
||||
return model_response
|
||||
|
||||
242
litellm/llms/fal_ai/image_generation/imagen4_transformation.py
Normal file
242
litellm/llms/fal_ai/image_generation/imagen4_transformation.py
Normal file
|
|
@ -0,0 +1,242 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIImagen4Config(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Fal AI Imagen4 model.
|
||||
|
||||
Google's highest quality image generation model available through Fal AI.
|
||||
|
||||
Model variants:
|
||||
- fal-ai/imagen4/preview (Standard): $0.05 per image
|
||||
- fal-ai/imagen4/preview/fast (Fast): $0.04 per image
|
||||
- fal-ai/imagen4/preview/ultra (Ultra): $0.06 per image
|
||||
|
||||
Documentation: https://fal.ai/models/fal-ai/imagen4/preview
|
||||
"""
|
||||
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/imagen4/preview"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for Imagen4.
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to Imagen4 parameters.
|
||||
|
||||
Mappings:
|
||||
- n -> num_images (1-4, default 1)
|
||||
- size -> aspect_ratio (1:1, 16:9, 9:16, 3:4, 4:3)
|
||||
- response_format -> ignored (Imagen4 returns URLs)
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Map OpenAI params to Imagen4 params
|
||||
param_mapping = {
|
||||
"n": "num_images",
|
||||
"size": "aspect_ratio",
|
||||
}
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# Use mapped parameter name if exists
|
||||
mapped_key = param_mapping.get(k, k)
|
||||
mapped_value = non_default_params[k]
|
||||
|
||||
# Transform specific parameters
|
||||
if k == "response_format":
|
||||
# Imagen4 always returns URLs, so we can ignore this
|
||||
continue
|
||||
elif k == "size":
|
||||
# Map OpenAI size format to Imagen4 aspect ratio
|
||||
mapped_value = self._map_aspect_ratio(mapped_value)
|
||||
|
||||
optional_params[mapped_key] = mapped_value
|
||||
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
|
||||
|
||||
def _map_aspect_ratio(self, size: str) -> str:
|
||||
"""
|
||||
Map OpenAI size format to Imagen4 aspect ratio format.
|
||||
|
||||
OpenAI format: "1024x1024", "1792x1024", etc.
|
||||
Imagen4 format: "1:1", "16:9", "9:16", "3:4", "4:3"
|
||||
|
||||
Available aspect ratios:
|
||||
- 1:1 (default)
|
||||
- 16:9
|
||||
- 9:16
|
||||
- 3:4
|
||||
- 4:3
|
||||
"""
|
||||
# Map common OpenAI sizes to Imagen4 aspect ratios
|
||||
size_to_aspect_ratio = {
|
||||
"1024x1024": "1:1",
|
||||
"512x512": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1024x768": "4:3",
|
||||
"768x1024": "3:4",
|
||||
}
|
||||
|
||||
if size in size_to_aspect_ratio:
|
||||
return size_to_aspect_ratio[size]
|
||||
|
||||
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
|
||||
if "x" in size:
|
||||
try:
|
||||
width_str, height_str = size.split("x")
|
||||
width = int(width_str)
|
||||
height = int(height_str)
|
||||
|
||||
# Calculate aspect ratio and find closest match
|
||||
ratio = width / height
|
||||
|
||||
# Map to closest supported aspect ratio
|
||||
if 0.95 <= ratio <= 1.05: # Close to 1:1
|
||||
return "1:1"
|
||||
elif ratio >= 1.7: # Close to 16:9
|
||||
return "16:9"
|
||||
elif ratio <= 0.6: # Close to 9:16
|
||||
return "9:16"
|
||||
elif ratio >= 1.2: # Close to 4:3
|
||||
return "4:3"
|
||||
elif ratio <= 0.8: # Close to 3:4
|
||||
return "3:4"
|
||||
except (ValueError, AttributeError, ZeroDivisionError):
|
||||
pass
|
||||
|
||||
# Default to 1:1
|
||||
return "1:1"
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to Imagen4 request body.
|
||||
|
||||
Required parameters:
|
||||
- prompt: The text prompt describing what you want to see
|
||||
|
||||
Optional parameters:
|
||||
- aspect_ratio: "1:1", "16:9", "9:16", "3:4", "4:3" (default: "1:1")
|
||||
- num_images: Number of images (1-4, default: 1)
|
||||
- resolution: "1K" or "2K" (default: "1K")
|
||||
- seed: Random seed for reproducibility
|
||||
- negative_prompt: Description of what to discourage (default: "")
|
||||
"""
|
||||
imagen4_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
return imagen4_request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the Imagen4 response to litellm ImageResponse format.
|
||||
|
||||
Expected response format:
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"url": "https://...",
|
||||
"content_type": "image/png",
|
||||
"file_name": "z9RV14K95DvU.png",
|
||||
"file_size": 4404019
|
||||
}
|
||||
],
|
||||
"seed": 42
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle Imagen4 response format
|
||||
images = response_data.get("images", [])
|
||||
if isinstance(images, list):
|
||||
for image_data in images:
|
||||
if isinstance(image_data, dict):
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=None, # Imagen4 returns URLs only
|
||||
)
|
||||
)
|
||||
elif isinstance(image_data, str):
|
||||
# If images is just a list of URLs
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data,
|
||||
b64_json=None,
|
||||
)
|
||||
)
|
||||
|
||||
# Add seed metadata from Imagen4 response
|
||||
if hasattr(model_response, "_hidden_params"):
|
||||
if "seed" in response_data:
|
||||
model_response._hidden_params["seed"] = response_data["seed"]
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -0,0 +1,226 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIRecraftV3Config(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Fal AI Recraft v3 Text-to-Image model.
|
||||
|
||||
Recraft v3 is a text-to-image model with multiple style options including
|
||||
realistic images, digital illustrations, and vector illustrations.
|
||||
|
||||
Model endpoint: fal-ai/recraft/v3/text-to-image
|
||||
Documentation: https://fal.ai/models/fal-ai/recraft/v3/text-to-image
|
||||
"""
|
||||
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/recraft/v3/text-to-image"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for Recraft v3.
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to Recraft v3 parameters.
|
||||
|
||||
Mappings:
|
||||
- size -> image_size (can be preset or custom width/height)
|
||||
- response_format -> ignored (Recraft returns URLs)
|
||||
- n -> ignored (Recraft doesn't support multiple images)
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Map OpenAI params to Recraft v3 params
|
||||
param_mapping = {
|
||||
"size": "image_size",
|
||||
}
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# Use mapped parameter name if exists
|
||||
mapped_key = param_mapping.get(k, k)
|
||||
mapped_value = non_default_params[k]
|
||||
|
||||
# Transform specific parameters
|
||||
if k == "response_format":
|
||||
# Recraft always returns URLs, so we can ignore this
|
||||
continue
|
||||
elif k == "n":
|
||||
# Recraft doesn't support multiple images, ignore
|
||||
continue
|
||||
elif k == "size":
|
||||
# Map OpenAI size format to Recraft image_size
|
||||
mapped_value = self._map_image_size(mapped_value)
|
||||
|
||||
optional_params[mapped_key] = mapped_value
|
||||
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
|
||||
|
||||
def _map_image_size(self, size: str) -> Any:
|
||||
"""
|
||||
Map OpenAI size format to Recraft v3 image_size format.
|
||||
|
||||
OpenAI format: "1024x1024", "1792x1024", etc.
|
||||
Recraft format: Can be preset strings or {"width": int, "height": int}
|
||||
|
||||
Available presets:
|
||||
- square_hd (default)
|
||||
- square
|
||||
- portrait_4_3
|
||||
- portrait_16_9
|
||||
- landscape_4_3
|
||||
- landscape_16_9
|
||||
"""
|
||||
# Map common OpenAI sizes to Recraft presets
|
||||
size_mapping = {
|
||||
"1024x1024": "square_hd",
|
||||
"512x512": "square",
|
||||
"768x1024": "portrait_4_3",
|
||||
"576x1024": "portrait_16_9",
|
||||
"1024x768": "landscape_4_3",
|
||||
"1024x576": "landscape_16_9",
|
||||
}
|
||||
|
||||
if size in size_mapping:
|
||||
return size_mapping[size]
|
||||
|
||||
# Parse custom size format "WIDTHxHEIGHT"
|
||||
if "x" in size:
|
||||
try:
|
||||
width, height = size.split("x")
|
||||
return {
|
||||
"width": int(width),
|
||||
"height": int(height),
|
||||
}
|
||||
except (ValueError, AttributeError):
|
||||
pass
|
||||
|
||||
# Default to square_hd
|
||||
return "square_hd"
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to Recraft v3 request body.
|
||||
|
||||
Required parameters:
|
||||
- prompt: Text prompt (max 1000 characters)
|
||||
|
||||
Optional parameters:
|
||||
- image_size: Preset or {"width": int, "height": int} (default: "square_hd")
|
||||
- style: Style preset (default: "realistic_image")
|
||||
Options: "any", "realistic_image", "digital_illustration", "vector_illustration", etc.
|
||||
- colors: Array of RGB color objects [{"r": 0-255, "g": 0-255, "b": 0-255}]
|
||||
- enable_safety_checker: Enable safety checker (default: false)
|
||||
- style_id: UUID for custom style reference
|
||||
|
||||
Note: Vector illustrations cost 2X as much.
|
||||
"""
|
||||
recraft_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
return recraft_request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the Recraft v3 response to litellm ImageResponse format.
|
||||
|
||||
Expected response format:
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"url": "https://...",
|
||||
"content_type": "image/webp",
|
||||
"file_name": "...",
|
||||
"file_size": 123456
|
||||
}
|
||||
]
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle Recraft v3 response format
|
||||
images = response_data.get("images", [])
|
||||
if isinstance(images, list):
|
||||
for image_data in images:
|
||||
if isinstance(image_data, dict):
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=None, # Recraft returns URLs only
|
||||
)
|
||||
)
|
||||
elif isinstance(image_data, str):
|
||||
# If images is just a list of URLs
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data,
|
||||
b64_json=None,
|
||||
)
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -0,0 +1,281 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from .transformation import FalAIBaseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIStableDiffusionConfig(FalAIBaseConfig):
|
||||
"""
|
||||
Configuration for Fal AI Stable Diffusion models.
|
||||
|
||||
Supports Stable Diffusion v3.5 variants and other Stable Diffusion models on Fal AI.
|
||||
|
||||
Example models:
|
||||
- fal-ai/stable-diffusion-v35-medium
|
||||
- fal-ai/stable-diffusion-v35-large
|
||||
|
||||
Documentation: https://fal.ai/models/fal-ai/stable-diffusion-v35-medium
|
||||
"""
|
||||
IMAGE_GENERATION_ENDPOINT: str = "" # Will be set from model name
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete url for the request.
|
||||
|
||||
For Stable Diffusion models, extract the endpoint from the model name.
|
||||
"""
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
complete_url: str = (
|
||||
api_base
|
||||
or get_secret_str("FAL_AI_API_BASE")
|
||||
or self.DEFAULT_BASE_URL
|
||||
)
|
||||
|
||||
complete_url = complete_url.rstrip("/")
|
||||
|
||||
# Extract endpoint from model name
|
||||
# e.g., "fal-ai/stable-diffusion-v35-medium" or "stable-diffusion-v35-medium"
|
||||
endpoint = model
|
||||
if "/" in model and not model.startswith("fal-ai/"):
|
||||
# If model is like "custom/stable-diffusion-v35-medium", use full path
|
||||
endpoint = model
|
||||
elif not model.startswith("fal-ai/"):
|
||||
# If model is just "stable-diffusion-v35-medium", prepend fal-ai
|
||||
endpoint = f"fal-ai/{model}"
|
||||
|
||||
complete_url = f"{complete_url}/{endpoint}"
|
||||
return complete_url
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for Stable Diffusion models.
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to Stable Diffusion parameters.
|
||||
|
||||
Mappings:
|
||||
- n -> num_images (1-4, default 1)
|
||||
- response_format -> output_format (jpeg or png)
|
||||
- size -> image_size (can be preset or custom width/height)
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Map OpenAI params to Stable Diffusion params
|
||||
param_mapping = {
|
||||
"n": "num_images",
|
||||
"response_format": "output_format",
|
||||
"size": "image_size",
|
||||
}
|
||||
|
||||
for k in non_default_params.keys():
|
||||
if k not in optional_params.keys():
|
||||
if k in supported_params:
|
||||
# Use mapped parameter name if exists
|
||||
mapped_key = param_mapping.get(k, k)
|
||||
mapped_value = non_default_params[k]
|
||||
|
||||
# Transform specific parameters
|
||||
if k == "response_format":
|
||||
# Map OpenAI response formats to image formats
|
||||
if mapped_value in ["b64_json", "url"]:
|
||||
mapped_value = "jpeg"
|
||||
elif k == "size":
|
||||
# Map OpenAI size format to Stable Diffusion image_size
|
||||
mapped_value = self._map_image_size(mapped_value)
|
||||
|
||||
optional_params[mapped_key] = mapped_value
|
||||
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
|
||||
|
||||
def _map_image_size(self, size: str) -> Any:
|
||||
"""
|
||||
Map OpenAI size format to Stable Diffusion image_size format.
|
||||
|
||||
OpenAI format: "1024x1024", "1792x1024", etc.
|
||||
Stable Diffusion format: Can be preset strings or {"width": int, "height": int}
|
||||
|
||||
Available presets:
|
||||
- square_hd
|
||||
- square
|
||||
- portrait_4_3
|
||||
- portrait_16_9
|
||||
- landscape_4_3 (default)
|
||||
- landscape_16_9
|
||||
"""
|
||||
# Map common OpenAI sizes to Stable Diffusion presets
|
||||
size_mapping = {
|
||||
"1024x1024": "square_hd",
|
||||
"512x512": "square",
|
||||
"768x1024": "portrait_4_3",
|
||||
"576x1024": "portrait_16_9",
|
||||
"1024x768": "landscape_4_3",
|
||||
"1024x576": "landscape_16_9",
|
||||
}
|
||||
|
||||
if size in size_mapping:
|
||||
return size_mapping[size]
|
||||
|
||||
# Parse custom size format "WIDTHxHEIGHT"
|
||||
if "x" in size:
|
||||
try:
|
||||
width, height = size.split("x")
|
||||
return {
|
||||
"width": int(width),
|
||||
"height": int(height),
|
||||
}
|
||||
except (ValueError, AttributeError):
|
||||
pass
|
||||
|
||||
# Default to landscape_4_3
|
||||
return "landscape_4_3"
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to Stable Diffusion request body.
|
||||
|
||||
Required parameters:
|
||||
- prompt: The prompt to generate an image from
|
||||
|
||||
Optional parameters:
|
||||
- num_images: Number of images (1-4, default: 1)
|
||||
- image_size: Size preset or {"width": int, "height": int} (default: landscape_4_3)
|
||||
- output_format: "jpeg" or "png" (default: jpeg)
|
||||
- sync_mode: Wait for image upload before returning (default: false)
|
||||
- guidance_scale: CFG scale 0-20 (default: 4.5)
|
||||
- num_inference_steps: Inference steps 1-50 (default: 40)
|
||||
- seed: Random seed for reproducibility
|
||||
- negative_prompt: Negative prompt string (default: "")
|
||||
- enable_safety_checker: Enable safety checker (default: true)
|
||||
"""
|
||||
stable_diffusion_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
return stable_diffusion_request_body
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the Stable Diffusion response to litellm ImageResponse format.
|
||||
|
||||
Expected response format:
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"url": "https://...",
|
||||
"width": 1024,
|
||||
"height": 768,
|
||||
"content_type": "image/jpeg"
|
||||
}
|
||||
],
|
||||
"timings": {"inference": 2.5, ...},
|
||||
"seed": 42,
|
||||
"has_nsfw_concepts": [false],
|
||||
"prompt": "original prompt"
|
||||
}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle Stable Diffusion response format
|
||||
images = response_data.get("images", [])
|
||||
if isinstance(images, list):
|
||||
for image_data in images:
|
||||
if isinstance(image_data, dict):
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=None, # Stable Diffusion returns URLs only
|
||||
)
|
||||
)
|
||||
elif isinstance(image_data, str):
|
||||
# If images is just a list of URLs
|
||||
model_response.data.append(
|
||||
ImageObject(
|
||||
url=image_data,
|
||||
b64_json=None,
|
||||
)
|
||||
)
|
||||
|
||||
# Add additional metadata from Stable Diffusion response
|
||||
if hasattr(model_response, "_hidden_params"):
|
||||
if "seed" in response_data:
|
||||
model_response._hidden_params["seed"] = response_data["seed"]
|
||||
if "timings" in response_data:
|
||||
model_response._hidden_params["timings"] = response_data["timings"]
|
||||
if "has_nsfw_concepts" in response_data:
|
||||
model_response._hidden_params["has_nsfw_concepts"] = response_data[
|
||||
"has_nsfw_concepts"
|
||||
]
|
||||
|
||||
return model_response
|
||||
|
||||
176
litellm/llms/fal_ai/image_generation/transformation.py
Normal file
176
litellm/llms/fal_ai/image_generation/transformation.py
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIImageGenerationOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class FalAIBaseConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
Base configuration for Fal AI image generation models.
|
||||
Handles common functionality like URL construction and authentication.
|
||||
"""
|
||||
DEFAULT_BASE_URL: str = "https://fal.run"
|
||||
IMAGE_GENERATION_ENDPOINT: str = ""
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete url for the request
|
||||
|
||||
Some providers need `model` in `api_base`
|
||||
"""
|
||||
complete_url: str = (
|
||||
api_base
|
||||
or get_secret_str("FAL_AI_API_BASE")
|
||||
or self.DEFAULT_BASE_URL
|
||||
)
|
||||
|
||||
complete_url = complete_url.rstrip("/")
|
||||
if self.IMAGE_GENERATION_ENDPOINT:
|
||||
complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}"
|
||||
return complete_url
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
final_api_key: Optional[str] = (
|
||||
api_key or
|
||||
get_secret_str("FAL_AI_API_KEY")
|
||||
)
|
||||
if not final_api_key:
|
||||
raise ValueError("FAL_AI_API_KEY is not set")
|
||||
|
||||
headers["Authorization"] = f"Key {final_api_key}"
|
||||
return headers
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform the image generation response to the litellm image response
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error transforming image generation response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
# Handle fal.ai response format
|
||||
images = response_data.get("images", [])
|
||||
if isinstance(images, list):
|
||||
for image_data in images:
|
||||
if isinstance(image_data, dict):
|
||||
model_response.data.append(ImageObject(
|
||||
url=image_data.get("url", None),
|
||||
b64_json=image_data.get("b64_json", None),
|
||||
))
|
||||
elif isinstance(image_data, str):
|
||||
# If images is just a list of URLs
|
||||
model_response.data.append(ImageObject(
|
||||
url=image_data,
|
||||
b64_json=None,
|
||||
))
|
||||
|
||||
return model_response
|
||||
|
||||
|
||||
class FalAIImageGenerationConfig(FalAIBaseConfig):
|
||||
"""
|
||||
Default Fal AI image generation configuration for generic models.
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Get supported OpenAI parameters for fal.ai image generation
|
||||
"""
|
||||
return [
|
||||
"n",
|
||||
"response_format",
|
||||
"size",
|
||||
]
|
||||
|
||||
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
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the image generation request to the fal.ai image generation request body
|
||||
"""
|
||||
fal_ai_image_generation_request_body = {
|
||||
"prompt": prompt,
|
||||
**optional_params,
|
||||
}
|
||||
return fal_ai_image_generation_request_body
|
||||
|
||||
|
|
@ -8259,6 +8259,46 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"fal_ai/bria/text-to-image/3.2": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/flux-pro/v1.1-ultra": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/imagen4/preview": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/recraft/v3/text-to-image": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/stable-diffusion-v35-medium": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"featherless_ai/featherless-ai/Qwerky-72B": {
|
||||
"litellm_provider": "featherless_ai",
|
||||
"max_input_tokens": 32768,
|
||||
|
|
|
|||
BIN
litellm/proxy/_experimental/out/assets/logos/fal_ai.jpg
Normal file
BIN
litellm/proxy/_experimental/out/assets/logos/fal_ai.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 8.1 KiB |
|
|
@ -2554,6 +2554,7 @@ class LlmProviders(str, Enum):
|
|||
PG_VECTOR = "pg_vector"
|
||||
HYPERBOLIC = "hyperbolic"
|
||||
RECRAFT = "recraft"
|
||||
FAL_AI = "fal_ai"
|
||||
HEROKU = "heroku"
|
||||
AIML = "aiml"
|
||||
COMETAPI = "cometapi"
|
||||
|
|
|
|||
|
|
@ -9,9 +9,9 @@
|
|||
|
||||
import ast
|
||||
import asyncio
|
||||
import contextvars
|
||||
import base64
|
||||
import binascii
|
||||
import contextvars
|
||||
import copy
|
||||
import datetime
|
||||
import hashlib
|
||||
|
|
@ -387,7 +387,15 @@ def print_verbose(
|
|||
####### CLIENT ###################
|
||||
# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking
|
||||
def load_custom_provider_entrypoints():
|
||||
found_entry_points = tuple(entry_points().select(group="litellm")) # type: ignore
|
||||
# Handle both Python 3.9 (returns dict) and Python 3.10+ (returns object with select method)
|
||||
eps = entry_points()
|
||||
if hasattr(eps, "select"):
|
||||
# Python 3.10+
|
||||
found_entry_points = tuple(eps.select(group="litellm")) # type: ignore
|
||||
else:
|
||||
# Python 3.9 and earlier - entry_points() returns a dict
|
||||
found_entry_points = eps.get("litellm", ()) # type: ignore
|
||||
|
||||
for entry_point in found_entry_points:
|
||||
# types are ignored because of circular dependency issues importing CustomLLM and CustomLLMItem
|
||||
HandlerClass = entry_point.load()
|
||||
|
|
@ -7658,6 +7666,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return LiteLLMProxyImageGenerationConfig()
|
||||
elif LlmProviders.FAL_AI == provider:
|
||||
from litellm.llms.fal_ai.image_generation import (
|
||||
get_fal_ai_image_generation_config,
|
||||
)
|
||||
|
||||
return get_fal_ai_image_generation_config(model)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -8259,6 +8259,46 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"fal_ai/bria/text-to-image/3.2": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/flux-pro/v1.1-ultra": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/imagen4/preview": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/recraft/v3/text-to-image": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"fal_ai/fal-ai/stable-diffusion-v35-medium": {
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.0398,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
]
|
||||
},
|
||||
"featherless_ai/featherless-ai/Qwerky-72B": {
|
||||
"litellm_provider": "featherless_ai",
|
||||
"max_input_tokens": 32768,
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ class BaseImageGenTest(ABC):
|
|||
async def test_basic_image_generation(self):
|
||||
"""Test basic image generation"""
|
||||
try:
|
||||
litellm._turn_on_debug()
|
||||
custom_logger = TestCustomLogger()
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
litellm.callbacks = [custom_logger]
|
||||
|
|
@ -55,7 +56,7 @@ class BaseImageGenTest(ABC):
|
|||
response = await litellm.aimage_generation(
|
||||
**base_image_generation_call_args, prompt="A image of a otter"
|
||||
)
|
||||
print(response)
|
||||
print("FAL AI RESPONSE: ", response)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
|
||||
|
|
|
|||
|
|
@ -168,7 +168,7 @@ def test_get_request_body_stability3():
|
|||
model = "stability.sd3-large"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["prompt"] == prompt
|
||||
|
|
@ -181,7 +181,7 @@ def test_get_request_body_stability():
|
|||
model = "stability.stable-diffusion-xl-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["text_prompts"][0]["text"] == prompt
|
||||
|
|
@ -239,7 +239,7 @@ def test_get_request_body_nova_canvas_default():
|
|||
model = "amazon.nova-canvas-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "TEXT_IMAGE"
|
||||
|
|
@ -254,7 +254,7 @@ def test_get_request_body_nova_canvas_text_image():
|
|||
model = "amazon.nova-canvas-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "TEXT_IMAGE"
|
||||
|
|
@ -273,7 +273,7 @@ def test_get_request_body_nova_canvas_color_guided_generation():
|
|||
model = "amazon.nova-canvas-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "COLOR_GUIDED_GENERATION"
|
||||
|
|
@ -434,7 +434,7 @@ def test_get_request_body_nova_canvas_inference_profile_arn():
|
|||
nova_model = "us.amazon.nova-canvas-v1:0"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=nova_model, prompt=prompt, optional_params=optional_params
|
||||
model=nova_model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "TEXT_IMAGE"
|
||||
|
|
@ -450,7 +450,7 @@ def test_get_request_body_nova_canvas_with_model_id_param():
|
|||
model = "amazon.nova-canvas-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
# After fix, model_id should not appear in the result
|
||||
|
|
@ -487,7 +487,7 @@ def test_get_request_body_cross_region_inference_profile():
|
|||
|
||||
# This should work after the fix - cross-region format should be detected as 'nova'
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "TEXT_IMAGE"
|
||||
|
|
@ -502,7 +502,7 @@ def test_backward_compatibility_regular_nova_model():
|
|||
model = "amazon.nova-canvas-v1"
|
||||
|
||||
result = handler._get_request_body(
|
||||
model=model, prompt=prompt, optional_params=optional_params
|
||||
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
|
||||
)
|
||||
|
||||
assert result["taskType"] == "TEXT_IMAGE"
|
||||
|
|
|
|||
68
tests/image_gen_tests/test_fal_ai_image_generation.py
Normal file
68
tests/image_gen_tests/test_fal_ai_image_generation.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
from litellm import aimage_generation
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
"fal_ai/fal-ai/recraft/v3/text-to-image",
|
||||
"fal_ai/bria/text-to-image/3.2",
|
||||
"fal_ai/fal-ai/stable-diffusion-v35-medium"
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_fal_ai_image_generation_basic(model):
|
||||
"""
|
||||
Test basic image generation for various Fal AI models.
|
||||
|
||||
Tests that each model can:
|
||||
- Accept a basic text prompt
|
||||
- Return a valid response with image data
|
||||
- Handle the response properly through litellm
|
||||
"""
|
||||
try:
|
||||
litellm.set_verbose = True
|
||||
|
||||
response = await aimage_generation(
|
||||
model=model,
|
||||
prompt="A cute baby sea otter",
|
||||
)
|
||||
|
||||
print(f"\nResponse from {model}:")
|
||||
print(f" Number of images: {len(response.data)}")
|
||||
print(f" First image URL: {response.data[0].url if response.data else 'None'}")
|
||||
|
||||
# Basic assertions
|
||||
assert response is not None, f"Response should not be None for {model}"
|
||||
assert hasattr(response, "data"), f"Response should have data attribute for {model}"
|
||||
assert len(response.data) > 0, f"Response should have at least one image for {model}"
|
||||
|
||||
# Check that we got a URL or b64_json
|
||||
first_image = response.data[0]
|
||||
assert (
|
||||
first_image.url is not None or first_image.b64_json is not None
|
||||
), f"Image should have either url or b64_json for {model}"
|
||||
|
||||
print(f"✓ Test passed for {model}")
|
||||
|
||||
except litellm.RateLimitError as e:
|
||||
pytest.skip(f"Rate limit error for {model}: {str(e)}")
|
||||
except litellm.ContentPolicyViolationError as e:
|
||||
pytest.skip(f"Content policy violation for {model}: {str(e)}")
|
||||
except litellm.InternalServerError as e:
|
||||
pytest.skip(f"Internal server error for {model}: {str(e)}")
|
||||
except Exception as e:
|
||||
if "Your task failed as a result of our safety system" in str(e):
|
||||
pytest.skip(f"Safety system rejection for {model}")
|
||||
else:
|
||||
pytest.fail(f"Test failed for {model}: {str(e)}")
|
||||
|
||||
|
|
@ -164,6 +164,9 @@ class TestAimlImageGeneration(BaseImageGenTest):
|
|||
def get_base_image_generation_call_args(self) -> dict:
|
||||
return {"model": "aiml/flux-pro/v1.1"}
|
||||
|
||||
class TestFAL_AI_ImageGeneration(BaseImageGenTest):
|
||||
def get_base_image_generation_call_args(self) -> dict:
|
||||
return {"model": "fal_ai/fal-ai/imagen4/preview"}
|
||||
|
||||
class TestGoogleImageGen(BaseImageGenTest):
|
||||
def get_base_image_generation_call_args(self) -> dict:
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ import traceback
|
|||
|
||||
import openai
|
||||
import pytest
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
|
|
@ -540,13 +539,14 @@ async def test_simple_aembedding():
|
|||
}
|
||||
|
||||
|
||||
def test_custom_llm_provider_entrypoint(mocker: MockerFixture):
|
||||
def test_custom_llm_provider_entrypoint():
|
||||
# This test mocks the use of entry-points in pyproject.toml:
|
||||
# [project.entry-point.litellm]
|
||||
# custom_llm = <module>:MyCustomLLM
|
||||
# another-custom-llm = <module>:AnotherCustomLLM
|
||||
|
||||
from litellm.utils import custom_llm_setup
|
||||
from importlib.metadata import EntryPoints, EntryPoint
|
||||
|
||||
class AnotherCustomLLM(CustomLLM):
|
||||
pass
|
||||
|
|
@ -559,25 +559,22 @@ def test_custom_llm_provider_entrypoint(mocker: MockerFixture):
|
|||
def load(self):
|
||||
return providers[self.name]
|
||||
|
||||
mocker.patch("importlib.metadata.EntryPoint.load", load)
|
||||
from importlib.metadata import EntryPoints, EntryPoint
|
||||
|
||||
entry_points = EntryPoints([
|
||||
EntryPoint(group="litellm", name="custom_llm", value="package.module:MyCustomLLM"),
|
||||
EntryPoint(group="litellm", name="another-custom-llm", value="package.module:AnotherCustomLLM"),
|
||||
])
|
||||
mocked = mocker.patch("litellm.utils.entry_points")
|
||||
mocked.return_value = entry_points
|
||||
|
||||
assert litellm.custom_provider_map == []
|
||||
assert litellm._custom_providers == []
|
||||
with patch("importlib.metadata.EntryPoint.load", load):
|
||||
with patch("litellm.utils.entry_points", return_value=entry_points):
|
||||
assert litellm.custom_provider_map == []
|
||||
assert litellm._custom_providers == []
|
||||
|
||||
custom_llm_setup()
|
||||
custom_llm_setup()
|
||||
|
||||
assert litellm._custom_providers == ['custom_llm', 'another-custom-llm']
|
||||
assert litellm._custom_providers == ['custom_llm', 'another-custom-llm']
|
||||
|
||||
assert litellm.custom_provider_map[0]["provider"] == "custom_llm"
|
||||
assert isinstance(litellm.custom_provider_map[0]["custom_handler"], CustomLLM)
|
||||
assert litellm.custom_provider_map[0]["provider"] == "custom_llm"
|
||||
assert isinstance(litellm.custom_provider_map[0]["custom_handler"], CustomLLM)
|
||||
|
||||
assert litellm.custom_provider_map[1]["provider"] == "another-custom-llm"
|
||||
assert isinstance(litellm.custom_provider_map[1]["custom_handler"], AnotherCustomLLM)
|
||||
assert litellm.custom_provider_map[1]["provider"] == "another-custom-llm"
|
||||
assert isinstance(litellm.custom_provider_map[1]["custom_handler"], AnotherCustomLLM)
|
||||
|
|
|
|||
BIN
ui/litellm-dashboard/public/assets/logos/fal_ai.jpg
Normal file
BIN
ui/litellm-dashboard/public/assets/logos/fal_ai.jpg
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 8.1 KiB |
|
|
@ -575,6 +575,14 @@ const PROVIDER_CREDENTIAL_FIELDS: Record<Providers, ProviderCredentialField[]> =
|
|||
placeholder: "http://localhost:7997",
|
||||
},
|
||||
],
|
||||
[Providers.FalAI]: [
|
||||
{
|
||||
key: "api_key",
|
||||
label: "API Key",
|
||||
type: "password",
|
||||
required: true,
|
||||
}
|
||||
],
|
||||
};
|
||||
|
||||
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selectedProvider, uploadProps }) => {
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ export enum Providers {
|
|||
Deepgram = "Deepgram",
|
||||
Deepseek = "Deepseek",
|
||||
ElevenLabs = "ElevenLabs",
|
||||
FalAI = "Fal AI",
|
||||
FireworksAI = "Fireworks AI",
|
||||
Google_AI_Studio = "Google AI Studio",
|
||||
GradientAI = "GradientAI",
|
||||
|
|
@ -73,6 +74,7 @@ export const provider_map: Record<string, string> = {
|
|||
Triton: "triton",
|
||||
Deepgram: "deepgram",
|
||||
ElevenLabs: "elevenlabs",
|
||||
FalAI: "fal_ai",
|
||||
SageMaker: "sagemaker_chat",
|
||||
Voyage: "voyage",
|
||||
JinaAI: "jina_ai",
|
||||
|
|
@ -120,6 +122,7 @@ export const providerLogoMap: Record<string, string> = {
|
|||
[Providers.Triton]: `${asset_logos_folder}nvidia_triton.png`,
|
||||
[Providers.Deepgram]: `${asset_logos_folder}deepgram.png`,
|
||||
[Providers.ElevenLabs]: `${asset_logos_folder}elevenlabs.png`,
|
||||
[Providers.FalAI]: `${asset_logos_folder}fal_ai.jpg`,
|
||||
[Providers.Voyage]: `${asset_logos_folder}voyage.webp`,
|
||||
[Providers.JinaAI]: `${asset_logos_folder}jina.png`,
|
||||
[Providers.VolcEngine]: `${asset_logos_folder}volcengine.png`,
|
||||
|
|
@ -183,6 +186,8 @@ export const getPlaceholder = (selectedProvider: string): string => {
|
|||
return "volcengine/<any-model-on-volcengine>";
|
||||
} else if (selectedProvider == Providers.DeepInfra) {
|
||||
return "deepinfra/<any-model-on-deepinfra>";
|
||||
} else if (selectedProvider == Providers.FalAI) {
|
||||
return "fal_ai/fal-ai/flux-pro/v1.1-ultra";
|
||||
} else {
|
||||
return "gpt-3.5-turbo";
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue