diff --git a/README.md b/README.md index 1d4f2c5..8a8323a 100755 --- a/README.md +++ b/README.md @@ -71,6 +71,8 @@ source venv_agent_lab/bin/activate ```bash pip install -r requirements.txt ``` +- This install includes the dependencies used by the optional AgentRxiv web workflow. +- TensorFlow is not required for the core Agent Laboratory workflows and is not installed by default. 4. **Install pdflatex [OPTIONAL]** ```bash diff --git a/ai_lab_repo.py b/ai_lab_repo.py index d8e7c1e..2e5d697 100755 --- a/ai_lab_repo.py +++ b/ai_lab_repo.py @@ -1,6 +1,5 @@ -import PyPDF2 +import importlib import threading -from app import * from agents import * from copy import copy from pathlib import Path @@ -8,6 +7,7 @@ from common_imports import * from mlesolver import MLESolver import argparse, pickle, yaml +from pypdf import PdfReader GLOBAL_AGENTRXIV = None DEFAULT_LLM_BACKBONE = "o3-mini" @@ -16,6 +16,31 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false" +def load_agentrxiv_app(): + optional_dependencies = { + "flask": "flask", + "flask_sqlalchemy": "flask-sqlalchemy", + "sentence_transformers": "sentence-transformers", + } + missing_packages = [] + for module_name, package_name in optional_dependencies.items(): + try: + importlib.import_module(module_name) + except ModuleNotFoundError: + missing_packages.append(package_name) + + if missing_packages: + missing_list = ", ".join(sorted(missing_packages)) + raise RuntimeError( + "AgentRxiv requires additional web dependencies that are not installed: " + f"{missing_list}. Install them with `pip install -r requirements.txt`." + ) + + from app import app, run_app, update_papers_from_uploads + + return app, run_app, update_papers_from_uploads + + class LaboratoryWorkflow: def __init__(self, research_topic, openai_api_key, max_steps=100, num_papers_lit_review=5, agent_model_backbone=f"{DEFAULT_LLM_BACKBONE}", notes=list(), human_in_loop_flag=None, compile_pdf=True, mlesolver_max_steps=3, papersolver_max_steps=5, paper_index=0, except_if_fail=False, parallelized=False, lab_dir=None, lab_index=0, agentRxiv=False, agentrxiv_papers=5): """ @@ -601,13 +626,15 @@ def retrieve_full_text(self, arxiv_id): return "Paper ID not found?" @staticmethod - def read_pdf_pypdf2(pdf_path): + def read_pdf_text(pdf_path): with open(pdf_path, 'rb') as pdf_file: - reader = PyPDF2.PdfReader(pdf_file) + reader = PdfReader(pdf_file) text = '' for page_num in range(len(reader.pages)): page = reader.pages[page_num] - text += page.extract_text() + page_text = page.extract_text() + if page_text: + text += page_text return text def search_agentrxiv(self, search_query, num_papers): @@ -615,6 +642,7 @@ def search_agentrxiv(self, search_query, num_papers): url = f'http://127.0.0.1:{5000 + self.lab_index}/api/search?q={search_query}' return_str = str() try: + app, _, update_papers_from_uploads = load_agentrxiv_app() with app.app_context(): update_papers_from_uploads() response = requests.get(url) @@ -628,7 +656,7 @@ def search_agentrxiv(self, search_query, num_papers): filename = Path(f'_tmp_{self.lab_index}.pdf') response = requests.get(result['pdf_url']) filename.write_bytes(response.content) - self.pdf_text[arxiv_id] = self.read_pdf_pypdf2(f'_tmp_{self.lab_index}.pdf') + self.pdf_text[arxiv_id] = self.read_pdf_text(f'_tmp_{self.lab_index}.pdf') self.summaries[arxiv_id] = query_model( prompt=self.pdf_text[arxiv_id], system_prompt="Please provide a 5 sentence summary of this paper.", @@ -647,6 +675,7 @@ def search_agentrxiv(self, search_query, num_papers): return return_str def run_server(self, port): + _, run_app, _ = load_agentrxiv_app() run_app(port=port) @@ -886,6 +915,3 @@ def run_lab(parallel_lab_index): - - - diff --git a/app.py b/app.py index 08114ee..7a63871 100644 --- a/app.py +++ b/app.py @@ -1,11 +1,11 @@ import random, time +from functools import lru_cache from flask import Flask, render_template, request, redirect, url_for, flash, send_from_directory, jsonify from werkzeug.utils import secure_filename import os -from PyPDF2 import PdfReader +from pypdf import PdfReader from flask_sqlalchemy import SQLAlchemy -from sentence_transformers import SentenceTransformer from sklearn.metrics.pairwise import cosine_similarity import numpy as np @@ -58,8 +58,11 @@ def update_papers_from_uploads(): return #raise Exception("FAILED TO UPDATE") -# Load a pre-trained sentence transformer model -model = SentenceTransformer('all-MiniLM-L6-v2') +@lru_cache(maxsize=1) +def get_embedding_model(): + from sentence_transformers import SentenceTransformer + + return SentenceTransformer('all-MiniLM-L6-v2') @app.route('/update', methods=['GET']) def update_on_demand(): @@ -106,6 +109,7 @@ def upload(): def search(): query = request.args.get('q', '') if query: + model = get_embedding_model() papers = Paper.query.all() query_embedding = model.encode([query]) paper_texts = [paper.text for paper in papers if paper.text] @@ -123,6 +127,7 @@ def api_search(): query = request.args.get('q', '') if not query: return jsonify({'error': 'No query provided'}), 400 + model = get_embedding_model() papers = Paper.query.all() if not papers: return jsonify({'query': query, 'results': []}) @@ -168,4 +173,4 @@ def run_app(port=5000): app.run(debug=False, port=port) if __name__ == '__main__': - run_app() \ No newline at end of file + run_app() diff --git a/common_imports.py b/common_imports.py index 7d968f3..a1f5d80 100755 --- a/common_imports.py +++ b/common_imports.py @@ -34,13 +34,31 @@ # Hugging Face & Transformers import transformers + +class MissingOptionalDependency: + def __init__(self, package_name, feature_description): + self.package_name = package_name + self.feature_description = feature_description + + def __bool__(self): + return False + + def __getattr__(self, name): + raise ImportError( + f"Optional dependency '{self.package_name}' is required for {self.feature_description}. " + f"Install it separately before using '{name}'." + ) + # Deep learning frameworks import torch import torch.nn as nn import torch.optim as optim import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset, random_split -import tensorflow as tf +try: + import tensorflow as tf +except ImportError: + tf = MissingOptionalDependency("tensorflow", "TensorFlow-backed experiments") #import keras # NLP Libraries diff --git a/inference.py b/inference.py index c2ccd5f..7ae4504 100755 --- a/inference.py +++ b/inference.py @@ -1,13 +1,21 @@ +import json +import os +import time + +import anthropic import openai -import time, tiktoken from openai import OpenAI -import os, anthropic, json -import google.generativeai as genai + +from tokenization import count_text_tokens, normalize_model_name TOKENS_IN = dict() TOKENS_OUT = dict() -encoding = tiktoken.get_encoding("cl100k_base") +OPENAI_MODELS = {"gpt-4o", "gpt-4o-mini", "o1", "o1-mini", "o1-preview", "o3-mini"} +ANTHROPIC_MODELS = {"claude-3-5-sonnet"} +GEMINI_MODELS = {"gemini-1.5-pro", "gemini-2.0-pro"} +DEEPSEEK_MODELS = {"deepseek-chat"} + def curr_cost_est(): costmap_in = { @@ -21,7 +29,7 @@ def curr_cost_est(): "o3-mini": 1.10 / 1000000, } costmap_out = { - "gpt-4o": 10.00/ 1000000, + "gpt-4o": 10.00 / 1000000, "gpt-4o-mini": 0.6 / 1000000, "o1-preview": 60.00 / 1000000, "o1-mini": 12.00 / 1000000, @@ -30,14 +38,72 @@ def curr_cost_est(): "o1": 60.00 / 1000000, "o3-mini": 4.40 / 1000000, } - return sum([costmap_in[_]*TOKENS_IN[_] for _ in TOKENS_IN]) + sum([costmap_out[_]*TOKENS_OUT[_] for _ in TOKENS_OUT]) - -def query_model(model_str, prompt, system_prompt, openai_api_key=None, gemini_api_key=None, anthropic_api_key=None, tries=5, timeout=5.0, temp=None, print_cost=True, version="1.5"): - preloaded_api = os.getenv('OPENAI_API_KEY') - if openai_api_key is None and preloaded_api is not None: - openai_api_key = preloaded_api - if openai_api_key is None and anthropic_api_key is None: - raise Exception("No API key provided in query_model function") + + total_cost = 0.0 + for model_name, token_count in TOKENS_IN.items(): + if model_name in costmap_in: + total_cost += costmap_in[model_name] * token_count + for model_name, token_count in TOKENS_OUT.items(): + if model_name in costmap_out: + total_cost += costmap_out[model_name] * token_count + return total_cost + + +def _resolve_api_key(explicit_key, env_var_name): + return explicit_key or os.getenv(env_var_name) + + +def _load_genai(): + import google.generativeai as genai + + return genai + + +def _validate_credentials(model_name, openai_api_key, gemini_api_key, anthropic_api_key, deepseek_api_key): + if model_name in OPENAI_MODELS and openai_api_key is None: + raise Exception("OPENAI_API_KEY must be provided for OpenAI-backed models") + if model_name in GEMINI_MODELS and gemini_api_key is None: + raise Exception("GEMINI_API_KEY must be provided for Gemini-backed models") + if model_name in ANTHROPIC_MODELS and anthropic_api_key is None: + raise Exception("ANTHROPIC_API_KEY must be provided for Anthropic-backed models") + if model_name in DEEPSEEK_MODELS and deepseek_api_key is None: + raise Exception("DEEPSEEK_API_KEY must be provided for DeepSeek-backed models") + + +def _record_token_usage(model_name, system_prompt, prompt, answer, print_cost): + try: + TOKENS_IN.setdefault(model_name, 0) + TOKENS_OUT.setdefault(model_name, 0) + TOKENS_IN[model_name] += count_text_tokens(system_prompt + prompt, model_name) + TOKENS_OUT[model_name] += count_text_tokens(answer, model_name) + if print_cost: + print(f"Current experiment cost = ${curr_cost_est()}, ** Approximate values, may not reflect true cost") + except Exception as e: + if print_cost: + print(f"Cost approximation has an error? {e}") + + +def query_model( + model_str, + prompt, + system_prompt, + openai_api_key=None, + gemini_api_key=None, + anthropic_api_key=None, + tries=5, + timeout=5.0, + temp=None, + print_cost=True, + version="1.5", +): + model_str = normalize_model_name(model_str) + openai_api_key = _resolve_api_key(openai_api_key, "OPENAI_API_KEY") + gemini_api_key = _resolve_api_key(gemini_api_key, "GEMINI_API_KEY") + anthropic_api_key = _resolve_api_key(anthropic_api_key, "ANTHROPIC_API_KEY") + deepseek_api_key = os.getenv("DEEPSEEK_API_KEY") + + _validate_credentials(model_str, openai_api_key, gemini_api_key, anthropic_api_key, deepseek_api_key) + if openai_api_key is not None: openai.api_key = openai_api_key os.environ["OPENAI_API_KEY"] = openai_api_key @@ -45,164 +111,144 @@ def query_model(model_str, prompt, system_prompt, openai_api_key=None, gemini_ap os.environ["ANTHROPIC_API_KEY"] = anthropic_api_key if gemini_api_key is not None: os.environ["GEMINI_API_KEY"] = gemini_api_key + for _ in range(tries): try: - if model_str == "gpt-4o-mini" or model_str == "gpt4omini" or model_str == "gpt-4omini" or model_str == "gpt4o-mini": - model_str = "gpt-4o-mini" + if model_str == "gpt-4o-mini": messages = [ {"role": "system", "content": system_prompt}, - {"role": "user", "content": prompt}] + {"role": "user", "content": prompt}, + ] if version == "0.28": - if temp is None: - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages - ) - else: - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages, temperature=temp - ) + request_kwargs = {"model": model_str, "messages": messages} + if temp is not None: + request_kwargs["temperature"] = temp + completion = openai.ChatCompletion.create(**request_kwargs) else: - client = OpenAI() - if temp is None: - completion = client.chat.completions.create( - model="gpt-4o-mini-2024-07-18", messages=messages, ) - else: - completion = client.chat.completions.create( - model="gpt-4o-mini-2024-07-18", messages=messages, temperature=temp) + client = OpenAI(api_key=openai_api_key) + request_kwargs = {"model": "gpt-4o-mini-2024-07-18", "messages": messages} + if temp is not None: + request_kwargs["temperature"] = temp + completion = client.chat.completions.create(**request_kwargs) answer = completion.choices[0].message.content elif model_str == "gemini-2.0-pro": + genai = _load_genai() genai.configure(api_key=gemini_api_key) - model = genai.GenerativeModel(model_name="gemini-2.0-pro-exp-02-05", system_instruction=system_prompt) + model = genai.GenerativeModel( + model_name="gemini-2.0-pro-exp-02-05", + system_instruction=system_prompt, + ) answer = model.generate_content(prompt).text + elif model_str == "gemini-1.5-pro": + genai = _load_genai() genai.configure(api_key=gemini_api_key) - model = genai.GenerativeModel(model_name="gemini-1.5-pro", system_instruction=system_prompt) + model = genai.GenerativeModel( + model_name="gemini-1.5-pro", + system_instruction=system_prompt, + ) answer = model.generate_content(prompt).text + elif model_str == "o3-mini": - model_str = "o3-mini" - messages = [ - {"role": "user", "content": system_prompt + prompt}] + messages = [{"role": "user", "content": system_prompt + prompt}] if version == "0.28": - completion = openai.ChatCompletion.create( - model=f"{model_str}", messages=messages) + completion = openai.ChatCompletion.create(model=model_str, messages=messages) else: - client = OpenAI() + client = OpenAI(api_key=openai_api_key) completion = client.chat.completions.create( - model="o3-mini-2025-01-31", messages=messages) + model="o3-mini-2025-01-31", + messages=messages, + ) answer = completion.choices[0].message.content - elif model_str == "claude-3.5-sonnet": - client = anthropic.Anthropic(api_key=os.environ["ANTHROPIC_API_KEY"]) + elif model_str == "claude-3-5-sonnet": + client = anthropic.Anthropic(api_key=anthropic_api_key) message = client.messages.create( model="claude-3-5-sonnet-latest", system=system_prompt, - messages=[{"role": "user", "content": prompt}]) + messages=[{"role": "user", "content": prompt}], + ) answer = json.loads(message.to_json())["content"][0]["text"] - elif model_str == "gpt4o" or model_str == "gpt-4o": - model_str = "gpt-4o" + + elif model_str == "gpt-4o": messages = [ {"role": "system", "content": system_prompt}, - {"role": "user", "content": prompt}] + {"role": "user", "content": prompt}, + ] if version == "0.28": - if temp is None: - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages - ) - else: - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages, temperature=temp) + request_kwargs = {"model": model_str, "messages": messages} + if temp is not None: + request_kwargs["temperature"] = temp + completion = openai.ChatCompletion.create(**request_kwargs) else: - client = OpenAI() - if temp is None: - completion = client.chat.completions.create( - model="gpt-4o-2024-08-06", messages=messages, ) - else: - completion = client.chat.completions.create( - model="gpt-4o-2024-08-06", messages=messages, temperature=temp) + client = OpenAI(api_key=openai_api_key) + request_kwargs = {"model": "gpt-4o-2024-08-06", "messages": messages} + if temp is not None: + request_kwargs["temperature"] = temp + completion = client.chat.completions.create(**request_kwargs) answer = completion.choices[0].message.content + elif model_str == "deepseek-chat": - model_str = "deepseek-chat" messages = [ {"role": "system", "content": system_prompt}, - {"role": "user", "content": prompt}] + {"role": "user", "content": prompt}, + ] if version == "0.28": raise Exception("Please upgrade your OpenAI version to use DeepSeek client") - else: - deepseek_client = OpenAI( - api_key=os.getenv('DEEPSEEK_API_KEY'), - base_url="https://api.deepseek.com/v1" - ) - if temp is None: - completion = deepseek_client.chat.completions.create( - model="deepseek-chat", - messages=messages) - else: - completion = deepseek_client.chat.completions.create( - model="deepseek-chat", - messages=messages, - temperature=temp) + deepseek_client = OpenAI( + api_key=deepseek_api_key, + base_url="https://api.deepseek.com/v1", + ) + request_kwargs = {"model": "deepseek-chat", "messages": messages} + if temp is not None: + request_kwargs["temperature"] = temp + completion = deepseek_client.chat.completions.create(**request_kwargs) answer = completion.choices[0].message.content + elif model_str == "o1-mini": - model_str = "o1-mini" - messages = [ - {"role": "user", "content": system_prompt + prompt}] + messages = [{"role": "user", "content": system_prompt + prompt}] if version == "0.28": - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages) + completion = openai.ChatCompletion.create(model=model_str, messages=messages) else: - client = OpenAI() + client = OpenAI(api_key=openai_api_key) completion = client.chat.completions.create( - model="o1-mini-2024-09-12", messages=messages) + model="o1-mini-2024-09-12", + messages=messages, + ) answer = completion.choices[0].message.content + elif model_str == "o1": - model_str = "o1" - messages = [ - {"role": "user", "content": system_prompt + prompt}] + messages = [{"role": "user", "content": system_prompt + prompt}] if version == "0.28": completion = openai.ChatCompletion.create( - model="o1-2024-12-17", # engine = "deployment_name". - messages=messages) + model="o1-2024-12-17", + messages=messages, + ) else: - client = OpenAI() + client = OpenAI(api_key=openai_api_key) completion = client.chat.completions.create( - model="o1-2024-12-17", messages=messages) + model="o1-2024-12-17", + messages=messages, + ) answer = completion.choices[0].message.content + elif model_str == "o1-preview": - model_str = "o1-preview" - messages = [ - {"role": "user", "content": system_prompt + prompt}] + messages = [{"role": "user", "content": system_prompt + prompt}] if version == "0.28": - completion = openai.ChatCompletion.create( - model=f"{model_str}", # engine = "deployment_name". - messages=messages) + completion = openai.ChatCompletion.create(model=model_str, messages=messages) else: - client = OpenAI() + client = OpenAI(api_key=openai_api_key) completion = client.chat.completions.create( - model="o1-preview", messages=messages) + model="o1-preview", + messages=messages, + ) answer = completion.choices[0].message.content - try: - if model_str in ["o1-preview", "o1-mini", "claude-3.5-sonnet", "o1", "o3-mini"]: - encoding = tiktoken.encoding_for_model("gpt-4o") - elif model_str in ["deepseek-chat"]: - encoding = tiktoken.encoding_for_model("cl100k_base") - else: - encoding = tiktoken.encoding_for_model(model_str) - if model_str not in TOKENS_IN: - TOKENS_IN[model_str] = 0 - TOKENS_OUT[model_str] = 0 - TOKENS_IN[model_str] += len(encoding.encode(system_prompt + prompt)) - TOKENS_OUT[model_str] += len(encoding.encode(answer)) - if print_cost: - print(f"Current experiment cost = ${curr_cost_est()}, ** Approximate values, may not reflect true cost") - except Exception as e: - if print_cost: print(f"Cost approximation has an error? {e}") + else: + raise ValueError(f"Unsupported model '{model_str}'") + + _record_token_usage(model_str, system_prompt, prompt, answer, print_cost) return answer except Exception as e: print("Inference Exception:", e) @@ -211,4 +257,4 @@ def query_model(model_str, prompt, system_prompt, openai_api_key=None, gemini_ap raise Exception("Max retries: timeout") -#print(query_model(model_str="o1-mini", prompt="hi", system_prompt="hey")) \ No newline at end of file +# print(query_model(model_str="o1-mini", prompt="hi", system_prompt="hey")) diff --git a/requirements.txt b/requirements.txt index b3f57e4..dae773b 100755 --- a/requirements.txt +++ b/requirements.txt @@ -32,6 +32,8 @@ fonttools==4.55.0 frozenlist==1.5.0 fsspec==2024.9.0 gast==0.6.0 +flask==3.1.3 +flask-sqlalchemy==3.1.1 google-pasta==0.2.0 grpcio==1.68.0 h11==0.14.0 @@ -104,6 +106,7 @@ shellingham==1.5.4 six==1.16.0 smart-open==7.0.5 sniffio==1.3.1 +sentence-transformers==5.6.0 spacy==3.8.2 spacy-legacy==3.0.12 spacy-loggers==1.0.5 @@ -132,4 +135,3 @@ xxhash==3.5.0 yarl==1.18.0 zipp==3.21.0 google-generativeai -PyPDF2 diff --git a/tests/test_tokenization.py b/tests/test_tokenization.py new file mode 100644 index 0000000..1985935 --- /dev/null +++ b/tests/test_tokenization.py @@ -0,0 +1,32 @@ +import unittest + +import inference +from tokenization import encoding_name_for_model, normalize_model_name + + +class TokenizationTests(unittest.TestCase): + def test_normalize_model_name_handles_repo_aliases(self): + self.assertEqual(normalize_model_name("gpt4o"), "gpt-4o") + self.assertEqual(normalize_model_name("gpt4omini"), "gpt-4o-mini") + self.assertEqual(normalize_model_name("claude-3.5-sonnet"), "claude-3-5-sonnet") + + def test_encoding_name_for_model_uses_expected_fallbacks(self): + self.assertEqual(encoding_name_for_model("deepseek-chat"), "cl100k_base") + self.assertEqual(encoding_name_for_model("o3-mini"), "o200k_base") + self.assertEqual(encoding_name_for_model("cl100k_base"), "cl100k_base") + + +class CostEstimationTests(unittest.TestCase): + def setUp(self): + inference.TOKENS_IN.clear() + inference.TOKENS_OUT.clear() + + def test_curr_cost_est_ignores_unknown_models(self): + inference.TOKENS_IN["unknown-model"] = 1000 + inference.TOKENS_OUT["unknown-model"] = 500 + + self.assertEqual(inference.curr_cost_est(), 0.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tokenization.py b/tokenization.py new file mode 100644 index 0000000..baa1d2f --- /dev/null +++ b/tokenization.py @@ -0,0 +1,73 @@ +from functools import lru_cache + +import tiktoken + + +KNOWN_ENCODINGS = { + "cl100k_base", + "o200k_base", + "p50k_base", + "r50k_base", + "p50k_edit", + "gpt2", +} + +MODEL_ALIASES = { + "gpt4o": "gpt-4o", + "gpt4omini": "gpt-4o-mini", + "gpt-4omini": "gpt-4o-mini", + "gpt4o-mini": "gpt-4o-mini", + "claude-3.5-sonnet": "claude-3-5-sonnet", +} + +EXPLICIT_MODEL_ENCODINGS = { + "claude-3-5-sonnet": "cl100k_base", + "deepseek-chat": "cl100k_base", + "gemini-1.5-pro": "cl100k_base", + "gemini-2.0-pro": "cl100k_base", + "o1": "o200k_base", + "o1-mini": "o200k_base", + "o1-preview": "o200k_base", + "o3-mini": "o200k_base", +} + + +def normalize_model_name(model_name): + return MODEL_ALIASES.get(model_name, model_name) + + +def encoding_name_for_model(model_name): + normalized_model = normalize_model_name(model_name) + + if normalized_model in KNOWN_ENCODINGS: + return normalized_model + + explicit_encoding = EXPLICIT_MODEL_ENCODINGS.get(normalized_model) + if explicit_encoding is not None: + return explicit_encoding + + if normalized_model.startswith(("gpt-4o", "o1", "o3", "o4")): + return "o200k_base" + if normalized_model.startswith(("gpt-4", "gpt-3.5", "claude", "deepseek", "gemini")): + return "cl100k_base" + + return normalized_model + + +@lru_cache(maxsize=None) +def get_encoding(model_name): + encoding_name = encoding_name_for_model(model_name) + + if encoding_name in KNOWN_ENCODINGS: + return tiktoken.get_encoding(encoding_name) + + return tiktoken.encoding_for_model(encoding_name) + + +def count_text_tokens(text, model_name): + return len(get_encoding(model_name).encode(text)) + + +def count_message_tokens(messages, model_name): + encoding = get_encoding(model_name) + return sum(len(encoding.encode(message["content"])) for message in messages) diff --git a/utils.py b/utils.py index b424bde..6331d75 100755 --- a/utils.py +++ b/utils.py @@ -1,11 +1,17 @@ import os, re import shutil import time -import tiktoken, openai +import openai import subprocess, string from openai import OpenAI -import google.generativeai as genai from huggingface_hub import InferenceClient +from tokenization import count_message_tokens, get_encoding + + +def load_genai(): + import google.generativeai as genai + + return genai def query_deepseekv3(prompt, system, api_key, attempt=0, temperature=0.0): @@ -23,7 +29,7 @@ def query_deepseekv3(prompt, system, api_key, attempt=0, temperature=0.0): except Exception as e: print(f"Query qwen error: {e}") if attempt >= 10: return f"Your attempt to query deepseekv3 failed: {e}" - return query_deepseekv3(prompt, system, attempt+1) + return query_deepseekv3(prompt, system, api_key, attempt=attempt+1, temperature=temperature) def query_qwen(prompt, system, api_key, attempt=0, temperature=0.0): @@ -47,7 +53,7 @@ def query_qwen(prompt, system, api_key, attempt=0, temperature=0.0): except Exception as e: print(f"Query qwen error: {e}") if attempt >= 10: return f"Your attempt to inference gemini failed: {e}" - return query_qwen(prompt, system, attempt+1) + return query_qwen(prompt, system, api_key, attempt=attempt+1, temperature=temperature) def query_gpt4omini(prompt, system, api_key, attempt=0, temperature=0.0): @@ -69,7 +75,7 @@ def query_gpt4omini(prompt, system, api_key, attempt=0, temperature=0.0): except Exception as e: print(f"Query 4o-mini error: {e}") if attempt >= 10: return f"Your attempt to inference gemini failed: {e}" - return query_gpt4omini(prompt, system, attempt+1) + return query_gpt4omini(prompt, system, api_key, attempt=attempt+1, temperature=temperature) @@ -91,12 +97,13 @@ def query_gpt4o(prompt, system, api_key, attempt=0, temperature=0.0): except Exception as e: print(f"Query gpr-4o error: {e}") if attempt >= 10: return f"Your attempt to inference gemini failed: {e}" - return query_gpt4o(prompt, system, attempt+1) + return query_gpt4o(prompt, system, api_key, attempt=attempt+1, temperature=temperature) def query_gemini(prompt, system, api_key, attempt=0, temperature=0.0): try: + genai = load_genai() genai.configure(api_key=api_key) model = genai.GenerativeModel(model_name="gemini-1.5-pro", system_instruction=system) response = model.generate_content(prompt, generation_config=genai.types.GenerationConfig(temperature=temperature)).text.strip() @@ -106,12 +113,13 @@ def query_gemini(prompt, system, api_key, attempt=0, temperature=0.0): print(f"Gemini error: {e}") if attempt >= 10: return f"Your attempt to inference gemini failed: {e}" time.sleep(1) - return query_gemini(prompt, system, attempt+1) + return query_gemini(prompt, system, api_key, attempt=attempt+1, temperature=temperature) def query_gemini2p0(prompt, system, api_key, attempt=0, temperature=0.0,): try: + genai = load_genai() genai.configure(api_key=api_key) model = genai.GenerativeModel(model_name="gemini-2.0-flash", system_instruction=system) response = model.generate_content(prompt, generation_config=genai.types.GenerationConfig(temperature=temperature)).text.strip() @@ -121,7 +129,7 @@ def query_gemini2p0(prompt, system, api_key, attempt=0, temperature=0.0,): print(f"Gemini error: {e}") if attempt >= 10: return f"Your attempt to inference gemini failed: {e}" time.sleep(1) - return query_gemini2p0(prompt, system, attempt+1) + return query_gemini2p0(prompt, system, api_key, attempt=attempt+1, temperature=temperature) def compile_latex(latex_code, output_path, compile=True, timeout=30): @@ -161,9 +169,7 @@ def compile_latex(latex_code, output_path, compile=True, timeout=30): def count_tokens(messages, model="gpt-4"): - enc = tiktoken.encoding_for_model(model) - num_tokens = sum([len(enc.encode(message["content"])) for message in messages]) - return num_tokens + return count_message_tokens(messages, model) def remove_figures(): """Remove a directory if it exists.""" @@ -195,8 +201,8 @@ def save_to_file(location, filename, data): def clip_tokens(messages, model="gpt-4", max_tokens=100000): - enc = tiktoken.encoding_for_model(model) - total_tokens = sum([len(enc.encode(message["content"])) for message in messages]) + enc = get_encoding(model) + total_tokens = count_message_tokens(messages, model) if total_tokens <= max_tokens: return messages # No need to clip if under the limit @@ -441,7 +447,7 @@ def strip_string(string): # remove percentage string = string.replace("\\%", "") - string = string.replace("\%", "") # noqa: W605 + string = string.replace("%", "") # " 0." equivalent to " ." and "{0." equivalent to "{." Alternatively, add "0" if "." is the start of the string string = string.replace(" .", " 0.")