From 3967d92b8d7a8ccb520afa6f4cb2ea7156cc5c4b Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Thu, 13 Mar 2025 14:46:46 +0000 Subject: [PATCH] Add model detection for inference profiles in US and EU --- .../model/text_completion/bedrock/llm.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/trustgraph-bedrock/trustgraph/model/text_completion/bedrock/llm.py b/trustgraph-bedrock/trustgraph/model/text_completion/bedrock/llm.py index fde5036b..caea444c 100755 --- a/trustgraph-bedrock/trustgraph/model/text_completion/bedrock/llm.py +++ b/trustgraph-bedrock/trustgraph/model/text_completion/bedrock/llm.py @@ -35,6 +35,7 @@ default_profile = os.getenv("AWS_PROFILE", None) default_region = os.getenv("AWS_DEFAULT_REGION", None) # Variant API handling depends on the model type +# FIXME: Missing, Amazon models, Deepseek class ModelVariant(enum.Enum): MISTRAL = enum.auto() # Mistral META = enum.auto() # Llama 3.1 @@ -122,6 +123,9 @@ class Processor(ConsumerProducer): def determine_model(self, model): + # FIXME: Missing, Amazon models, Deepseek + + # This set of conditions deals with normal bedrock on-demand usage if self.model.startswith("mistral"): return ModelVariant.MISTRAL elif self.model.startswith("meta"): @@ -132,8 +136,18 @@ class Processor(ConsumerProducer): return ModelVariant.AI21 elif self.model.startswith("cohere"): return ModelVariant.COHERE - else: - return ModelVariant.DEFAULT + + # The inference profiles + if self.model.startswith("us.meta"): + return ModelVariant.META + elif self.model.startswith("us.anthropic"): + return ModelVariant.ANTHROPIC + elif self.model.startswith("eu.meta"): + return ModelVariant.META + elif self.model.startswith("eu.anthropic"): + return ModelVariant.ANTHROPIC + + return ModelVariant.DEFAULT async def handle(self, msg):