mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
build: extract <think>..</think> block for amazon deepseek r1 and put in reasoning_content
This commit is contained in:
parent
1dfdad1707
commit
9470f57e86
9 changed files with 189 additions and 14 deletions
|
|
@ -887,6 +887,9 @@ from .llms.bedrock.chat.invoke_transformations.amazon_cohere_transformation impo
|
|||
from .llms.bedrock.chat.invoke_transformations.amazon_llama_transformation import (
|
||||
AmazonLlamaConfig,
|
||||
)
|
||||
from .llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation import (
|
||||
AmazonDeepSeekR1Config,
|
||||
)
|
||||
from .llms.bedrock.chat.invoke_transformations.amazon_mistral_transformation import (
|
||||
AmazonMistralConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Choices
|
||||
from litellm.types.utils import GenericStreamingChunk as GChunk
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
from litellm.types.utils import ModelResponse, ModelResponseStream, Usage
|
||||
from litellm.utils import CustomStreamWrapper, get_secret
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
|
|
@ -226,6 +226,14 @@ async def make_call(
|
|||
completion_stream = decoder.aiter_bytes(
|
||||
response.aiter_bytes(chunk_size=1024)
|
||||
)
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(
|
||||
response.aiter_bytes(chunk_size=1024)
|
||||
)
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model)
|
||||
completion_stream = decoder.aiter_bytes(
|
||||
|
|
@ -302,6 +310,12 @@ def make_sync_call(
|
|||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=1024))
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=1024))
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=1024))
|
||||
|
|
@ -1331,7 +1345,7 @@ class AWSEventStreamDecoder:
|
|||
except Exception as e:
|
||||
raise Exception("Received streaming error - {}".format(str(e)))
|
||||
|
||||
def _chunk_parser(self, chunk_data: dict) -> GChunk:
|
||||
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream]:
|
||||
text = ""
|
||||
is_finished = False
|
||||
finish_reason = ""
|
||||
|
|
@ -1389,7 +1403,9 @@ class AWSEventStreamDecoder:
|
|||
tool_use=None,
|
||||
)
|
||||
|
||||
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk]:
|
||||
def iter_bytes(
|
||||
self, iterator: Iterator[bytes]
|
||||
) -> Iterator[Union[GChunk, ModelResponseStream]]:
|
||||
"""Given an iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
|
|
@ -1405,7 +1421,7 @@ class AWSEventStreamDecoder:
|
|||
|
||||
async def aiter_bytes(
|
||||
self, iterator: AsyncIterator[bytes]
|
||||
) -> AsyncIterator[GChunk]:
|
||||
) -> AsyncIterator[Union[GChunk, ModelResponseStream]]:
|
||||
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
|
|
@ -1474,6 +1490,27 @@ class AmazonAnthropicClaudeStreamDecoder(AWSEventStreamDecoder):
|
|||
return self.anthropic_model_response_iterator.chunk_parser(chunk=chunk_data)
|
||||
|
||||
|
||||
class AmazonDeepSeekR1StreamDecoder(AWSEventStreamDecoder):
|
||||
def __init__(
|
||||
self,
|
||||
model: str,
|
||||
sync_stream: bool,
|
||||
) -> None:
|
||||
|
||||
super().__init__(model=model)
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.amazon_deepseek_transformation import (
|
||||
AmazonDeepseekR1ResponseIterator,
|
||||
)
|
||||
|
||||
self.deepseek_model_response_iterator = AmazonDeepseekR1ResponseIterator(
|
||||
streaming_response=None,
|
||||
sync_stream=sync_stream,
|
||||
)
|
||||
|
||||
def _chunk_parser(self, chunk_data: dict) -> Union[GChunk, ModelResponseStream]:
|
||||
return self.deepseek_model_response_iterator.chunk_parser(chunk=chunk_data)
|
||||
|
||||
|
||||
class MockResponseIterator: # for returning ai21 streaming responses
|
||||
def __init__(self, model_response, json_mode: Optional[bool] = False):
|
||||
self.model_response = model_response
|
||||
|
|
|
|||
|
|
@ -0,0 +1,128 @@
|
|||
from typing import Any, List, Optional, cast
|
||||
|
||||
from httpx import Response
|
||||
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
_parse_content_for_reasoning,
|
||||
)
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.types.llms.bedrock import AmazonDeepSeekR1StreamingResponse
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionUsageBlock,
|
||||
Choices,
|
||||
Delta,
|
||||
Message,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
from .amazon_llama_transformation import AmazonLlamaConfig
|
||||
|
||||
|
||||
class AmazonDeepSeekR1Config(AmazonLlamaConfig):
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Extract the reasoning content, and return it as a separate field in the response.
|
||||
"""
|
||||
response = super().transform_response(
|
||||
model,
|
||||
raw_response,
|
||||
model_response,
|
||||
logging_obj,
|
||||
request_data,
|
||||
messages,
|
||||
optional_params,
|
||||
litellm_params,
|
||||
encoding,
|
||||
api_key,
|
||||
json_mode,
|
||||
)
|
||||
prompt = cast(Optional[str], request_data.get("prompt"))
|
||||
message_content = cast(
|
||||
Optional[str], cast(Choices, response.choices[0]).message.get("content")
|
||||
)
|
||||
if prompt and prompt.strip().endswith("<think>") and message_content:
|
||||
message_content_with_reasoning_token = "<think>" + message_content
|
||||
reasoning, content = _parse_content_for_reasoning(
|
||||
message_content_with_reasoning_token
|
||||
)
|
||||
provider_specific_fields = (
|
||||
cast(Choices, response.choices[0]).message.provider_specific_fields
|
||||
or {}
|
||||
)
|
||||
if reasoning:
|
||||
provider_specific_fields["reasoning_content"] = reasoning
|
||||
|
||||
message = Message(
|
||||
**{
|
||||
**cast(Choices, response.choices[0]).message.model_dump(),
|
||||
"content": content,
|
||||
"provider_specific_fields": provider_specific_fields,
|
||||
}
|
||||
)
|
||||
cast(Choices, response.choices[0]).message = message
|
||||
return response
|
||||
|
||||
|
||||
class AmazonDeepseekR1ResponseIterator(BaseModelResponseIterator):
|
||||
def __init__(self, streaming_response: Any, sync_stream: bool) -> None:
|
||||
super().__init__(streaming_response=streaming_response, sync_stream=sync_stream)
|
||||
self.has_finished_thinking = False
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
|
||||
"""
|
||||
Deepseek r1 starts by thinking, then it generates the response.
|
||||
"""
|
||||
try:
|
||||
typed_chunk = AmazonDeepSeekR1StreamingResponse(**chunk) # type: ignore
|
||||
if "</think>" in typed_chunk["generation"]:
|
||||
self.has_finished_thinking = True
|
||||
|
||||
prompt_token_count = typed_chunk.get("prompt_token_count") or 0
|
||||
generation_token_count = typed_chunk.get("generation_token_count") or 0
|
||||
usage = ChatCompletionUsageBlock(
|
||||
prompt_tokens=prompt_token_count,
|
||||
completion_tokens=generation_token_count,
|
||||
total_tokens=prompt_token_count + generation_token_count,
|
||||
)
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=typed_chunk["stop_reason"],
|
||||
delta=Delta(
|
||||
content=(
|
||||
typed_chunk["generation"]
|
||||
if self.has_finished_thinking
|
||||
else None
|
||||
),
|
||||
reasoning_content=(
|
||||
typed_chunk["generation"]
|
||||
if not self.has_finished_thinking
|
||||
else None
|
||||
),
|
||||
),
|
||||
)
|
||||
],
|
||||
usage=usage,
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -215,6 +215,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config = ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=litellm.LlmProviders(custom_llm_provider)
|
||||
)
|
||||
|
||||
# get config from model, custom llm provider
|
||||
headers = provider_config.validate_environment(
|
||||
api_key=api_key,
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -4,7 +4,9 @@ model_list:
|
|||
model: openai/gpt-3.5-turbo
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
api_base: http://0.0.0.0:8090
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["prometheus"]
|
||||
# custom_prometheus_metadata_labels: ["metadata.foo", "metadata.bar"]
|
||||
- model_name: deepseek-r1
|
||||
litellm_params:
|
||||
model: bedrock/deepseek_r1/arn:aws:bedrock:us-west-2:888602223428:imported-model/bnnr6463ejgf
|
||||
- model_name: deepseek-r1-api
|
||||
litellm_params:
|
||||
model: deepseek/deepseek-reasoner
|
||||
|
|
|
|||
|
|
@ -690,15 +690,10 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
if user_api_key is None:
|
||||
return
|
||||
|
||||
verbose_proxy_logger.info("ENTERS FAILURE LOG EVENT")
|
||||
|
||||
## decrement call count if call failed
|
||||
if CommonProxyErrors.max_parallel_request_limit_reached.value in str(
|
||||
kwargs["exception"]
|
||||
):
|
||||
verbose_proxy_logger.info(
|
||||
"IGNORE FAILED CALLS DUE TO MAX LIMIT BEING REACHED"
|
||||
)
|
||||
pass # ignore failed calls due to max limit being reached
|
||||
else:
|
||||
# ------------
|
||||
|
|
|
|||
|
|
@ -413,3 +413,10 @@ class BedrockRerankRequest(TypedDict):
|
|||
queries: List[BedrockRerankQuery]
|
||||
rerankingConfiguration: BedrockRerankConfiguration
|
||||
sources: List[BedrockRerankSource]
|
||||
|
||||
|
||||
class AmazonDeepSeekR1StreamingResponse(TypedDict):
|
||||
generation: str
|
||||
generation_token_count: int
|
||||
stop_reason: Optional[str]
|
||||
prompt_token_count: int
|
||||
|
|
|
|||
|
|
@ -6138,6 +6138,7 @@ class ProviderConfigManager:
|
|||
bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(
|
||||
model
|
||||
)
|
||||
|
||||
if bedrock_route == "converse" or bedrock_route == "converse_like":
|
||||
return litellm.AmazonConverseConfig()
|
||||
elif bedrock_invoke_provider == "amazon": # amazon titan llms
|
||||
|
|
@ -6152,6 +6153,8 @@ class ProviderConfigManager:
|
|||
return litellm.AmazonCohereConfig()
|
||||
elif bedrock_invoke_provider == "mistral": # mistral models on bedrock
|
||||
return litellm.AmazonMistralConfig()
|
||||
elif bedrock_invoke_provider == "deepseek_r1": # deepseek models on bedrock
|
||||
return litellm.AmazonDeepSeekR1Config()
|
||||
else:
|
||||
return litellm.AmazonInvokeConfig()
|
||||
return litellm.OpenAIGPTConfig()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue