InstaFuel_Chatbot_public/scripts/token_analysis.py
2026-04-21 11:57:37 +05:30

149 lines
5.2 KiB
Python

import asyncio
import os
import sys
from pathlib import Path
# Add src to path
sys.path.append(str(Path(__file__).resolve().parents[1] / "src"))
from dotenv import load_dotenv
from database.vector import QdrantManager
from llm.models import Model
from query_modes.query_engine import QueryEngine
from utils import setup_logging
load_dotenv()
setup_logging(verbose=False)
def build_query_engine():
qdrant_manager = QdrantManager(
url=os.getenv("QDRANT_URL"),
api_key=os.getenv("QDRANT_API_KEY"),
host=os.getenv("QDRANT_HOST"),
)
embed_base_url = os.getenv("EMBED_BASE_URL", "http://127.0.0.1:11434/v1")
embed_api_key = os.getenv("EMBED_API_KEY", "ollama")
dense_embedding_model = Model(
base_url=embed_base_url,
api_key=embed_api_key,
model_name="qwen3-embedding:0.6b",
provider="ollama",
embed_model=True,
use_provider_prefix=False,
)
sparse_embedding_model = Model(
base_url=embed_base_url,
api_key=embed_api_key,
model_name="qwen3-embedding:0.6b",
provider="ollama",
embed_model=True,
use_provider_prefix=False,
)
chat_model = Model(
model_name=os.environ.get("MODEL_NAME", "gpt-3.5-turbo"),
provider=os.environ.get("MODEL_PROVIDER", "openai"),
use_provider_prefix=False,
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
api_key=os.getenv("GEMINI_API_KEY"),
)
return QueryEngine(
chat_model, dense_embedding_model, sparse_embedding_model, qdrant_manager
)
TEST_QUERIES = [
("Hello there!", "greeting"),
("What are your shipping policies?", "faq"),
("I want to return my order #12345", "complaint"),
("Where is my order #12345?", "order_status"),
("Do you have any protein powder?", "product_enquiry"),
("Recommend me a pre-workout for beginners", "recommendation_request"),
("Thanks, bye!", "goodbye"),
("What is the capital of France?", "general_question"),
]
async def run_analysis():
print("Initializing Query Engine...")
try:
engine = build_query_engine()
except Exception as e:
print(f"Failed to initialize engine: {e}")
return
results: dict[str, dict[str, list[int]]] = {}
print("\nRunning token analysis...")
for query, expected_mode in TEST_QUERIES:
print(f"Testing: {query} ({expected_mode})")
try:
# We use get_response_with_suggestions to get full usage
# We pass an empty history for independent tests
response = engine.get_response_with_suggestions(query, conv_history=[])
usage = response.get("token_usage") or {}
total = (usage.get("total") or {}).get("total_tokens", 0)
stages = usage.get("stages", {})
mode_result = results.setdefault(expected_mode, {"total": [], "stages": {}})
mode_result["total"].append(total)
for stage, metrics in stages.items():
stage_totals = mode_result["stages"].setdefault(stage, [])
stage_totals.append(metrics.get("total_tokens", 0))
stage_breakdown = " | ".join(
f"{stage}: {metrics.get('total_tokens', 0)}"
for stage, metrics in stages.items()
)
print(
f" -> Tokens: {total} (Prompt: {(usage.get('total') or {}).get('prompt_tokens', 0)}, Completion: {(usage.get('total') or {}).get('completion_tokens', 0)})"
)
if stage_breakdown:
print(f" Stage breakdown: {stage_breakdown}")
except Exception as e:
print(f" -> Error: {e}")
await asyncio.sleep(10)
print("\n--- Analysis Report ---")
overall_tokens = []
overall_stage_totals: dict[str, list[int]] = {}
for mode, data in results.items():
tokens = data.get("total", [])
if tokens:
avg = sum(tokens) / len(tokens)
print(f"Mode: {mode:<25} | Avg Tokens: {avg:.2f} | Samples: {len(tokens)}")
overall_tokens.extend(tokens)
stage_summary = []
for stage, samples in data.get("stages", {}).items():
if not samples:
continue
stage_avg = sum(samples) / len(samples)
stage_summary.append(f"{stage}~{stage_avg:.0f}")
overall_stage_totals.setdefault(stage, []).extend(samples)
if stage_summary:
print(f" Stage avgs: {', '.join(stage_summary)}")
else:
print(f"Mode: {mode:<25} | No successful samples")
if overall_tokens:
overall_avg = sum(overall_tokens) / len(overall_tokens)
print(f"\nOverall Average Tokens per Turn: {overall_avg:.2f}")
if overall_stage_totals:
print("Stage contribution averages:")
for stage, samples in overall_stage_totals.items():
avg = sum(samples) / len(samples)
print(f" - {stage}: {avg:.2f}")
if __name__ == "__main__":
asyncio.run(run_analysis())