From b3f1a10000cfca1063f9530f952ea23b3874741e Mon Sep 17 00:00:00 2001 From: Xu Yang Date: Tue, 8 Oct 2024 19:34:19 +0800 Subject: [PATCH] use main thread to get embedding to speed up cache (#415) --- rdagent/core/conf.py | 1 + rdagent/oai/llm_utils.py | 57 +++++++------------ .../factor_experiment_loader/pdf_loader.py | 4 +- 3 files changed, 24 insertions(+), 38 deletions(-) diff --git a/rdagent/core/conf.py b/rdagent/core/conf.py index 6d193c7d..00f92685 100644 --- a/rdagent/core/conf.py +++ b/rdagent/core/conf.py @@ -51,6 +51,7 @@ class RDAgentSettings(BaseSettings): embedding_azure_api_base: str = "" embedding_azure_api_version: str = "" embedding_model: str = "" + embedding_max_str_num: int = 50 # offline llama2 related config use_llama2: bool = False diff --git a/rdagent/oai/llm_utils.py b/rdagent/oai/llm_utils.py index 22524ec8..2f5db594 100644 --- a/rdagent/oai/llm_utils.py +++ b/rdagent/oai/llm_utils.py @@ -238,7 +238,7 @@ class APIBackend: """ This is a unified interface for different backends. - (xiao) thinks integerate all kinds of API in a single class is not a good design. + (xiao) thinks integrate all kinds of API in a single class is not a good design. So we should split them into different classes in `oai/backends/` in the future. """ @@ -543,21 +543,25 @@ class APIBackend: filtered_input_content_list = input_content_list if len(filtered_input_content_list) > 0: - if self.use_azure: - response = self.embedding_client.embeddings.create( - model=self.embedding_model, - input=filtered_input_content_list, - ) - else: - response = self.embedding_client.embeddings.create( - model=self.embedding_model, - input=filtered_input_content_list, - ) - for index, data in enumerate(response.data): - content_to_embedding_dict[filtered_input_content_list[index]] = data.embedding + for sliced_filtered_input_content_list in [ + filtered_input_content_list[i : i + self.cfg.embedding_max_str_num] + for i in range(0, len(filtered_input_content_list), self.cfg.embedding_max_str_num) + ]: + if self.use_azure: + response = self.embedding_client.embeddings.create( + model=self.embedding_model, + input=sliced_filtered_input_content_list, + ) + else: + response = self.embedding_client.embeddings.create( + model=self.embedding_model, + input=sliced_filtered_input_content_list, + ) + for index, data in enumerate(response.data): + content_to_embedding_dict[sliced_filtered_input_content_list[index]] = data.embedding - if self.dump_embedding_cache: - self.cache.embedding_set(content_to_embedding_dict) + if self.dump_embedding_cache: + self.cache.embedding_set(content_to_embedding_dict) return [content_to_embedding_dict[content] for content in input_content_list] def _build_log_messages(self, messages: list[dict]) -> str: @@ -735,26 +739,6 @@ class APIBackend: return self.calculate_token_from_messages(messages) -def calculate_embedding_process(str_list: list) -> list: - return APIBackend().create_embedding(str_list) - - -def create_embedding_with_multiprocessing(str_list: list, slice_count: int = 50, nproc: int = 8) -> list: - embeddings = [] - - pool = multiprocessing.Pool(nproc) - result_list = [ - pool.apply_async(calculate_embedding_process, (str_list[index : index + slice_count],)) - for index in range(0, len(str_list), slice_count) - ] - pool.close() - pool.join() - - for res in result_list: - embeddings.extend(res.get()) - return embeddings - - def calculate_embedding_distance_between_str_list( source_str_list: list[str], target_str_list: list[str], @@ -762,7 +746,8 @@ def calculate_embedding_distance_between_str_list( if not source_str_list or not target_str_list: return [[]] - embeddings = create_embedding_with_multiprocessing(source_str_list + target_str_list, slice_count=50, nproc=8) + embeddings = APIBackend().create_embedding(source_str_list + target_str_list) + source_embeddings = embeddings[: len(source_str_list)] target_embeddings = embeddings[len(source_str_list) :] diff --git a/rdagent/scenarios/qlib/factor_experiment_loader/pdf_loader.py b/rdagent/scenarios/qlib/factor_experiment_loader/pdf_loader.py index 3820ab36..74e86cd6 100644 --- a/rdagent/scenarios/qlib/factor_experiment_loader/pdf_loader.py +++ b/rdagent/scenarios/qlib/factor_experiment_loader/pdf_loader.py @@ -22,7 +22,7 @@ from rdagent.core.conf import RD_AGENT_SETTINGS from rdagent.core.prompts import Prompts from rdagent.core.utils import multiprocessing_wrapper from rdagent.log import rdagent_logger as logger -from rdagent.oai.llm_utils import APIBackend, create_embedding_with_multiprocessing +from rdagent.oai.llm_utils import APIBackend from rdagent.scenarios.qlib.factor_experiment_loader.json_loader import ( FactorExperimentLoaderFromDict, ) @@ -479,7 +479,7 @@ Factor variables: {variables} """ full_str_list = [factor_name_to_full_str[factor_name] for factor_name in factor_names] - embeddings = create_embedding_with_multiprocessing(full_str_list) + embeddings = APIBackend.create_embedding(full_str_list) target_k = None if len(full_str_list) < RD_AGENT_SETTINGS.max_input_duplicate_factor_group: