Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions app/configuration_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,7 +314,6 @@ def generate_all_combinations(data):

if experiment_count >= MAX_VALID_EXPERIMENTS:
break # Exit the loop when the limit is reached

configuration = {
**combination
}
Expand All @@ -327,6 +326,8 @@ def generate_all_combinations(data):
"embedding_model": configuration["embedding"]["model"],
"retrieval_service": configuration["retrieval"]["service"],
"retrieval_model": configuration["retrieval"]["model"],
"retrieval_model_input_token_cost": configuration["retrieval"].get("input_token_cost", 0),
"retrieval_model_output_token_cost": configuration["retrieval"].get("output_token_cost", 0),
"eval_service": configuration["evaluation"]["service"],
"eval_embedding_model": configuration["evaluation"]["embedding_model"],
"eval_retrieval_model": configuration["evaluation"]["retrieval_model"],
Expand Down Expand Up @@ -360,7 +361,9 @@ def generate_all_combinations(data):

#Calculate the inferencing price - doesn't include OpenSearch pricing
if configuration.get('gateway_enabled', False):
configuration["inferencing_cost_estimate"] += 0
configuration["inferencing_cost_estimate"] += estimate_retrieval_model_bedrock_price(bedrock_price_df, configuration, avg_prompt_length, num_prompts,
input_token_cost=configuration.get("retrieval_model_input_token_cost"),
output_token_cost=configuration.get("retrieval_model_output_token_cost"))
else:
if configuration["retrieval_service"] == "bedrock":
inferencing_price = estimate_retrieval_model_bedrock_price(bedrock_price_df, configuration, avg_prompt_length, num_prompts)
Expand Down Expand Up @@ -392,7 +395,6 @@ def generate_all_combinations(data):
configuration["directional_pricing"] = configuration["indexing_cost_estimate"] + configuration["retrieval_cost_estimate"] + configuration["inferencing_cost_estimate"] + configuration["eval_cost_estimate"]
configuration["directional_pricing"] +=configuration["directional_pricing"]*0.05 #extra
configuration["directional_pricing"] = round(configuration["directional_pricing"],2)

return valid_configurations

def generate_all_combinations_in_background(execution_id: str, execution_config_data):
Expand Down
45 changes: 24 additions & 21 deletions app/price_calculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,11 +48,8 @@ def estimate_embedding_model_bedrock_price(file_path, configuration, effective_k


def estimate_retrieval_model_bedrock_price(file_path, configuration, avg_prompt_length,
num_prompts):
num_prompts, input_token_cost=None, output_token_cost=None):
try:
if configuration.get('gateway_enabled', False):
return 0

df = file_path.copy()
# return df
except Exception as e:
Expand All @@ -72,25 +69,31 @@ def estimate_retrieval_model_bedrock_price(file_path, configuration, avg_prompt_
gen_model = configuration["retrieval_model"]
n_shot_prompts = configuration["n_shot_prompts"]
k = configuration["knn_num"]

gen_model_price = df[(df['model'] == gen_model) & (df['Region'] == region)]['input_price']
if gen_model_price.empty:
logger.warning("Returning price as Zero, as model is not present in Sheet")
return 0
if input_token_cost != None and output_token_cost != None:
gen_model_input_price = float(input_token_cost)
gen_model_output_price = float(output_token_cost)
else:
gen_model_price = float(gen_model_price.values[0]) # this price is in millions of tokens
gen_model_out_price = df[(df['model'] == gen_model) & (df['Region'] == region)]['output_price']
gen_model_out_price = float(gen_model_out_price.values[0]) # this price is in millions of tokens
context_len = k * chunk_size
prompt_len = (n_shot_prompts + 1) * avg_prompt_length
total_input_tokens = (context_len + prompt_len) * num_prompts
retrieval_input_price = gen_model_price * float(total_input_tokens) / 1000000
total_output_tokens = avg_prompt_length
retrieval_output_price = gen_model_out_price * float(total_output_tokens) / 1000000
retrieval_price = retrieval_input_price + retrieval_output_price
return retrieval_price

gen_model_price_series = df[(df['model'] == gen_model) & (df['Region'] == region)]['input_price']
if gen_model_price_series.empty:
logger.warning("Returning price as Zero, as input price for model is not present in Sheet")
return 0
else:
gen_model_input_price = float(gen_model_price_series.values[0])

gen_model_out_price_series = df[(df['model'] == gen_model) & (df['Region'] == region)]['output_price']
if gen_model_out_price_series.empty:
logger.warning("Returning price as Zero, as output price for model is not present in Sheet")
return 0
else:
gen_model_output_price = float(gen_model_out_price_series.values[0])
context_len = k * chunk_size
prompt_len = (n_shot_prompts + 1) * avg_prompt_length
total_input_tokens = (context_len + prompt_len) * num_prompts
retrieval_input_price = gen_model_input_price * float(total_input_tokens) / 1000000
total_output_tokens = avg_prompt_length
retrieval_output_price = gen_model_output_price * float(total_output_tokens) / 1000000
retrieval_price = retrieval_input_price + retrieval_output_price
return retrieval_price

def estimate_fargate_price(total_time, vpc=8, mem=16):
#Fargate pricing
Expand Down
Empty file.
42 changes: 42 additions & 0 deletions cost_compute_handler/fargate/base_task_processor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
import json
from abc import ABC, abstractmethod
import sys
from flotorch_core.logger.global_logger import get_logger

logger = get_logger()

class BaseFargateTaskProcessor(ABC):
"""
Abstract base class for Fargate task processors.
"""

def __init__(self, input_data: dict):
"""
Initializes the task processor with input data.
Args:
input_data (dict): The input data for the task.
"""
self.input_data = input_data

@abstractmethod
def process(self):
"""
Abstract method to be implemented by subclasses for processing tasks.
"""
raise NotImplementedError("Subclasses must implement the process method.")

def send_task_success(self):
"""
Sends task success signal.
"""
print(json.dumps({"status": "success", "output": "Retrieval completed successfully"}))
sys.exit(0)

def send_task_failure(self, error_message: str):
"""
Sends task failure signal.
Args:
error_message (str): The error message to send.
"""
print(json.dumps({"status": "failure", "output": "Retrieval Failed", "error": error_message}))
sys.exit(1)
185 changes: 185 additions & 0 deletions cost_compute_handler/fargate/cost_compute_processor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
import json
import math
import os
from fargate.base_task_processor import BaseFargateTaskProcessor
from decimal import Decimal
from .pricing import compute_actual_price_breakdown, calculate_experiment_duration
from flotorch_core.config.config import Config
from flotorch_core.config.env_config_provider import EnvConfigProvider
from flotorch_core.storage.db.postgresdb import PostgresDB
from flotorch_core.storage.db.dynamodb import DynamoDB
from flotorch_core.logger.global_logger import get_logger


MILLION = 1_000_000
THOUSAND = 1_000
SECONDS_IN_MINUTE = 60
MINUTES_IN_HOUR = 60
HOURS_IN_DAY = 24
DAYS_IN_MONTH = 30

logger = get_logger()
env_config_provider = EnvConfigProvider()
config = Config(env_config_provider)
db_type = config.get_db_type()

def fetch_data_from_db(table_name, key, value, index_name=None):
"""
Fetch items with the specified key and value from DynamoDB or postgres based on db type.
"""
try:
if db_type == "DYNAMODB":
db_client = DynamoDB(table_name=table_name, region_name=config.get_region())
if index_name:
items = db_client.read(keys={index_name: value})
else:
items = db_client.read(keys={key: value})

elif db_type == "POSTGRESDB":
db_client = PostgresDB(dbname=config.get_postgres_db(), user=config.get_postgres_user(), password=config.get_postgres_password(), table_name=table_name, host=config.get_postgres_host(), port=config.get_postgres_port())
items = db_client.read(key={key: value})
return items
except Exception as e:
logger.error(f"Error fetching data from DB: {e}")
raise


def validate_event(event):
"""
Validate the input event to ensure required fields are present.
"""
required_fields = ["experiment_id"]
for field in required_fields:
if field not in event:
raise ValueError(f"Missing required field: {field}")

if not isinstance(event["experiment_id"], str):
raise ValueError("'experiment_id' must be a string")


class RetrieverProcessor(BaseFargateTaskProcessor):
"""
Processor for retriever tasks in Fargate.
"""

def process(self):
logger.info("Starting retriever process.")
try:
logger.info(f"Experiment Configuration received: {self.input_data}")

# Validate input event
validate_event(self.input_data)

experiment_id = self.input_data["experiment_id"]
experiment_table = config.get_experiment_table_name()
experiment_question_metrics_table = config.get_experiment_question_metrics_table()
experiment_question_metrics_index = os.getenv("experiment_question_metrics_index")

if not experiment_table:
raise EnvironmentError("Environment variable 'experiment_table' is not set")

if not experiment_question_metrics_table:
raise EnvironmentError("Environment variable 'experiment_question_metrics_table' is not set")

# Initialize variables
total_query_embed_tokens = 0
total_answer_input_tokens = 0
total_answer_output_tokens = 0

experiment_items = fetch_data_from_db(experiment_table, 'id', experiment_id)
experiment_question_metrics_items = fetch_data_from_db(experiment_question_metrics_table, 'experiment_id', experiment_id, experiment_question_metrics_index)
total_duration = 0
indexing_time = 0
retrieval_time = 0
eval_time = 0
total_index_embed_tokens = 0

if experiment_items:
experiment = experiment_items[0]
indexing_time, retrieval_time, eval_time = calculate_experiment_duration(experiment)
indexing_time_in_min = math.ceil(indexing_time / SECONDS_IN_MINUTE)
retrieval_time_in_min = math.ceil(retrieval_time / SECONDS_IN_MINUTE)
eval_time_in_min = math.ceil(eval_time / SECONDS_IN_MINUTE)
total_duration = indexing_time + retrieval_time + eval_time
total_duration_in_min = indexing_time_in_min + retrieval_time_in_min + eval_time_in_min
logger.info(f"Experiment {experiment_id} Total Time (in minutes): {total_duration_in_min} Indexing Time: {indexing_time_in_min}, Retrieval: {retrieval_time_in_min}, Evaluation: {eval_time_in_min}")

total_index_embed_tokens = experiment.get("index_embed_tokens", 0)
total_query_embed_tokens = experiment.get("retrieval_query_embed_tokens", 0)
total_answer_input_tokens = experiment.get("retrieval_input_tokens", 1)
total_answer_output_tokens = experiment.get("retrieval_output_tokens", 1)

overall_metadata, indexing_metadata, retriever_metadata, inferencer_metadata, eval_metadata = compute_actual_price_breakdown(
experiment,
input_tokens=total_answer_input_tokens,
output_tokens=total_answer_output_tokens,
index_embed_tokens=total_index_embed_tokens,
query_embed_tokens=total_query_embed_tokens,
total_time=total_duration,
indexing_time=indexing_time,
retrieval_time=retrieval_time,
eval_time=eval_time,
experiment_question_metrics_items=experiment_question_metrics_items
)

total_cost = overall_metadata['total_cost']
indexing_cost = indexing_metadata['total_cost']
retriever_cost = retriever_metadata['total_cost']
inferencer_cost = inferencer_metadata['total_cost']
eval_cost = eval_metadata['total_cost']
logger.info(f"Experiment {experiment_id} Actual Cost (in $): {total_cost}, Indexing: {indexing_cost}, Retrieval: {retriever_cost}, Inferencing: {inferencer_cost}, Evaluation : {eval_cost}")

# Update DynamoDB with the new cost
if total_cost is None:
logger.error(f"Experiment {experiment_id} Actual Cost is None")
total_cost = 0

try:
key={"id": experiment_id}
data = {
"cost": str(total_cost),
"indexing_time": str(indexing_time),
"retrieval_time": str(retrieval_time),
"eval_time": str(eval_time),
"total_time": str(total_duration),
"indexing_cost": str(indexing_cost),
"retrieval_cost": str(retriever_cost),
"inferencing_cost": str(inferencer_cost),
"eval_cost": str(eval_cost),
"indexing_metadata": convert_floats_to_decimal(indexing_metadata),
"retriever_metadata": convert_floats_to_decimal(retriever_metadata),
"inferencer_metadata": convert_floats_to_decimal(inferencer_metadata),
"eval_metadata": convert_floats_to_decimal(eval_metadata),
"overall_metadata": convert_floats_to_decimal(overall_metadata)
}
if db_type == "DYNAMODB":
db_client = DynamoDB(table_name=experiment_table, region_name=config.get_region())
elif db_type == "POSTGRESDB":
db_client = PostgresDB(dbname=config.get_postgres_db(), user=config.get_postgres_user(), password=config.get_postgres_password(), table_name=experiment_table, host=config.get_postgres_host(), port=config.get_postgres_port())

db_client.update(key, data)

except Exception as e:
logger.error(f"Error updating DB: {e}")
raise
self.send_task_success()

except ValueError as ve:
logger.error(f"Validation error: {ve}")
self.send_task_failure(str(ve))
except EnvironmentError as ee:
logger.error(f"Environment error: {ee}")
self.send_task_failure(str(ee))
except Exception as e:
logger.error(f"Unhandled error: {e}")
self.send_task_failure(str(e))

def convert_floats_to_decimal(obj):
if isinstance(obj, float):
return Decimal(str(obj)) # Convert float to string first to prevent precision loss
elif isinstance(obj, dict):
return {k: convert_floats_to_decimal(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [convert_floats_to_decimal(i) for i in obj]
else:
return obj
Empty file.
28 changes: 28 additions & 0 deletions cost_compute_handler/fargate/handler/cost_compute/Dockerfile
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
# Use an official Python runtime as a parent image
FROM python:3.9-slim-buster

# Set the working directory in the container
WORKDIR /app

# Install FloTorch-core
RUN pip install FloTorch-core

# Copy the application code
COPY fargate/ fargate/
COPY fargate/handler/cost_compute/fargate_cost_compute_handler.py .

# Set environment variables (these can be overridden at runtime)
ENV experiment_table=your_experiment_table_name
ENV experiment_question_metrics_table=your_experiment_question_metrics_table_name
ENV experiment_question_metrics_index=your_experiment_question_metrics_index_name
ENV bedrock_limit_csv=your_bedrock_limit_csv_file
ENV s3_bucket=your_s3_bucket_name
ENV AWS_REGION=your_aws_region
ENV AWS_ACCESS_KEY_ID=your_aws_access_key_id
ENV AWS_SECRET_ACCESS_KEY=your_aws_secret_access_key
ENV AWS_SESSION_TOKEN=your_aws_session_token
ENV AWS_DEFAULT_REGION=your_aws_default_region
ENV db_type=type_of_your_database

# Command to run the main application
CMD ["python", "fargate_cost_compute_handler.py"]
Empty file.
Loading