mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-21 19:21:03 +02:00
Break out enums for different model types
This commit is contained in:
parent
88768a791b
commit
b6dcc484fc
1 changed files with 37 additions and 11 deletions
|
|
@ -8,6 +8,7 @@ import boto3
|
||||||
import json
|
import json
|
||||||
from prometheus_client import Histogram
|
from prometheus_client import Histogram
|
||||||
import os
|
import os
|
||||||
|
import enum
|
||||||
|
|
||||||
from .... schema import TextCompletionRequest, TextCompletionResponse, Error
|
from .... schema import TextCompletionRequest, TextCompletionResponse, Error
|
||||||
from .... schema import text_completion_request_queue
|
from .... schema import text_completion_request_queue
|
||||||
|
|
@ -33,6 +34,15 @@ default_session_token = os.getenv("AWS_SESSION_TOKEN", None)
|
||||||
default_profile = os.getenv("AWS_PROFILE", None)
|
default_profile = os.getenv("AWS_PROFILE", None)
|
||||||
default_region = os.getenv("AWS_DEFAULT_REGION", None)
|
default_region = os.getenv("AWS_DEFAULT_REGION", None)
|
||||||
|
|
||||||
|
# Variant API handling depends on the model type
|
||||||
|
class ModelVariant(enum.Enum):
|
||||||
|
MISTRAL = enum.auto() # Mistral
|
||||||
|
META = enum.auto() # Llama 3.1
|
||||||
|
ANTHROPIC = enum.auto() # Anthropic
|
||||||
|
AI21 = enum.auto() # Jamba
|
||||||
|
COHERE = enum.auto() # Cohere
|
||||||
|
DEFAULT = enum.auto() # Default (use Mistral)
|
||||||
|
|
||||||
class Processor(ConsumerProducer):
|
class Processor(ConsumerProducer):
|
||||||
|
|
||||||
def __init__(self, **params):
|
def __init__(self, **params):
|
||||||
|
|
@ -94,6 +104,7 @@ class Processor(ConsumerProducer):
|
||||||
)
|
)
|
||||||
|
|
||||||
self.model = model
|
self.model = model
|
||||||
|
self.variant = self.determine_model(self.model)
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.max_output = max_output
|
self.max_output = max_output
|
||||||
|
|
||||||
|
|
@ -109,6 +120,21 @@ class Processor(ConsumerProducer):
|
||||||
|
|
||||||
print("Initialised", flush=True)
|
print("Initialised", flush=True)
|
||||||
|
|
||||||
|
def determine_model(self, model):
|
||||||
|
|
||||||
|
if self.model.startswith("mistral"):
|
||||||
|
return ModelVariant.MISTRAL
|
||||||
|
elif self.model.startswith("meta"):
|
||||||
|
return ModelVariant.META
|
||||||
|
elif self.model.startswith("anthropic"):
|
||||||
|
return ModelVariant.ANTHROPIC
|
||||||
|
elif self.model.startswith("ai21"):
|
||||||
|
return ModelVariant.AI21
|
||||||
|
elif self.model.startswith("cohere"):
|
||||||
|
return ModelVariant.COHERE
|
||||||
|
else:
|
||||||
|
return ModelVariant.DEFAULT
|
||||||
|
|
||||||
async def handle(self, msg):
|
async def handle(self, msg):
|
||||||
|
|
||||||
v = msg.value()
|
v = msg.value()
|
||||||
|
|
@ -124,7 +150,7 @@ class Processor(ConsumerProducer):
|
||||||
try:
|
try:
|
||||||
|
|
||||||
# Mistral Input Format
|
# Mistral Input Format
|
||||||
if self.model.startswith("mistral"):
|
if self.variant == ModelVariant.MISTRAL:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
"max_tokens": self.max_output,
|
"max_tokens": self.max_output,
|
||||||
|
|
@ -134,7 +160,7 @@ class Processor(ConsumerProducer):
|
||||||
})
|
})
|
||||||
|
|
||||||
# Llama 3.1 Input Format
|
# Llama 3.1 Input Format
|
||||||
elif self.model.startswith("meta"):
|
elif self.variant == ModelVariant.META:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
"max_gen_len": self.max_output,
|
"max_gen_len": self.max_output,
|
||||||
|
|
@ -143,7 +169,7 @@ class Processor(ConsumerProducer):
|
||||||
})
|
})
|
||||||
|
|
||||||
# Anthropic Input Format
|
# Anthropic Input Format
|
||||||
elif self.model.startswith("anthropic"):
|
elif self.variant == ModelVariant.ANTHROPIC:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"anthropic_version": "bedrock-2023-05-31",
|
"anthropic_version": "bedrock-2023-05-31",
|
||||||
"max_tokens": self.max_output,
|
"max_tokens": self.max_output,
|
||||||
|
|
@ -163,7 +189,7 @@ class Processor(ConsumerProducer):
|
||||||
})
|
})
|
||||||
|
|
||||||
# Jamba Input Format
|
# Jamba Input Format
|
||||||
elif self.model.startswith("ai21"):
|
elif self.variant == ModelVariant.AI21:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"max_tokens": self.max_output,
|
"max_tokens": self.max_output,
|
||||||
"temperature": self.temperature,
|
"temperature": self.temperature,
|
||||||
|
|
@ -177,14 +203,14 @@ class Processor(ConsumerProducer):
|
||||||
})
|
})
|
||||||
|
|
||||||
# Cohere Input Format
|
# Cohere Input Format
|
||||||
elif self.model.startswith("cohere"):
|
elif self.variant == ModelVariant.COHERE:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"max_tokens": self.max_output,
|
"max_tokens": self.max_output,
|
||||||
"temperature": self.temperature,
|
"temperature": self.temperature,
|
||||||
"message": prompt
|
"message": prompt
|
||||||
})
|
})
|
||||||
|
|
||||||
# Use Mistral format as defualt
|
# Use Mistral format as default
|
||||||
else:
|
else:
|
||||||
promptbody = json.dumps({
|
promptbody = json.dumps({
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
|
|
@ -204,29 +230,29 @@ class Processor(ConsumerProducer):
|
||||||
)
|
)
|
||||||
|
|
||||||
# Mistral Response Structure
|
# Mistral Response Structure
|
||||||
if self.model.startswith("mistral"):
|
if self.variant == ModelVariant.MISTRAL:
|
||||||
response_body = json.loads(response.get("body").read())
|
response_body = json.loads(response.get("body").read())
|
||||||
outputtext = response_body['outputs'][0]['text']
|
outputtext = response_body['outputs'][0]['text']
|
||||||
|
|
||||||
# Claude Response Structure
|
# Claude Response Structure
|
||||||
elif self.model.startswith("anthropic"):
|
elif self.variant == ModelVariant.ANTHROPIC:
|
||||||
model_response = json.loads(response["body"].read())
|
model_response = json.loads(response["body"].read())
|
||||||
outputtext = model_response['content'][0]['text']
|
outputtext = model_response['content'][0]['text']
|
||||||
|
|
||||||
# Llama 3.1 Response Structure
|
# Llama 3.1 Response Structure
|
||||||
elif self.model.startswith("meta"):
|
elif self.variant == ModelVariant.META:
|
||||||
model_response = json.loads(response["body"].read())
|
model_response = json.loads(response["body"].read())
|
||||||
outputtext = model_response["generation"]
|
outputtext = model_response["generation"]
|
||||||
|
|
||||||
# Jamba Response Structure
|
# Jamba Response Structure
|
||||||
elif self.model.startswith("ai21"):
|
elif self.variant == ModelVariant.AI21:
|
||||||
content = response['body'].read()
|
content = response['body'].read()
|
||||||
content_str = content.decode('utf-8')
|
content_str = content.decode('utf-8')
|
||||||
content_json = json.loads(content_str)
|
content_json = json.loads(content_str)
|
||||||
outputtext = content_json['choices'][0]['message']['content']
|
outputtext = content_json['choices'][0]['message']['content']
|
||||||
|
|
||||||
# Cohere Input Format
|
# Cohere Input Format
|
||||||
elif self.model.startswith("cohere"):
|
elif self.variant == ModelVariant.COHERE:
|
||||||
content = response['body'].read()
|
content = response['body'].read()
|
||||||
content_str = content.decode('utf-8')
|
content_str = content.decode('utf-8')
|
||||||
content_json = json.loads(content_str)
|
content_json = json.loads(content_str)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue