mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-03 02:23:51 +00:00
fix: apply default model params to outbound payloads
This commit is contained in:
parent
1a97751e37
commit
be7b69403d
5 changed files with 204 additions and 5 deletions
|
|
@ -265,6 +265,10 @@ async def generate_function_chat_completion(request, form_data, user, models: di
|
|||
}
|
||||
extra_params['__tools__'] = metadata.get('tools', {})
|
||||
|
||||
default_params = getattr(request.app.state.config, 'DEFAULT_MODEL_PARAMS', None) or {}
|
||||
if default_params:
|
||||
form_data = apply_model_params_to_body_openai(default_params, form_data, overwrite=False)
|
||||
|
||||
if model_info:
|
||||
if model_info.base_model_id:
|
||||
form_data['model'] = model_info.base_model_id
|
||||
|
|
|
|||
|
|
@ -1018,6 +1018,10 @@ async def generate_chat_completion(
|
|||
model_id = payload['model']
|
||||
model_info = await Models.get_model_by_id(model_id)
|
||||
|
||||
default_params = getattr(request.app.state.config, 'DEFAULT_MODEL_PARAMS', None) or {}
|
||||
if default_params:
|
||||
payload = apply_model_params_to_body_ollama(default_params, payload, overwrite=False)
|
||||
|
||||
if model_info is not None:
|
||||
if model_info.base_model_id:
|
||||
base_model_id = request.base_model_id if hasattr(request, 'base_model_id') else model_info.base_model_id
|
||||
|
|
@ -1110,6 +1114,11 @@ async def generate_openai_completion(
|
|||
|
||||
model_id = form_data.model
|
||||
model_info = await Models.get_model_by_id(model_id)
|
||||
|
||||
default_params = getattr(request.app.state.config, 'DEFAULT_MODEL_PARAMS', None) or {}
|
||||
if default_params:
|
||||
payload = apply_model_params_to_body_openai(default_params, payload, overwrite=False)
|
||||
|
||||
if model_info is not None:
|
||||
if model_info.base_model_id:
|
||||
payload['model'] = model_info.base_model_id
|
||||
|
|
@ -1163,6 +1172,11 @@ async def generate_openai_chat_completion(
|
|||
|
||||
model_id = form_data.model
|
||||
model_info = await Models.get_model_by_id(model_id)
|
||||
|
||||
default_params = getattr(request.app.state.config, 'DEFAULT_MODEL_PARAMS', None) or {}
|
||||
if default_params:
|
||||
payload = apply_model_params_to_body_openai(default_params, payload, overwrite=False)
|
||||
|
||||
if model_info is not None:
|
||||
if model_info.base_model_id:
|
||||
payload['model'] = model_info.base_model_id
|
||||
|
|
|
|||
|
|
@ -1077,6 +1077,10 @@ async def generate_chat_completion(
|
|||
model_id = form_data.get('model')
|
||||
model_info = await Models.get_model_by_id(model_id)
|
||||
|
||||
default_params = getattr(request.app.state.config, 'DEFAULT_MODEL_PARAMS', None) or {}
|
||||
if default_params:
|
||||
payload = apply_model_params_to_body_openai(default_params, payload, overwrite=False)
|
||||
|
||||
# Check model info and override the payload
|
||||
if model_info:
|
||||
if model_info.base_model_id:
|
||||
|
|
|
|||
|
|
@ -41,12 +41,20 @@ async def apply_system_prompt_to_body(
|
|||
|
||||
|
||||
# inplace function: form_data is modified
|
||||
def apply_model_params_to_body(params: dict, form_data: dict, mappings: dict[str, Callable]) -> dict:
|
||||
def apply_model_params_to_body(
|
||||
params: dict,
|
||||
form_data: dict,
|
||||
mappings: dict[str, Callable],
|
||||
overwrite: bool = True,
|
||||
) -> dict:
|
||||
if not params:
|
||||
return form_data
|
||||
|
||||
for key, value in params.items():
|
||||
if value is not None:
|
||||
if not overwrite and key in form_data:
|
||||
continue
|
||||
|
||||
if key in mappings:
|
||||
cast_func = mappings[key]
|
||||
if isinstance(cast_func, Callable):
|
||||
|
|
@ -83,7 +91,11 @@ def remove_open_webui_params(params: dict) -> dict:
|
|||
|
||||
|
||||
# inplace function: form_data is modified
|
||||
def apply_model_params_to_body_openai(params: dict, form_data: dict) -> dict:
|
||||
OPENAI_TOKEN_LIMIT_PARAMS = {'max_tokens', 'max_completion_tokens'}
|
||||
|
||||
|
||||
def apply_model_params_to_body_openai(params: dict, form_data: dict, overwrite: bool = True) -> dict:
|
||||
params = copy.deepcopy(params or {})
|
||||
params = remove_open_webui_params(params)
|
||||
|
||||
custom_params = params.pop('custom_params', {})
|
||||
|
|
@ -101,11 +113,15 @@ def apply_model_params_to_body_openai(params: dict, form_data: dict) -> dict:
|
|||
# If there are custom parameters, we need to apply them first
|
||||
params = deep_update(params, custom_params)
|
||||
|
||||
if not overwrite and any(key in form_data for key in OPENAI_TOKEN_LIMIT_PARAMS):
|
||||
params = {key: value for key, value in params.items() if key not in OPENAI_TOKEN_LIMIT_PARAMS}
|
||||
|
||||
mappings = {
|
||||
'temperature': float,
|
||||
'top_p': float,
|
||||
'min_p': float,
|
||||
'max_tokens': int,
|
||||
'max_completion_tokens': int,
|
||||
'frequency_penalty': float,
|
||||
'presence_penalty': float,
|
||||
'reasoning_effort': str,
|
||||
|
|
@ -114,10 +130,11 @@ def apply_model_params_to_body_openai(params: dict, form_data: dict) -> dict:
|
|||
'logit_bias': lambda x: x,
|
||||
'response_format': dict,
|
||||
}
|
||||
return apply_model_params_to_body(params, form_data, mappings)
|
||||
return apply_model_params_to_body(params, form_data, mappings, overwrite=overwrite)
|
||||
|
||||
|
||||
def apply_model_params_to_body_ollama(params: dict, form_data: dict) -> dict:
|
||||
def apply_model_params_to_body_ollama(params: dict, form_data: dict, overwrite: bool = True) -> dict:
|
||||
params = copy.deepcopy(params or {})
|
||||
params = remove_open_webui_params(params)
|
||||
|
||||
custom_params = params.pop('custom_params', {})
|
||||
|
|
@ -188,12 +205,21 @@ def apply_model_params_to_body_ollama(params: dict, form_data: dict) -> dict:
|
|||
|
||||
for key, value in ollama_root_params.items():
|
||||
if (param := params.get(key, None)) is not None:
|
||||
if not overwrite and key in form_data:
|
||||
del params[key]
|
||||
continue
|
||||
|
||||
# Copy the parameter to new name then delete it, to prevent Ollama warning of invalid option provided
|
||||
form_data[key] = value(param)
|
||||
del params[key]
|
||||
|
||||
# Unlike OpenAI, Ollama does not support params directly in the body
|
||||
form_data['options'] = apply_model_params_to_body(params, (form_data.get('options', {}) or {}), mappings)
|
||||
form_data['options'] = apply_model_params_to_body(
|
||||
params,
|
||||
(form_data.get('options', {}) or {}),
|
||||
mappings,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
return form_data
|
||||
|
||||
|
||||
|
|
|
|||
151
test/test_default_model_params_payload.py
Normal file
151
test/test_default_model_params_payload.py
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def payload(monkeypatch):
|
||||
root = Path(__file__).resolve().parents[1]
|
||||
|
||||
open_webui = types.ModuleType('open_webui')
|
||||
open_webui_utils = types.ModuleType('open_webui.utils')
|
||||
|
||||
misc = types.ModuleType('open_webui.utils.misc')
|
||||
misc.add_or_update_system_message = lambda system, messages: messages
|
||||
misc.replace_system_message_content = lambda system, messages: messages
|
||||
|
||||
def deep_update(target, update):
|
||||
for key, value in update.items():
|
||||
if isinstance(value, dict) and isinstance(target.get(key), dict):
|
||||
deep_update(target[key], value)
|
||||
else:
|
||||
target[key] = value
|
||||
return target
|
||||
|
||||
misc.deep_update = deep_update
|
||||
|
||||
task = types.ModuleType('open_webui.utils.task')
|
||||
|
||||
async def prompt_template(system, user=None):
|
||||
return system
|
||||
|
||||
task.prompt_template = prompt_template
|
||||
task.prompt_variables_template = lambda system, variables: system
|
||||
|
||||
monkeypatch.setitem(sys.modules, 'open_webui', open_webui)
|
||||
monkeypatch.setitem(sys.modules, 'open_webui.utils', open_webui_utils)
|
||||
monkeypatch.setitem(sys.modules, 'open_webui.utils.misc', misc)
|
||||
monkeypatch.setitem(sys.modules, 'open_webui.utils.task', task)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
'payload_under_test',
|
||||
root / 'backend/open_webui/utils/payload.py',
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def test_openai_default_model_params_fill_outbound_payload(payload):
|
||||
form_data = {
|
||||
'model': 'gpt-test',
|
||||
'messages': [{'role': 'user', 'content': 'hello'}],
|
||||
}
|
||||
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_tokens': '256', 'temperature': '0.7'},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
|
||||
assert result['max_tokens'] == 256
|
||||
assert result['temperature'] == 0.7
|
||||
|
||||
|
||||
def test_openai_explicit_payload_params_override_defaults(payload):
|
||||
form_data = {
|
||||
'model': 'gpt-test',
|
||||
'messages': [{'role': 'user', 'content': 'hello'}],
|
||||
'max_tokens': 64,
|
||||
'temperature': 0.2,
|
||||
}
|
||||
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_tokens': 256, 'temperature': 0.7},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
|
||||
assert result['max_tokens'] == 64
|
||||
assert result['temperature'] == 0.2
|
||||
|
||||
|
||||
def test_openai_explicit_max_tokens_overrides_default_max_completion_tokens(payload):
|
||||
form_data = {
|
||||
'model': 'gpt-test',
|
||||
'messages': [{'role': 'user', 'content': 'hello'}],
|
||||
'max_tokens': 64,
|
||||
}
|
||||
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_completion_tokens': 256},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
|
||||
assert result['max_tokens'] == 64
|
||||
assert 'max_completion_tokens' not in result
|
||||
|
||||
|
||||
def test_openai_explicit_max_completion_tokens_overrides_default_max_tokens(payload):
|
||||
form_data = {
|
||||
'model': 'gpt-test',
|
||||
'messages': [{'role': 'user', 'content': 'hello'}],
|
||||
'max_completion_tokens': 64,
|
||||
}
|
||||
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_tokens': 256},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
|
||||
assert result['max_completion_tokens'] == 64
|
||||
assert 'max_tokens' not in result
|
||||
|
||||
|
||||
def test_model_params_can_override_pre_applied_defaults(payload):
|
||||
form_data = {'model': 'gpt-test', 'messages': []}
|
||||
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_tokens': 256, 'temperature': 0.7},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
result = payload.apply_model_params_to_body_openai(
|
||||
{'max_tokens': 128, 'temperature': 0.4},
|
||||
result,
|
||||
)
|
||||
|
||||
assert result['max_tokens'] == 128
|
||||
assert result['temperature'] == 0.4
|
||||
|
||||
|
||||
def test_ollama_default_model_params_fill_options_without_overriding_explicit_options(payload):
|
||||
form_data = {
|
||||
'model': 'llama-test',
|
||||
'messages': [{'role': 'user', 'content': 'hello'}],
|
||||
'options': {'temperature': 0.2},
|
||||
}
|
||||
|
||||
result = payload.apply_model_params_to_body_ollama(
|
||||
{'max_tokens': '256', 'temperature': '0.7'},
|
||||
form_data,
|
||||
overwrite=False,
|
||||
)
|
||||
|
||||
assert result['options']['num_predict'] == 256
|
||||
assert result['options']['temperature'] == 0.2
|
||||
Loading…
Add table
Reference in a new issue