mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-04 23:00:30 +08:00
Fix: handle llm_setting switch (#17745)
This commit is contained in:
@@ -18,6 +18,48 @@ from copy import deepcopy
|
||||
|
||||
GENERATION_CONFIG_KEYS = ("temperature", "top_p", "frequency_penalty", "presence_penalty", "max_tokens")
|
||||
|
||||
# Default values for the four LLM generation parameters stored in
|
||||
# search_config.llm_setting. When the corresponding ``_enabled`` flag is
|
||||
# ``False`` (or the key is absent), the default is used instead of whatever
|
||||
# the user may have stored in ``llm_setting``.
|
||||
LLM_SETTING_DEFAULTS = {
|
||||
"temperature": 0.1,
|
||||
"top_p": 0.3,
|
||||
"frequency_penalty": 0.7,
|
||||
"presence_penalty": 0.4,
|
||||
}
|
||||
|
||||
|
||||
def resolve_llm_setting(llm_setting):
|
||||
"""Resolve *llm_setting* values according to their enable flags.
|
||||
|
||||
For each of the four generation parameters the dictionary may carry a
|
||||
``{key}_enabled`` boolean. When the flag is ``True`` and the value
|
||||
exists in *llm_setting*, the user-configured value is kept; otherwise
|
||||
the default from :data:`LLM_SETTING_DEFAULTS` is substituted.
|
||||
|
||||
Keys whose name ends with ``_enabled`` are stripped from the result so
|
||||
they are never forwarded to the downstream LLM call.
|
||||
"""
|
||||
if not llm_setting:
|
||||
return dict(LLM_SETTING_DEFAULTS)
|
||||
|
||||
resolved = {}
|
||||
for key, default_val in LLM_SETTING_DEFAULTS.items():
|
||||
enabled_key = f"{key}_enabled"
|
||||
if llm_setting.get(enabled_key, True) and key in llm_setting:
|
||||
resolved[key] = llm_setting[key]
|
||||
else:
|
||||
resolved[key] = default_val
|
||||
|
||||
# Carry over any extra keys that are not generation parameters and not
|
||||
# enable flags (e.g. ``llm_id``, ``model_type``).
|
||||
for key, val in llm_setting.items():
|
||||
if key not in resolved and not key.endswith("_enabled"):
|
||||
resolved[key] = val
|
||||
|
||||
return resolved
|
||||
|
||||
|
||||
def extract_generation_config(req):
|
||||
return {key: req[key] for key in GENERATION_CONFIG_KEYS if key in req and req[key] is not None}
|
||||
|
||||
@@ -36,6 +36,7 @@ from common.metadata_utils import apply_meta_data_filter
|
||||
from api.db.services.search_service import SearchService
|
||||
from api.db.services.user_service import UserTenantService
|
||||
from api.db.joint_services.tenant_model_service import get_tenant_default_model_by_type, resolve_model_config
|
||||
from api.apps.restful_apis._generation_params import resolve_llm_setting
|
||||
from common.misc_utils import thread_pool_exec
|
||||
from api.utils.api_utils import get_error_data_result, get_json_result, add_tenant_id_to_kwargs, get_result, get_request_json, server_error_response, validate_request
|
||||
from rag.app.tag import label_question
|
||||
@@ -493,7 +494,7 @@ async def related_questions_embedded(tenant_id=None):
|
||||
chat_model_config = await thread_pool_exec(get_tenant_default_model_by_type, tenant_id, LLMType.CHAT)
|
||||
chat_mdl = LLMBundle(tenant_id, chat_model_config)
|
||||
|
||||
gen_conf = search_config.get("llm_setting", {"temperature": 0.9})
|
||||
gen_conf = resolve_llm_setting(search_config.get("llm_setting"))
|
||||
prompt = load_prompt("related_question")
|
||||
ans = await chat_mdl.async_chat(
|
||||
prompt,
|
||||
|
||||
@@ -27,7 +27,7 @@ from quart import Response, request
|
||||
from werkzeug.exceptions import BadRequest
|
||||
|
||||
from api.apps import current_user, login_required
|
||||
from api.apps.restful_apis._generation_params import merge_generation_config, pop_generation_config
|
||||
from api.apps.restful_apis._generation_params import merge_generation_config, pop_generation_config, resolve_llm_setting
|
||||
from api.db.joint_services.tenant_model_service import (
|
||||
get_api_key,
|
||||
get_composite_model_name_by_id,
|
||||
@@ -1176,7 +1176,7 @@ async def recommendation():
|
||||
chat_model_config = get_tenant_default_model_by_type(current_user.id, LLMType.CHAT)
|
||||
chat_mdl = LLMBundle(current_user.id, chat_model_config)
|
||||
|
||||
gen_conf = search_config.get("llm_setting", {"temperature": 0.9})
|
||||
gen_conf = resolve_llm_setting(search_config.get("llm_setting"))
|
||||
if "parameter" in gen_conf:
|
||||
del gen_conf["parameter"]
|
||||
prompt = load_prompt("related_question")
|
||||
|
||||
@@ -1206,13 +1206,16 @@ class Search(DataBaseModel):
|
||||
"top_k": 1024,
|
||||
# chat settings
|
||||
"summary": False,
|
||||
"chat_id": "",
|
||||
# Leave it here for reference, don't need to set default values
|
||||
"chat_id": "", # id of chat model in tenant_model table
|
||||
"llm_setting": {
|
||||
# "temperature": 0.1,
|
||||
# "top_p": 0.3,
|
||||
# "frequency_penalty": 0.7,
|
||||
# "presence_penalty": 0.4,
|
||||
"temperature": 0.1,
|
||||
"top_p": 0.3,
|
||||
"frequency_penalty": 0.7,
|
||||
"presence_penalty": 0.4,
|
||||
"temperature_enabled": True,
|
||||
"top_p_enabled": True,
|
||||
"frequency_penalty_enabled": True,
|
||||
"presence_penalty_enabled": True,
|
||||
},
|
||||
"chat_settingcross_languages": [],
|
||||
"highlight": False,
|
||||
|
||||
@@ -71,7 +71,9 @@ def _load_bot_api(monkeypatch, *, accessible, calls):
|
||||
return _gen()
|
||||
|
||||
_stub(monkeypatch, "quart", Response=lambda *a, **k: SimpleNamespace(headers=SimpleNamespace(add_header=lambda *aa, **kk: None)), request=SimpleNamespace())
|
||||
_stub(monkeypatch, "api.apps", AUTH_BETA="beta", login_required=lambda *_a, **_k: lambda func: func)
|
||||
_stub(monkeypatch, "api.apps", AUTH_BETA="beta", login_required=lambda *_a, **_k: lambda func: func, __path__=[])
|
||||
_stub(monkeypatch, "api.apps.restful_apis", __path__=[])
|
||||
_stub(monkeypatch, "api.apps.restful_apis._generation_params", resolve_llm_setting=lambda s: s or {"temperature": 0.1, "top_p": 0.3, "frequency_penalty": 0.7, "presence_penalty": 0.4})
|
||||
_stub(monkeypatch, "agent.canvas", Canvas=lambda *a, **k: SimpleNamespace(get_component_input_form=lambda _n: {}, get_prologue=lambda: "", get_mode=lambda: "agent"))
|
||||
_stub(monkeypatch, "api.db.db_models", APIToken=SimpleNamespace(query=lambda **_k: [SimpleNamespace(tenant_id="attacker-tenant")]))
|
||||
_stub(monkeypatch, "api.db.services.api_service", API4ConversationService=SimpleNamespace())
|
||||
|
||||
@@ -8,10 +8,7 @@ import { clearSensitiveFields } from './clear-sensitive-fields';
|
||||
* canvas export button and the version-history dialog so both emit
|
||||
* the same structure.
|
||||
*/
|
||||
export const downloadDsl = (
|
||||
dsl: DSL | Record<string, any>,
|
||||
title: string,
|
||||
) => {
|
||||
export const downloadDsl = (dsl: DSL | Record<string, any>, title: string) => {
|
||||
const sanitizedDsl = clearSensitiveFields(dsl);
|
||||
downloadJsonFile(
|
||||
{ ...sanitizedDsl, globals: { ...(sanitizedDsl.globals ?? {}) } },
|
||||
|
||||
@@ -130,7 +130,9 @@ export function GeneralForm() {
|
||||
ownerTenantId={useKnowledgeBaseContext().knowledgeBase?.tenant_id}
|
||||
></EmbeddingModelItem>
|
||||
<PageRankFormField></PageRankFormField>
|
||||
<CompilationTemplateFormField horizontal={true}></CompilationTemplateFormField>
|
||||
<CompilationTemplateFormField
|
||||
horizontal={true}
|
||||
></CompilationTemplateFormField>
|
||||
|
||||
<TagItems></TagItems>
|
||||
</>
|
||||
|
||||
@@ -160,10 +160,10 @@ const SearchSetting: React.FC<SearchSettingProps> = ({
|
||||
: undefined,
|
||||
},
|
||||
},
|
||||
temperatureEnabled: llm_setting?.temperature !== undefined,
|
||||
topPEnabled: llm_setting?.top_p !== undefined,
|
||||
presencePenaltyEnabled: llm_setting?.presence_penalty !== undefined,
|
||||
frequencyPenaltyEnabled: llm_setting?.frequency_penalty !== undefined,
|
||||
temperatureEnabled: llm_setting?.temperature_enabled ?? true,
|
||||
topPEnabled: llm_setting?.top_p_enabled ?? true,
|
||||
presencePenaltyEnabled: llm_setting?.presence_penalty_enabled ?? true,
|
||||
frequencyPenaltyEnabled: llm_setting?.frequency_penalty_enabled ?? true,
|
||||
});
|
||||
}, [data, search_config, llm_setting, formMethods, descriptionDefaultValue]);
|
||||
|
||||
@@ -273,10 +273,6 @@ const SearchSetting: React.FC<SearchSettingProps> = ({
|
||||
frequencyPenaltyEnabled?: boolean;
|
||||
maxTokensEnabled?: boolean;
|
||||
};
|
||||
void _temperatureEnabled;
|
||||
void _topPEnabled;
|
||||
void _presencePenaltyEnabled;
|
||||
void _frequencyPenaltyEnabled;
|
||||
void _maxTokensEnabled;
|
||||
const {
|
||||
llm_setting,
|
||||
@@ -292,6 +288,10 @@ const SearchSetting: React.FC<SearchSettingProps> = ({
|
||||
top_p: llm_setting.top_p,
|
||||
frequency_penalty: llm_setting.frequency_penalty,
|
||||
presence_penalty: llm_setting.presence_penalty,
|
||||
temperature_enabled: _temperatureEnabled,
|
||||
top_p_enabled: _topPEnabled,
|
||||
frequency_penalty_enabled: _frequencyPenaltyEnabled,
|
||||
presence_penalty_enabled: _presencePenaltyEnabled,
|
||||
} as IllmSettingProps;
|
||||
const referenceMetadata = other_config.reference_metadata;
|
||||
const normalizedReferenceMetadata = referenceMetadata
|
||||
|
||||
@@ -155,6 +155,10 @@ export interface IllmSettingProps {
|
||||
top_p?: number;
|
||||
frequency_penalty?: number;
|
||||
presence_penalty?: number;
|
||||
temperature_enabled?: boolean;
|
||||
top_p_enabled?: boolean;
|
||||
frequency_penalty_enabled?: boolean;
|
||||
presence_penalty_enabled?: boolean;
|
||||
}
|
||||
interface IllmSettingEnableProps {
|
||||
temperatureEnabled?: boolean;
|
||||
|
||||
Reference in New Issue
Block a user