diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index 40abf13662..e62616ce37 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -2175,6 +2175,14 @@ TASK_MODEL = os.getenv('TASK_MODEL', '') TASK_MODEL_EXTERNAL = os.getenv('TASK_MODEL_EXTERNAL', '') +try: + task_model_params = JSONCodec.loads(os.getenv('TASK_MODEL_PARAMS', '{}')) +except Exception as e: + log.exception(f'Error loading TASK_MODEL_PARAMS: {e}') + task_model_params = {} + +TASK_MODEL_PARAMS = task_model_params + CONTEXT_COMPACTION_MODEL = os.getenv('CONTEXT_COMPACTION_MODEL', '') ENABLE_CONTEXT_COMPACTION = os.getenv('ENABLE_CONTEXT_COMPACTION', 'False').lower() == 'true' @@ -3104,6 +3112,7 @@ DEFAULT_CONFIG = { 'auth.admin.email': ADMIN_EMAIL, 'task.model.default': TASK_MODEL, 'task.model.external': TASK_MODEL_EXTERNAL, + 'task.model.params': TASK_MODEL_PARAMS, 'chat.context_compaction.model': CONTEXT_COMPACTION_MODEL, 'chat.context_compaction.enable': ENABLE_CONTEXT_COMPACTION, 'chat.context_compaction.token_threshold': CONTEXT_COMPACTION_TOKEN_THRESHOLD, diff --git a/backend/open_webui/models/config.py b/backend/open_webui/models/config.py index 8888b4ed7c..011a58d9ee 100644 --- a/backend/open_webui/models/config.py +++ b/backend/open_webui/models/config.py @@ -32,6 +32,7 @@ DICT_CONFIG_KEY_ALIASES = { 'audio.tts.openai.params': ('AUDIO_TTS_OPENAI_PARAMS',), 'models.default_metadata': ('DEFAULT_MODEL_METADATA',), 'models.default_params': ('DEFAULT_MODEL_PARAMS',), + 'task.model.params': ('TASK_MODEL_PARAMS',), 'ui.default_interface_settings': ('DEFAULT_INTERFACE_SETTINGS',), 'user.permissions': ('USER_PERMISSIONS',), } diff --git a/backend/open_webui/routers/tasks.py b/backend/open_webui/routers/tasks.py index 01e0fa2158..1dd53d8ebe 100644 --- a/backend/open_webui/routers/tasks.py +++ b/backend/open_webui/routers/tasks.py @@ -20,6 +20,7 @@ from open_webui.models.config import Config from open_webui.routers.pipelines import process_pipeline_inlet_filter from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.chat import generate_chat_completion +from open_webui.utils.payload import apply_params_to_form_data from open_webui.utils.task import ( autocomplete_generation_template, emoji_generation_template, @@ -40,6 +41,7 @@ router = APIRouter() TASK_CONFIG_KEYS = { 'TASK_MODEL': 'task.model.default', 'TASK_MODEL_EXTERNAL': 'task.model.external', + 'TASK_MODEL_PARAMS': 'task.model.params', 'TITLE_GENERATION_PROMPT_TEMPLATE': 'task.title.prompt_template', 'IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE': 'task.image.prompt_template', 'ENABLE_AUTOCOMPLETE_GENERATION': 'task.autocomplete.enable', @@ -68,6 +70,34 @@ def config_updates(data: dict, key_map: dict[str, str]) -> dict: return {key_map[field]: value for field, value in data.items() if field in key_map} +def apply_task_model_params(payload: dict, models: dict, task_model_id: str, params: dict | None = None) -> dict: + model = models.get(payload.get('model')) or models.get(task_model_id) + if not model or (not params and not payload.get('params')): + return payload + return apply_params_to_form_data(payload, model, params or None) + + +async def get_task_model_generation_config(default_model_id: str, models) -> tuple[str, dict]: + config = await Config.get_many( + 'task.model.default', + 'task.model.external', + 'task.model.params', + ) + params = config.get('task.model.params') or {} + if not isinstance(params, dict): + params = {} + + return ( + get_task_model_id( + default_model_id, + config.get('task.model.default'), + config.get('task.model.external'), + models, + ), + {key: value for key, value in params.items() if value is not None and value != ''}, + ) + + ################################## # # Task Endpoints @@ -83,6 +113,7 @@ async def get_task_config(request: Request, user=Depends(get_verified_user)): class TaskConfigForm(BaseModel): TASK_MODEL: Optional[str] TASK_MODEL_EXTERNAL: Optional[str] + TASK_MODEL_PARAMS: dict | None = None ENABLE_TITLE_GENERATION: bool TITLE_GENERATION_PROMPT_TEMPLATE: str IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE: str @@ -135,14 +166,7 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating chat title using model %s for user %s ', task_model_id, user.email) @@ -153,20 +177,14 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver template = DEFAULT_TITLE_GENERATION_PROMPT_TEMPLATE content = await title_generation_template(template, form_data['messages'], user) - - max_tokens = models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) + task_model_params = task_model_params or { + 'max_tokens': models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) + } payload = { 'model': task_model_id, 'messages': [{'role': 'user', 'content': content}], 'stream': False, - **( - {'max_tokens': max_tokens} - if models[task_model_id].get('owned_by') == 'ollama' - else { - 'max_completion_tokens': max_tokens, - } - ), 'metadata': { **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 'task': str(TASKS.TITLE_GENERATION), @@ -181,6 +199,8 @@ async def generate_title(request: Request, form_data: dict, user=Depends(get_ver except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -214,14 +234,7 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating chat title using model %s for user %s ', task_model_id, user.email) @@ -251,6 +264,8 @@ async def generate_follow_ups(request: Request, form_data: dict, user=Depends(ge except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -284,14 +299,7 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating chat tags using model %s for user %s ', task_model_id, user.email) @@ -321,6 +329,8 @@ async def generate_chat_tags(request: Request, form_data: dict, user=Depends(get except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -348,14 +358,7 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends( detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating image prompt using model %s for user %s ', task_model_id, user.email) @@ -385,6 +388,8 @@ async def generate_image_prompt(request: Request, form_data: dict, user=Depends( except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -430,14 +435,7 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating %s queries using model %s for user %s', type, task_model_id, user.email) @@ -467,6 +465,8 @@ async def generate_queries(request: Request, form_data: dict, user=Depends(get_v except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -511,14 +511,7 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, task_model_params = await get_task_model_generation_config(model_id, models) log.debug('generating autocompletion using model %s for user %s', task_model_id, user.email) @@ -548,6 +541,8 @@ async def generate_autocompletion(request: Request, form_data: dict, user=Depend except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, task_model_params) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: @@ -575,14 +570,7 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver detail=ERROR_MESSAGES.MODEL_NOT_FOUND(), ) - # Check if the user has a custom task model - # If the user has a custom task model, use that model - task_model_id = get_task_model_id( - model_id, - await Config.get('task.model.default'), - await Config.get('task.model.external'), - models, - ) + task_model_id, _ = await get_task_model_generation_config(model_id, models) log.debug('generating emoji using model %s for user %s ', task_model_id, user.email) @@ -594,13 +582,6 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver 'model': task_model_id, 'messages': [{'role': 'user', 'content': content}], 'stream': False, - **( - {'max_tokens': 4} - if models[task_model_id].get('owned_by') == 'ollama' - else { - 'max_completion_tokens': 4, - } - ), 'metadata': { **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 'task': str(TASKS.EMOJI_GENERATION), @@ -615,6 +596,8 @@ async def generate_emoji(request: Request, form_data: dict, user=Depends(get_ver except Exception as e: raise e + payload = apply_task_model_params(payload, models, task_model_id, {'max_tokens': 4}) + try: return await generate_chat_completion(request, form_data=payload, user=user) except Exception as e: diff --git a/backend/open_webui/utils/context_compaction.py b/backend/open_webui/utils/context_compaction.py index 0af8ee5907..ff84481a5e 100644 --- a/backend/open_webui/utils/context_compaction.py +++ b/backend/open_webui/utils/context_compaction.py @@ -9,6 +9,7 @@ from open_webui.models.config import Config from open_webui.utils.chat_id import is_saved_chat_id from open_webui.utils.json_codec import JSONCodec from open_webui.utils.misc import get_content_from_message, get_last_user_message, get_message_list +from open_webui.utils.payload import apply_params_to_form_data from open_webui.utils.task import ( get_task_model_id, prompt_template, @@ -367,6 +368,7 @@ async def _generate_summary( task_config = await Config.get_many( 'task.model.default', 'task.model.external', + 'task.model.params', 'chat.context_compaction.model', ) context_compaction_model = task_config.get('chat.context_compaction.model') @@ -394,22 +396,25 @@ async def _generate_summary( prompt = prompt_variables_template(prompt, {'{{PREVIOUS_SUMMARY}}': previous_summary or ''}) prompt = await prompt_template(prompt, user) - max_tokens = models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) + task_model_params = task_config.get('task.model.params') or {} + if not isinstance(task_model_params, dict): + task_model_params = {} + task_model_params = {key: value for key, value in task_model_params.items() if value is not None and value != ''} + task_model_params = task_model_params or { + 'max_tokens': models[task_model_id].get('info', {}).get('params', {}).get('max_tokens', 1000) + } + payload = { 'model': task_model_id, 'messages': [{'role': 'user', 'content': prompt}], 'stream': False, - **( - {'max_tokens': max_tokens} - if models[task_model_id].get('owned_by') == 'ollama' - else {'max_completion_tokens': max_tokens} - ), 'metadata': { **(request.state.metadata if hasattr(request.state, 'metadata') else {}), 'task': 'context_compaction', }, } + payload = apply_params_to_form_data(payload, models[task_model_id], task_model_params) response = await generate_chat_completion(request, form_data=payload, user=user) summary = _response_text(response).strip() if summary: diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 5a10db3386..338b652d45 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -100,9 +100,7 @@ from open_webui.utils.memory import add_memory_context, review_memory_after_turn from open_webui.utils.misc import ( add_or_update_system_message, add_or_update_user_message, - convert_logit_bias_input_to_json, convert_output_to_messages, - deep_update, extract_urls, get_content_from_message, get_last_assistant_message, @@ -118,7 +116,7 @@ from open_webui.utils.misc import ( set_last_user_message_content, strip_empty_content_blocks, ) -from open_webui.utils.payload import apply_system_prompt_to_body, resolve_system_prompt +from open_webui.utils.payload import apply_params_to_form_data, apply_system_prompt_to_body, resolve_system_prompt from open_webui.utils.plugin import load_function_module_by_id from open_webui.utils.response import merge_usage, normalize_usage from open_webui.utils.sanitize import sanitize_code @@ -1942,59 +1940,6 @@ async def chat_completion_files_handler( return body, {'sources': sources} -def apply_params_to_form_data(form_data, model): - params = form_data.pop('params', {}) - custom_params = params.pop('custom_params', {}) - - open_webui_params = { - 'stream_response': bool, - 'stream_delta_chunk_size': int, - 'function_calling': str, - 'reasoning_tags': list, - 'compact_token_threshold': int, - 'system': str, - 'note_id': str, - } - - for key in list(params.keys()): - if key in open_webui_params: - del params[key] - - if custom_params: - # Attempt to parse custom_params if they are strings - for key, value in custom_params.items(): - if isinstance(value, str): - try: - # Attempt to parse the string as JSON - custom_params[key] = JSONCodec.loads(value) - except JSONCodec.JSONDecodeError: - # If it fails, keep the original string - pass - - # If custom_params are provided, merge them into params - params = deep_update(params, custom_params) - - if model.get('owned_by') == 'ollama': - # Ollama specific parameters - form_data['options'] = {**params, **(form_data.get('options') or {})} - else: - if isinstance(params, dict): - for key, value in params.items(): - if value is not None and key not in form_data: - form_data[key] = value - - if 'logit_bias' in params and params['logit_bias'] is not None and 'logit_bias' not in form_data: - try: - logit_bias = convert_logit_bias_input_to_json(params['logit_bias']) - - if logit_bias: - form_data['logit_bias'] = JSONCodec.loads(logit_bias) - except Exception as e: - log.exception(f'Error parsing logit_bias: {e}') - - return form_data - - async def convert_url_images_to_base64(form_data, user=None): messages = form_data.get('messages', []) diff --git a/backend/open_webui/utils/payload.py b/backend/open_webui/utils/payload.py index 30ca698be0..167ca76bf8 100644 --- a/backend/open_webui/utils/payload.py +++ b/backend/open_webui/utils/payload.py @@ -78,6 +78,17 @@ def apply_model_params_to_body(params: dict, form_data: dict, mappings: dict[str return form_data +def apply_params_to_form_data(form_data: dict, model: dict, params: dict | None = None) -> dict: + payload_params = form_data.pop('params', {}) or {} + params = dict(payload_params if params is None else params) + if not params: + return form_data + + if model.get('owned_by') == 'ollama': + return apply_model_params_to_body_ollama(params, form_data) + return apply_model_params_to_body_openai(params, form_data) + + def remove_open_webui_params(params: dict) -> dict: """ Removes OpenWebUI specific parameters from the provided dictionary. @@ -95,6 +106,7 @@ def remove_open_webui_params(params: dict) -> dict: 'reasoning_tags': list, 'compact_token_threshold': int, 'system': str, + 'note_id': str, } for key in list(params.keys()): diff --git a/src/lib/components/admin/Settings/Interface.svelte b/src/lib/components/admin/Settings/Interface.svelte index c1cae4b83b..5c54a52b6e 100644 --- a/src/lib/components/admin/Settings/Interface.svelte +++ b/src/lib/components/admin/Settings/Interface.svelte @@ -10,6 +10,7 @@ import Textarea from '$lib/components/common/Textarea.svelte'; import Spinner from '$lib/components/common/Spinner.svelte'; import SettingsSelect from '$lib/components/common/SettingsSelect.svelte'; + import AdvancedParams from '$lib/components/chat/Settings/Advanced/AdvancedParams.svelte'; import AdminSettingField from './AdminSettingField.svelte'; import AdminSettingRow from './AdminSettingRow.svelte'; import AdminSettingSection from './AdminSettingSection.svelte'; @@ -22,6 +23,7 @@ let taskConfig = { TASK_MODEL: '', TASK_MODEL_EXTERNAL: '', + TASK_MODEL_PARAMS: {}, ENABLE_TITLE_GENERATION: true, TITLE_GENERATION_PROMPT_TEMPLATE: '', ENABLE_FOLLOW_UP_GENERATION: true, @@ -48,10 +50,23 @@ CONTEXT_COMPACTION_RETENTION_PERCENTAGE: 40, CONTEXT_COMPACTION_PROMPT_TEMPLATE: '' }; + let showTaskParameters = false; + + const configuredParams = (params: Record = {}) => + Object.fromEntries( + Object.entries(params).filter( + ([_, value]) => value !== null && value !== '' && value !== undefined + ) + ); const updateInterfaceHandler = async () => { + const taskConfigPayload = { + ...taskConfig, + TASK_MODEL_PARAMS: configuredParams(taskConfig.TASK_MODEL_PARAMS) + }; + [taskConfig, chatConfig] = await Promise.all([ - updateTaskConfig(localStorage.token, taskConfig), + updateTaskConfig(localStorage.token, taskConfigPayload), updateChatConfig(localStorage.token, chatConfig) ]); appConfig.update((current) => @@ -104,6 +119,7 @@ getTaskConfig(localStorage.token), getChatConfig(localStorage.token) ]); + taskConfig.TASK_MODEL_PARAMS = taskConfig.TASK_MODEL_PARAMS ?? {}; workspaceModels = await getBaseModels(localStorage.token); baseModels = await getModels(localStorage.token, null, false); @@ -206,6 +222,33 @@ + +
+ + + {#if showTaskParameters} +
+ +
+ {/if} +