mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-29 00:17:44 +00:00
feat: Factor Implement Search Enhancement (#294)
* Search enhancement * refactor: reorganize imports for consistency with isort * reformatterd by black --------- Co-authored-by: Tim <illking@foxmail.com>
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import pickle
|
||||
from pathlib import Path
|
||||
|
||||
@@ -49,6 +50,11 @@ class FactorCoSTEER(Developer[FactorExperiment]):
|
||||
if FACTOR_IMPLEMENT_SETTINGS.new_knowledge_base_path is not None
|
||||
else None
|
||||
)
|
||||
self.data_tables_knowledge_path = (
|
||||
Path(FACTOR_IMPLEMENT_SETTINGS.data_tables_knowledge_path)
|
||||
if FACTOR_IMPLEMENT_SETTINGS.data_tables_knowledge_path is not None
|
||||
else None
|
||||
)
|
||||
self.with_knowledge = with_knowledge
|
||||
self.with_feedback = with_feedback
|
||||
self.knowledge_self_gen = knowledge_self_gen
|
||||
@@ -72,6 +78,7 @@ class FactorCoSTEER(Developer[FactorExperiment]):
|
||||
factor_knowledge_base = (
|
||||
FactorGraphKnowledgeBase(
|
||||
init_component_list=component_init_list,
|
||||
data_set_knowledge_path=self.data_tables_knowledge_path,
|
||||
)
|
||||
if self.evolving_version == 2
|
||||
else FactorKnowledgeBaseV1()
|
||||
|
||||
@@ -183,6 +183,19 @@ class FactorEvolvingStrategyWithGraph(MultiProcessEvolvingStrategy):
|
||||
self.num_loop = 0
|
||||
self.haveSelected = False
|
||||
|
||||
def _query_data_tables(self, user_prompt, session):
|
||||
for _ in range(10): # max attempt to reduce the length of user_prompt
|
||||
response = session.build_chat_completion(
|
||||
user_prompt=user_prompt,
|
||||
json_mode=True,
|
||||
)
|
||||
try:
|
||||
result = json.loads(response)
|
||||
return result
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
return None
|
||||
|
||||
def implement_one_factor(
|
||||
self,
|
||||
target_task: FactorTask,
|
||||
@@ -218,6 +231,42 @@ class FactorEvolvingStrategyWithGraph(MultiProcessEvolvingStrategy):
|
||||
queried_knowledge.former_traces[target_factor_task_information] if queried_knowledge is not None else []
|
||||
)
|
||||
|
||||
queried_data_tables = (
|
||||
queried_knowledge.data_set_knowledge_dict[target_factor_task_information]
|
||||
if queried_knowledge is not None
|
||||
else []
|
||||
)
|
||||
queried_data_tables_str = json.dumps(queried_data_tables, indent=2)
|
||||
|
||||
system_prompt = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
.from_string(
|
||||
implement_prompts["evolving_strategy_search_data_table_system_prompt"],
|
||||
)
|
||||
.render()
|
||||
)
|
||||
|
||||
user_prompt = (
|
||||
Environment(undefined=StrictUndefined)
|
||||
.from_string(
|
||||
implement_prompts["evolving_strategy_search_data_table"],
|
||||
)
|
||||
.render(
|
||||
scenario=self.scen.get_scenario_all_desc(),
|
||||
factor_information_str=target_factor_task_information,
|
||||
data_tables=queried_data_tables_str,
|
||||
)
|
||||
)
|
||||
session = APIBackend(use_chat_cache=FACTOR_IMPLEMENT_SETTINGS.coder_use_cache).build_chat_session(
|
||||
session_system_prompt=system_prompt,
|
||||
)
|
||||
|
||||
useful_data_table = self._query_data_tables(user_prompt, session)
|
||||
selected_knowledge_dict = {}
|
||||
for key in useful_data_table:
|
||||
if key in queried_knowledge.data_set_knowledge_dict:
|
||||
selected_knowledge_dict[key] = queried_knowledge.data_set_knowledge_dict[key]
|
||||
|
||||
queried_former_failed_knowledge_to_render = queried_former_failed_knowledge
|
||||
|
||||
system_prompt = (
|
||||
@@ -228,6 +277,7 @@ class FactorEvolvingStrategyWithGraph(MultiProcessEvolvingStrategy):
|
||||
.render(
|
||||
scenario=self.scen.get_scenario_all_desc(),
|
||||
queried_former_failed_knowledge=queried_former_failed_knowledge_to_render,
|
||||
selected_knowledge_dict=selected_knowledge_dict,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import heapq
|
||||
import json
|
||||
import random
|
||||
import re
|
||||
@@ -204,11 +205,13 @@ class FactorQueriedGraphKnowledge(FactorQueriedKnowledge):
|
||||
former_traces: dict = {},
|
||||
component_with_success_task: dict = {},
|
||||
error_with_success_task: dict = {},
|
||||
data_set_knowledge_dict: dict = {},
|
||||
**kwargs,
|
||||
) -> None:
|
||||
self.former_traces = former_traces
|
||||
self.component_with_success_task = component_with_success_task
|
||||
self.error_with_success_task = error_with_success_task
|
||||
self.data_set_knowledge_dict = data_set_knowledge_dict
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
@@ -308,6 +311,10 @@ class FactorGraphRAGStrategy(RAGStrategy):
|
||||
FACTOR_IMPLEMENT_SETTINGS.v2_query_error_limit,
|
||||
knowledge_sampler=conf_knowledge_sampler,
|
||||
)
|
||||
factor_implementation_queried_graph_knowledge = self.dataset_query(
|
||||
evo,
|
||||
factor_implementation_queried_graph_knowledge,
|
||||
)
|
||||
return factor_implementation_queried_graph_knowledge
|
||||
|
||||
def analyze_component(
|
||||
@@ -710,9 +717,37 @@ class FactorGraphRAGStrategy(RAGStrategy):
|
||||
|
||||
return factor_implementation_queried_graph_knowledge
|
||||
|
||||
def dataset_query(
|
||||
self,
|
||||
evo: EvolvableSubjects,
|
||||
factor_implementation_queried_graph_knowledge: FactorQueriedGraphKnowledge,
|
||||
) -> QueriedKnowledge | None:
|
||||
for task_index, target_factor_task in enumerate(evo.sub_tasks):
|
||||
target_factor_task_information = target_factor_task.get_task_information()
|
||||
related_info = {}
|
||||
|
||||
knowledge_dict = self.knowledgebase.data_set_knowledge_dict
|
||||
table_explanations = [f"{key}: {json.dumps(value)}" for key, value in knowledge_dict.items()]
|
||||
|
||||
similarity = calculate_embedding_distance_between_str_list(
|
||||
[target_factor_task_information], table_explanations
|
||||
)[0]
|
||||
|
||||
top_related_indexes = heapq.nlargest(10, range(len(similarity)), key=lambda i: similarity[i])
|
||||
|
||||
for index in top_related_indexes:
|
||||
key = list(knowledge_dict.keys())[index]
|
||||
related_info[key] = knowledge_dict[key]
|
||||
|
||||
factor_implementation_queried_graph_knowledge.data_set_knowledge_dict[
|
||||
target_factor_task_information
|
||||
] = related_info
|
||||
|
||||
return factor_implementation_queried_graph_knowledge
|
||||
|
||||
|
||||
class FactorGraphKnowledgeBase(KnowledgeBase):
|
||||
def __init__(self, init_component_list=None) -> None:
|
||||
def __init__(self, init_component_list=None, data_set_knowledge_path=None) -> None:
|
||||
"""
|
||||
Load knowledge, offer brief information of knowledge and common handle interfaces
|
||||
"""
|
||||
@@ -740,6 +775,12 @@ class FactorGraphKnowledgeBase(KnowledgeBase):
|
||||
# store the task description to component nodes
|
||||
self.task_to_component_nodes = {}
|
||||
|
||||
# data set: data set information
|
||||
self.data_set_knowledge_dict = {}
|
||||
if data_set_knowledge_path:
|
||||
with open(data_set_knowledge_path, "r") as f:
|
||||
self.data_set_knowledge_dict = json.load(f)
|
||||
|
||||
def get_all_nodes_by_label(self, label: str) -> list[UndirectedNode]:
|
||||
return self.graph.get_all_nodes_by_label(label)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user