mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
feat: enhance compatibility with more LLM models (#905)
* add try-except to avoid retry when using completion_cost * feat: add JSON prompt injection and think tag removal handling * refactor: simplify cost handling in LiteLLM backend and clean up style * docs: clarify purpose of reasoning_think_rm in LLMSettings * refactor: remove unused *args and redundant cost assignment --------- Co-authored-by: Young <afe.young@gmail.com>
This commit is contained in:
+40
-34
@@ -330,6 +330,7 @@ class APIBackend(ABC):
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> str | list[list[float]]:
|
||||
"""This function to share operation between embedding and chat completion"""
|
||||
assert not (chat_completion and embedding), "chat_completion and embedding cannot be True at the same time"
|
||||
max_retry = LLM_SETTINGS.max_retry if LLM_SETTINGS.max_retry is not None else max_retry
|
||||
timeout_count = 0
|
||||
@@ -397,38 +398,31 @@ class APIBackend(ABC):
|
||||
error_message = f"Failed to create chat completion after {max_retry} retries."
|
||||
raise RuntimeError(error_message)
|
||||
|
||||
def _create_chat_completion_add_json_in_prompt(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
add_json_in_prompt: bool = False,
|
||||
json_mode: bool = False,
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
) -> tuple[str, str | None]:
|
||||
def _add_json_in_prompt(self, messages: list[dict[str, Any]]) -> None:
|
||||
"""
|
||||
add json related content in the prompt if add_json_in_prompt is True
|
||||
"""
|
||||
if json_mode and add_json_in_prompt:
|
||||
for message in messages[::-1]:
|
||||
message["content"] = message["content"] + "\nPlease respond in json format."
|
||||
if message["role"] == LLM_SETTINGS.system_prompt_role:
|
||||
# NOTE: assumption: systemprompt is always the first message
|
||||
break
|
||||
return self._create_chat_completion_inner_function(messages=messages, json_mode=json_mode, *args, **kwargs) # type: ignore[misc]
|
||||
for message in messages[::-1]:
|
||||
message["content"] = message["content"] + "\nPlease respond in json format."
|
||||
if message["role"] == LLM_SETTINGS.system_prompt_role:
|
||||
# NOTE: assumption: systemprompt is always the first message
|
||||
break
|
||||
|
||||
def _create_chat_completion_auto_continue(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
*args: Any,
|
||||
json_mode: bool = False,
|
||||
chat_cache_prefix: str = "",
|
||||
seed: Optional[int] = None,
|
||||
json_target_type: Optional[str] = None,
|
||||
add_json_in_prompt: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
"""
|
||||
Call the chat completion function and automatically continue the conversation if the finish_reason is length.
|
||||
"""
|
||||
|
||||
# 0) return directly if cache is hit
|
||||
if seed is None and LLM_SETTINGS.use_auto_chat_cache_seed_gen:
|
||||
seed = LLM_CACHE_SEED_GEN.get_next_seed()
|
||||
input_content_json = json.dumps(messages)
|
||||
@@ -443,31 +437,43 @@ class APIBackend(ABC):
|
||||
logger.info(f"{LogColors.CYAN}Response:{cache_result}{LogColors.END}", tag="llm_messages")
|
||||
return cache_result
|
||||
|
||||
# 1) get a full response
|
||||
all_response = ""
|
||||
new_messages = deepcopy(messages)
|
||||
# Loop to get a full response
|
||||
try_n = 6
|
||||
for _ in range(try_n): # for some long code, 3 times may not enough for reasoning models
|
||||
if "json_mode" in kwargs:
|
||||
del kwargs["json_mode"]
|
||||
response, finish_reason = self._create_chat_completion_add_json_in_prompt(
|
||||
new_messages, json_mode=json_mode, *args, **kwargs
|
||||
) # type: ignore[misc]
|
||||
if json_mode and add_json_in_prompt:
|
||||
self._add_json_in_prompt(new_messages)
|
||||
response, finish_reason = self._create_chat_completion_inner_function(
|
||||
messages=new_messages,
|
||||
json_mode=json_mode,
|
||||
**kwargs,
|
||||
)
|
||||
all_response += response
|
||||
if finish_reason is None or finish_reason != "length":
|
||||
if json_mode:
|
||||
try:
|
||||
json.loads(all_response)
|
||||
except:
|
||||
match = re.search(r"```json(.*)```", all_response, re.DOTALL)
|
||||
all_response = match.groups()[0] if match else all_response
|
||||
json.loads(all_response)
|
||||
if json_target_type is not None:
|
||||
TypeAdapter(json_target_type).validate_json(all_response)
|
||||
if self.dump_chat_cache:
|
||||
self.cache.chat_set(input_content_json, all_response)
|
||||
return all_response
|
||||
break # we get a full response now.
|
||||
new_messages.append({"role": "assistant", "content": response})
|
||||
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
|
||||
else:
|
||||
raise RuntimeError(f"Failed to continue the conversation after {try_n} retries.")
|
||||
|
||||
# 2) refine the response and return
|
||||
if LLM_SETTINGS.reasoning_think_rm:
|
||||
match = re.search(r"<think>(.*?)</think>(.*)", all_response, re.DOTALL)
|
||||
_, all_response = match.groups() if match else ("", all_response)
|
||||
|
||||
if json_mode:
|
||||
try:
|
||||
json.loads(all_response)
|
||||
except json.decoder.JSONDecodeError:
|
||||
match = re.search(r"```json(.*)```", all_response, re.DOTALL)
|
||||
all_response = match.groups()[0] if match else all_response
|
||||
json.loads(all_response)
|
||||
if json_target_type is not None:
|
||||
TypeAdapter(json_target_type).validate_json(all_response)
|
||||
if self.dump_chat_cache:
|
||||
self.cache.chat_set(input_content_json, all_response)
|
||||
return all_response
|
||||
|
||||
def _create_embedding_with_cache(
|
||||
self, input_content_list: list[str], *args: Any, **kwargs: Any
|
||||
|
||||
@@ -135,11 +135,16 @@ class LiteLLMAPIBackend(APIBackend):
|
||||
logger.info(f"{LogColors.BLUE}assistant:{LogColors.END} {finish_reason_str}\n{content}", tag="llm_messages")
|
||||
|
||||
global ACC_COST
|
||||
cost = completion_cost(model=model, messages=messages, completion=content)
|
||||
ACC_COST += cost
|
||||
logger.info(
|
||||
f"Current Cost: ${float(cost):.10f}; Accumulated Cost: ${float(ACC_COST):.10f}; {finish_reason=}",
|
||||
)
|
||||
try:
|
||||
cost = completion_cost(model=model, messages=messages, completion=content)
|
||||
except Exception as e:
|
||||
logger.warning(f"Cost calculation failed for model {model}: {e}. Skip cost statistics.")
|
||||
else:
|
||||
ACC_COST += cost
|
||||
logger.info(
|
||||
f"Current Cost: ${float(cost):.10f}; Accumulated Cost: ${float(ACC_COST):.10f}; {finish_reason=}",
|
||||
)
|
||||
|
||||
prompt_tokens = token_counter(model=model, messages=messages)
|
||||
completion_tokens = token_counter(model=model, text=content)
|
||||
logger.log_object(
|
||||
|
||||
@@ -17,6 +17,13 @@ class LLMSettings(ExtendedBaseSettings):
|
||||
|
||||
reasoning_effort: Literal["low", "medium", "high"] | None = None
|
||||
|
||||
# Handling format
|
||||
reasoning_think_rm: bool = False
|
||||
"""
|
||||
Some LLMs include <think>...</think> tags in their responses, which can interfere with the main output.
|
||||
Set reasoning_think_rm to True to remove any <think>...</think> content from responses.
|
||||
"""
|
||||
|
||||
# TODO: most of the settings are only used on deprec.DeprecBackend.
|
||||
# So they should move the settings to that folder.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user