Files
NexQuant/rdagent/factor_implementation/task_loader/pdf_loader.py
T

585 lines
22 KiB
Python
Raw Normal View History

2024-05-21 22:48:41 +08:00
from __future__ import annotations
import json
import multiprocessing as mp
import re
from pathlib import Path
2024-06-12 15:12:11 +08:00
from typing import Mapping
2024-05-21 22:48:41 +08:00
import numpy as np
import pandas as pd
import tiktoken
2024-06-12 15:12:11 +08:00
from jinja2 import Template
from rdagent.core.conf import RD_AGENT_SETTINGS
from rdagent.core.log import RDAgentLog
2024-06-05 15:36:15 +08:00
from rdagent.core.prompts import Prompts
from rdagent.core.task import TaskLoader
from rdagent.document_reader.document_reader import load_and_process_pdfs_by_langchain
from rdagent.factor_implementation.task_loader.json_loader import FactorImplementationTaskLoaderFromDict
2024-06-05 15:36:15 +08:00
from rdagent.oai.llm_utils import APIBackend, create_embedding_with_multiprocessing
2024-05-21 22:48:41 +08:00
from sklearn.cluster import KMeans
from sklearn.metrics.pairwise import cosine_similarity
from sklearn.preprocessing import normalize
2024-06-05 15:36:15 +08:00
document_process_prompts = Prompts(file_path=Path(__file__).parent / "prompts.yaml")
2024-05-21 22:48:41 +08:00
def classify_report_from_dict(
report_dict: Mapping[str, str],
vote_time: int = 1,
substrings: tuple[str] = (),
) -> dict[str, dict[str, str]]:
"""
Parameters:
- report_dict (Dict[str, str]):
A dictionary where the key is the path of the report (ending with .pdf),
and the value is either the report content as a string.
- input_max_token (int): Specifying the maximum number of input tokens.
- vote_time (int): An integer specifying how many times to vote.
- substrings (list(str)): List of hardcode substrings.
Returns:
- Dict[str, Dict[str, str]]: A dictionary where each key is the path of the report,
with a single key 'class' and its value being the classification result (0 or 1).
"""
# if len(substrings) == 0:
# substrings = (
# "金融工程",
# "金工",
# "回测",
# "因子",
# "机器学习",
# "深度学习",
# "量化",
# )
2024-05-21 22:48:41 +08:00
res_dict = {}
2024-06-05 15:36:15 +08:00
classify_prompt = document_process_prompts["classify_system"]
2024-05-21 22:48:41 +08:00
for key, value in report_dict.items():
if not key.endswith(".pdf"):
continue
file_name = key
if isinstance(value, str):
content = value
else:
RDAgentLog().warning(f"输入格式不符合要求: {file_name}")
2024-05-21 22:48:41 +08:00
res_dict[file_name] = {"class": 0}
continue
# pre-filter document with key words is not necessary, skip this check for now
# if (
# not any(substring in content for substring in substrings) and False
# ):
# res_dict[file_name] = {"class": 0}
# else:
while (
APIBackend().build_messages_and_calculate_token(
user_prompt=content,
system_prompt=classify_prompt,
)
> RD_AGENT_SETTINGS.chat_token_limit
):
content = content[: -(RD_AGENT_SETTINGS.chat_token_limit // 100)]
vote_list = []
for _ in range(vote_time):
user_prompt = content
system_prompt = classify_prompt
res = APIBackend().build_messages_and_create_chat_completion(
user_prompt=user_prompt,
system_prompt=system_prompt,
json_mode=True,
)
try:
res = json.loads(res)
vote_list.append(int(res["class"]))
except json.JSONDecodeError:
RDAgentLog().warning(f"返回值无法解析: {file_name}")
res_dict[file_name] = {"class": 0}
count_0 = vote_list.count(0)
count_1 = vote_list.count(1)
if max(count_0, count_1) > int(vote_time / 2):
break
result = 1 if count_1 > count_0 else 0
res_dict[file_name] = {"class": result}
2024-05-21 22:48:41 +08:00
return res_dict
def __extract_factors_name_and_desc_from_content(
content: str,
) -> dict[str, dict[str, str]]:
session = APIBackend().build_chat_session(
2024-06-05 15:36:15 +08:00
session_system_prompt=document_process_prompts["extract_factors_system"],
2024-05-21 22:48:41 +08:00
)
extracted_factor_dict = {}
current_user_prompt = content
for _ in range(10):
extract_result_resp = session.build_chat_completion(
user_prompt=current_user_prompt,
json_mode=False,
)
re_search_res = re.search(r"```json(.*)```", extract_result_resp, re.S)
ret_json_str = re_search_res.group(1) if re_search_res is not None else ""
try:
ret_dict = json.loads(ret_json_str)
parse_success = bool(isinstance(ret_dict, dict)) and "factors" in ret_dict
except json.JSONDecodeError:
parse_success = False
if ret_json_str is None or not parse_success:
2024-06-05 15:36:15 +08:00
current_user_prompt = "Your response didn't follow the instruction might be wrong json format. Try again."
2024-05-21 22:48:41 +08:00
else:
factors = ret_dict["factors"]
if len(factors) == 0:
break
for factor_name, factor_description in factors.items():
extracted_factor_dict[factor_name] = factor_description
2024-06-05 15:36:15 +08:00
current_user_prompt = document_process_prompts["extract_factors_follow_user"]
2024-05-21 22:48:41 +08:00
return extracted_factor_dict
def __extract_factors_formulation_from_content(
content: str,
factor_dict: dict[str, str],
) -> dict[str, dict[str, str]]:
factor_dict_df = pd.DataFrame(
factor_dict.items(),
columns=["factor_name", "factor_description"],
)
2024-06-05 15:36:15 +08:00
system_prompt = document_process_prompts["extract_factor_formulation_system"]
2024-05-21 22:48:41 +08:00
current_user_prompt = Template(
2024-06-05 15:36:15 +08:00
document_process_prompts["extract_factor_formulation_user"],
2024-05-21 22:48:41 +08:00
).render(report_content=content, factor_dict=factor_dict_df.to_string())
session = APIBackend().build_chat_session(session_system_prompt=system_prompt)
factor_to_formulation = {}
for _ in range(10):
extract_result_resp = session.build_chat_completion(
user_prompt=current_user_prompt,
json_mode=False,
)
re_search_res = re.search(r"```json(.*)```", extract_result_resp, re.S)
ret_json_str = re_search_res.group(1) if re_search_res is not None else ""
try:
ret_dict = json.loads(ret_json_str)
parse_success = bool(isinstance(ret_dict, dict))
except json.JSONDecodeError:
parse_success = False
if ret_json_str is None or not parse_success:
2024-06-05 15:36:15 +08:00
current_user_prompt = "Your response didn't follow the instruction might be wrong json format. Try again."
2024-05-21 22:48:41 +08:00
else:
for name, formulation_and_description in ret_dict.items():
if name in factor_dict:
factor_to_formulation[name] = formulation_and_description
if len(factor_to_formulation) != len(factor_dict):
remain_df = factor_dict_df[~factor_dict_df["factor_name"].isin(factor_to_formulation)]
current_user_prompt = (
"Some factors are missing. Please check the following"
" factors and their descriptions and continue extraction.\n"
"==========================Remaining factors"
"==========================\n" + remain_df.to_string()
)
else:
break
return factor_to_formulation
2024-06-05 15:36:15 +08:00
def __extract_factor_and_formulation_from_one_report(
2024-05-21 22:48:41 +08:00
content: str,
) -> dict[str, dict[str, str]]:
final_factor_dict_to_one_report = {}
factor_dict = __extract_factors_name_and_desc_from_content(content)
if len(factor_dict) != 0:
factor_to_formulation = __extract_factors_formulation_from_content(
content,
factor_dict,
)
for factor_name in factor_dict:
2024-06-05 15:36:15 +08:00
if factor_name not in factor_to_formulation:
continue
2024-05-21 22:48:41 +08:00
final_factor_dict_to_one_report.setdefault(factor_name, {})
final_factor_dict_to_one_report[factor_name]["description"] = factor_dict[factor_name]
# use code to correct _ in formulation
formulation = factor_to_formulation[factor_name]["formulation"]
if factor_name in formulation:
target_factor_name = factor_name.replace("_", r"\_")
formulation = formulation.replace(factor_name, target_factor_name)
for variable in factor_to_formulation[factor_name]["variables"]:
if variable in formulation:
target_variable = variable.replace("_", r"\_")
formulation = formulation.replace(variable, target_variable)
final_factor_dict_to_one_report[factor_name]["formulation"] = formulation
final_factor_dict_to_one_report[factor_name]["variables"] = factor_to_formulation[factor_name]["variables"]
return final_factor_dict_to_one_report
2024-06-05 15:36:15 +08:00
def extract_factors_from_report_dict(
2024-05-21 22:48:41 +08:00
report_dict: dict[str, str],
useful_no_dict: dict[str, dict[str, str]],
n_proc: int = 11,
) -> dict[str, dict[str, dict[str, str]]]:
useful_report_dict = {}
for key, value in useful_no_dict.items():
if isinstance(value, dict):
if int(value.get("class")) == 1:
useful_report_dict[key] = report_dict[key]
else:
RDAgentLog().warning(f"Invalid input format: {key}")
2024-05-21 22:48:41 +08:00
final_report_factor_dict = {}
# for file_name, content in useful_report_dict.items():
# final_report_factor_dict.setdefault(file_name, {})
2024-06-05 15:36:15 +08:00
# final_report_factor_dict[file_name] = __extract_factor_and_formulation_from_one_report(content)
2024-05-21 22:48:41 +08:00
while len(final_report_factor_dict) != len(useful_report_dict):
pool = mp.Pool(n_proc)
pool_result_list = []
file_names = []
for file_name, content in useful_report_dict.items():
if file_name in final_report_factor_dict:
continue
file_names.append(file_name)
pool_result_list.append(
pool.apply_async(
2024-06-05 15:36:15 +08:00
__extract_factor_and_formulation_from_one_report,
2024-05-21 22:48:41 +08:00
(content,),
),
)
pool.close()
pool.join()
for index, result in enumerate(pool_result_list):
if result.get is not None:
file_name = file_names[index]
final_report_factor_dict.setdefault(file_name, {})
final_report_factor_dict[file_name] = result.get()
RDAgentLog().info(f"已经完成{len(final_report_factor_dict)}个报告的因子提取")
2024-05-21 22:48:41 +08:00
return final_report_factor_dict
2024-06-05 15:36:15 +08:00
def merge_file_to_factor_dict_to_factor_dict(
file_to_factor_dict: dict[str, dict],
) -> dict:
factor_dict = {}
for file_name in file_to_factor_dict:
for factor_name in file_to_factor_dict[file_name]:
factor_dict.setdefault(factor_name, [])
factor_dict[factor_name].append(file_to_factor_dict[file_name][factor_name])
factor_dict_simple_deduplication = {}
for factor_name in factor_dict:
if len(factor_dict[factor_name]) > 1:
factor_dict_simple_deduplication[factor_name] = max(
factor_dict[factor_name],
key=lambda x: len(x["formulation"]),
)
else:
factor_dict_simple_deduplication[factor_name] = factor_dict[factor_name][0]
return factor_dict_simple_deduplication
def __check_factor_dict_viability_simulate_json_mode(
2024-05-21 22:48:41 +08:00
factor_df_string: str,
) -> dict[str, dict[str, str]]:
session = APIBackend().build_chat_session(
2024-06-05 15:36:15 +08:00
session_system_prompt=document_process_prompts["factor_viability_system"],
2024-05-21 22:48:41 +08:00
)
current_user_prompt = factor_df_string
for _ in range(10):
extract_result_resp = session.build_chat_completion(
user_prompt=current_user_prompt,
json_mode=False,
)
re_search_res = re.search(r"```json(.*)```", extract_result_resp, re.S)
ret_json_str = re_search_res.group(1) if re_search_res is not None else ""
try:
ret_dict = json.loads(ret_json_str)
parse_success = bool(isinstance(ret_dict, dict))
except json.JSONDecodeError:
parse_success = False
if ret_json_str is None or not parse_success:
2024-06-05 15:36:15 +08:00
current_user_prompt = "Your response didn't follow the instruction might be wrong json format. Try again."
2024-05-21 22:48:41 +08:00
else:
return ret_dict
return {}
2024-06-05 15:36:15 +08:00
def check_factor_viability(
2024-05-21 22:48:41 +08:00
factor_dict: dict[str, dict[str, str]],
2024-06-05 15:36:15 +08:00
) -> tuple[dict[str, dict[str, str]], dict[str, dict[str, str]]]:
2024-05-21 22:48:41 +08:00
factor_viability_dict = {}
factor_df = pd.DataFrame(factor_dict).T
factor_df.index.names = ["factor_name"]
while factor_df.shape[0] > 0:
pool = mp.Pool(8)
result_list = []
for i in range(0, factor_df.shape[0], 50):
target_factor_df_string = factor_df.iloc[i : i + 50, :].to_string()
result_list.append(
pool.apply_async(
2024-06-05 15:36:15 +08:00
__check_factor_dict_viability_simulate_json_mode,
2024-05-21 22:48:41 +08:00
(target_factor_df_string,),
),
)
pool.close()
pool.join()
for result in result_list:
respond = result.get()
for factor_name, viability in respond.items():
factor_viability_dict[factor_name] = viability
factor_df = factor_df[~factor_df.index.isin(factor_viability_dict)]
2024-06-05 15:36:15 +08:00
# filtered_factor_dict = {
# factor_name: factor_dict[factor_name]
# for factor_name in factor_dict
# if factor_viability_dict[factor_name]["viability"]
# }
2024-05-21 22:48:41 +08:00
return factor_viability_dict
2024-06-05 15:36:15 +08:00
def __check_factor_duplication_simulate_json_mode(
2024-05-21 22:48:41 +08:00
factor_df: pd.DataFrame,
) -> list[list[str]]:
session = APIBackend().build_chat_session(
2024-06-05 15:36:15 +08:00
session_system_prompt=document_process_prompts["factor_duplicate_system"],
2024-05-21 22:48:41 +08:00
)
current_user_prompt = factor_df.to_string()
generated_duplicated_groups = []
for _ in range(20):
extract_result_resp = session.build_chat_completion(
user_prompt=current_user_prompt,
json_mode=False,
)
re_search_res = re.search(r"```json(.*)```", extract_result_resp, re.S)
ret_json_str = re_search_res.group(1) if re_search_res is not None else ""
try:
ret_dict = json.loads(ret_json_str)
parse_success = bool(isinstance(ret_dict, list))
except json.JSONDecodeError:
parse_success = False
if ret_json_str is None or not parse_success:
current_user_prompt = (
"Your previous response didn't follow"
" the instruction might be wrong json"
" format. Try reducing the factors."
)
elif len(ret_dict) == 0:
return generated_duplicated_groups
else:
generated_duplicated_groups.extend(ret_dict)
current_user_prompt = (
"Continue to extract duplicated"
" groups. If no more duplicated group"
" found please respond empty dict."
)
return generated_duplicated_groups
2024-06-05 15:36:15 +08:00
def __kmeans_embeddings(embeddings: np.ndarray, k: int = 20) -> list[list[str]]:
2024-05-21 22:48:41 +08:00
x_normalized = normalize(embeddings)
kmeans = KMeans(
n_clusters=k,
init="random",
max_iter=100,
n_init=10,
random_state=42,
)
# KMeans algorithm uses Euclidean distance, and we need to customize a function to find the most similar cluster center
2024-05-21 22:48:41 +08:00
def find_closest_cluster_cosine_similarity(
data: np.ndarray,
centroids: np.ndarray,
) -> np.ndarray:
similarity = cosine_similarity(data, centroids)
return np.argmax(similarity, axis=1)
# Initializes the cluster center
2024-05-21 22:48:41 +08:00
rng = np.random.default_rng()
centroids = rng.choice(x_normalized, size=k, replace=False)
# Iterate until convergence or the maximum number of iterations is reached
2024-05-21 22:48:41 +08:00
for _ in range(kmeans.max_iter):
# Assign the sample to the nearest cluster center
2024-05-21 22:48:41 +08:00
closest_clusters = find_closest_cluster_cosine_similarity(
x_normalized,
centroids,
)
# update the cluster center
2024-05-21 22:48:41 +08:00
new_centroids = np.array(
[x_normalized[closest_clusters == i].mean(axis=0) for i in range(k)],
)
new_centroids = normalize(new_centroids) # 归一化新的簇中心
# Check whether the cluster center has changed
2024-05-21 22:48:41 +08:00
if np.allclose(centroids, new_centroids):
break
centroids = new_centroids
clusters = find_closest_cluster_cosine_similarity(x_normalized, centroids)
cluster_to_index = {}
for index, cluster in enumerate(clusters):
cluster_to_index.setdefault(cluster, []).append(index)
return sorted(
cluster_to_index.values(),
key=lambda x: len(x),
reverse=True,
)
2024-06-05 15:36:15 +08:00
def __deduplicate_factor_dict(factor_dict: dict[str, dict[str, str]]) -> list[list[str]]:
if len(factor_dict) == 0:
return []
2024-05-21 22:48:41 +08:00
factor_df = pd.DataFrame(factor_dict).T
factor_df.index.names = ["factor_name"]
factor_names = sorted(factor_dict)
factor_name_to_full_str = {}
for factor_name in factor_dict:
description = factor_dict[factor_name]["description"]
formulation = factor_dict[factor_name]["formulation"]
variables = factor_dict[factor_name]["variables"]
factor_name_to_full_str[
factor_name
] = f"""Factor name: {factor_name}
Factor description: {description}
Factor formulation: {formulation}
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)
target_k = None
if len(full_str_list) < RD_AGENT_SETTINGS.max_input_duplicate_factor_group:
2024-05-21 22:48:41 +08:00
kmeans_index_group = [list(range(len(full_str_list)))]
target_k = 1
else:
for k in range(
len(full_str_list) // RD_AGENT_SETTINGS.max_input_duplicate_factor_group,
2024-05-21 22:48:41 +08:00
30,
):
2024-06-05 15:36:15 +08:00
kmeans_index_group = __kmeans_embeddings(embeddings=embeddings, k=k)
if len(kmeans_index_group[0]) < RD_AGENT_SETTINGS.max_input_duplicate_factor_group:
2024-05-21 22:48:41 +08:00
target_k = k
RDAgentLog().info(f"K-means group number: {k}")
2024-05-21 22:48:41 +08:00
break
factor_name_groups = [[factor_names[index] for index in index_group] for index_group in kmeans_index_group]
duplication_names_list = []
pool = mp.Pool(target_k)
result_list = []
result_list = [
pool.apply_async(
2024-06-05 15:36:15 +08:00
__check_factor_duplication_simulate_json_mode,
2024-05-21 22:48:41 +08:00
(factor_df.loc[factor_name_group, :],),
)
for factor_name_group in factor_name_groups
]
pool.close()
pool.join()
for result in result_list:
deduplication_factor_names_list = result.get()
for deduplication_factor_names in deduplication_factor_names_list:
filter_factor_names = [
factor_name for factor_name in set(deduplication_factor_names) if factor_name in factor_dict
]
if len(filter_factor_names) > 1:
duplication_names_list.append(filter_factor_names)
return duplication_names_list
def deduplicate_factors_by_llm( # noqa: C901, PLR0912
2024-05-21 22:48:41 +08:00
factor_dict: dict[str, dict[str, str]],
2024-06-12 15:12:11 +08:00
factor_viability_dict: dict[str, dict[str, str]] | None = None,
2024-05-21 22:48:41 +08:00
) -> list[list[str]]:
final_duplication_names_list = []
current_round_factor_dict = factor_dict
for _ in range(10):
2024-06-05 15:36:15 +08:00
duplication_names_list = __deduplicate_factor_dict(current_round_factor_dict)
2024-05-21 22:48:41 +08:00
new_round_names = []
for duplication_names in duplication_names_list:
if len(duplication_names) < RD_AGENT_SETTINGS.max_output_duplicate_factor_group:
2024-05-21 22:48:41 +08:00
final_duplication_names_list.append(duplication_names)
else:
new_round_names.extend(duplication_names)
if len(new_round_names) != 0:
current_round_factor_dict = {factor_name: factor_dict[factor_name] for factor_name in new_round_names}
else:
2024-06-05 15:36:15 +08:00
break
final_duplication_names_list = sorted(final_duplication_names_list, key=lambda x: len(x), reverse=True)
to_replace_dict = {}
for duplication_names in duplication_names_list:
if factor_viability_dict is not None:
viability_list = [factor_viability_dict[name]["viability"] for name in duplication_names]
if True not in viability_list:
continue
target_factor_name = duplication_names[viability_list.index(True)]
else:
target_factor_name = duplication_names[0]
for duplication_factor_name in duplication_names:
if duplication_factor_name == target_factor_name:
continue
to_replace_dict[duplication_factor_name] = target_factor_name
2024-06-12 15:12:11 +08:00
llm_deduplicated_factor_dict = {}
2024-06-05 15:36:15 +08:00
added_lower_name_set = set()
for factor_name in factor_dict:
if factor_name not in to_replace_dict and factor_name.lower() not in added_lower_name_set:
if factor_viability_dict is not None and not factor_viability_dict[factor_name]["viability"]:
continue
added_lower_name_set.add(factor_name.lower())
llm_deduplicated_factor_dict[factor_name] = factor_dict[factor_name]
return llm_deduplicated_factor_dict, final_duplication_names_list
class FactorImplementationTaskLoaderFromPDFfiles(TaskLoader):
def load(self, file_or_folder_path: Path) -> dict:
docs_dict = load_and_process_pdfs_by_langchain(Path(file_or_folder_path))
selected_report_dict = classify_report_from_dict(report_dict=docs_dict, vote_time=1)
file_to_factor_result = extract_factors_from_report_dict(docs_dict, selected_report_dict)
factor_dict = merge_file_to_factor_dict_to_factor_dict(file_to_factor_result)
factor_viability = check_factor_viability(factor_dict)
factor_dict, duplication_names_list = deduplicate_factors_by_llm(factor_dict, factor_viability)
return FactorImplementationTaskLoaderFromDict().load(factor_dict)