mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 01:23:55 +08:00
@@ -21,22 +21,22 @@ class LLMProvider(LLMProviderBase):
|
|||||||
self.timeout = int(timeout) if timeout else 300
|
self.timeout = int(timeout) if timeout else 300
|
||||||
|
|
||||||
param_defaults = {
|
param_defaults = {
|
||||||
"max_tokens": (500, int),
|
"max_tokens": int,
|
||||||
"temperature": (0.7, lambda x: round(float(x), 1)),
|
"temperature": lambda x: round(float(x), 1),
|
||||||
"top_p": (1.0, lambda x: round(float(x), 1)),
|
"top_p": lambda x: round(float(x), 1),
|
||||||
"frequency_penalty": (0, lambda x: round(float(x), 1)),
|
"frequency_penalty": lambda x: round(float(x), 1),
|
||||||
}
|
}
|
||||||
|
|
||||||
for param, (default, converter) in param_defaults.items():
|
for param, converter in param_defaults.items():
|
||||||
value = config.get(param)
|
value = config.get(param)
|
||||||
try:
|
try:
|
||||||
setattr(
|
setattr(
|
||||||
self,
|
self,
|
||||||
param,
|
param,
|
||||||
converter(value) if value not in (None, "") else default,
|
converter(value) if value not in (None, "") else None,
|
||||||
)
|
)
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
setattr(self, param, default)
|
setattr(self, param, None)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}"
|
f"意图识别参数初始化: {self.temperature}, {self.max_tokens}, {self.top_p}, {self.frequency_penalty}"
|
||||||
@@ -59,17 +59,25 @@ class LLMProvider(LLMProviderBase):
|
|||||||
try:
|
try:
|
||||||
dialogue = self.normalize_dialogue(dialogue)
|
dialogue = self.normalize_dialogue(dialogue)
|
||||||
|
|
||||||
responses = self.client.chat.completions.create(
|
request_params = {
|
||||||
model=self.model_name,
|
"model": self.model_name,
|
||||||
messages=dialogue,
|
"messages": dialogue,
|
||||||
stream=True,
|
"stream": True,
|
||||||
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
}
|
||||||
temperature=kwargs.get("temperature", self.temperature),
|
|
||||||
top_p=kwargs.get("top_p", self.top_p),
|
# 添加可选参数,只有当参数不为None时才添加
|
||||||
frequency_penalty=kwargs.get(
|
optional_params = {
|
||||||
"frequency_penalty", self.frequency_penalty
|
"max_tokens": kwargs.get("max_tokens", self.max_tokens),
|
||||||
),
|
"temperature": kwargs.get("temperature", self.temperature),
|
||||||
)
|
"top_p": kwargs.get("top_p", self.top_p),
|
||||||
|
"frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty),
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, value in optional_params.items():
|
||||||
|
if value is not None:
|
||||||
|
request_params[key] = value
|
||||||
|
|
||||||
|
responses = self.client.chat.completions.create(**request_params)
|
||||||
|
|
||||||
is_active = True
|
is_active = True
|
||||||
for chunk in responses:
|
for chunk in responses:
|
||||||
@@ -91,13 +99,29 @@ class LLMProvider(LLMProviderBase):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
||||||
|
|
||||||
def response_with_functions(self, session_id, dialogue, functions=None):
|
def response_with_functions(self, session_id, dialogue, functions=None, **kwargs):
|
||||||
try:
|
try:
|
||||||
dialogue = self.normalize_dialogue(dialogue)
|
dialogue = self.normalize_dialogue(dialogue)
|
||||||
|
|
||||||
stream = self.client.chat.completions.create(
|
request_params = {
|
||||||
model=self.model_name, messages=dialogue, stream=True, tools=functions
|
"model": self.model_name,
|
||||||
)
|
"messages": dialogue,
|
||||||
|
"stream": True,
|
||||||
|
"tools": functions,
|
||||||
|
}
|
||||||
|
|
||||||
|
optional_params = {
|
||||||
|
"max_tokens": kwargs.get("max_tokens", self.max_tokens),
|
||||||
|
"temperature": kwargs.get("temperature", self.temperature),
|
||||||
|
"top_p": kwargs.get("top_p", self.top_p),
|
||||||
|
"frequency_penalty": kwargs.get("frequency_penalty", self.frequency_penalty),
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, value in optional_params.items():
|
||||||
|
if value is not None:
|
||||||
|
request_params[key] = value
|
||||||
|
|
||||||
|
stream = self.client.chat.completions.create(**request_params)
|
||||||
|
|
||||||
for chunk in stream:
|
for chunk in stream:
|
||||||
if getattr(chunk, "choices", None):
|
if getattr(chunk, "choices", None):
|
||||||
|
|||||||
Reference in New Issue
Block a user