build: extract <think>..</think> block for amazon deepseek r1 and put in reasoning_content

This commit is contained in:
Krrish Dholakia 2025-02-19 21:10:38 -08:00
parent 1dfdad1707
commit 9470f57e86
9 changed files with 189 additions and 14 deletions

View file

@ -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,
)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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:
# ------------

View file

@ -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

View file

@ -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()