Files
NexQuant/rdagent/components/task_implementation/factor_implementation/CoSTEER.py
T

105 lines
4.4 KiB
Python
Raw Normal View History

2024-06-14 12:59:44 +08:00
import pickle
from pathlib import Path
from typing import List
from rdagent.components.task_implementation.factor_implementation.evolving.evaluators import (
FactorImplementationEvaluatorV1,
FactorImplementationsMultiEvaluator,
)
from rdagent.components.task_implementation.factor_implementation.evolving.evolving_strategy import (
FactorEvolvingStrategyWithGraph,
)
from rdagent.components.task_implementation.factor_implementation.evolving.factor import (
FactorEvovlingItem,
FactorImplementTask,
)
from rdagent.components.task_implementation.factor_implementation.evolving.knowledge_management import (
FactorImplementationGraphKnowledgeBase,
FactorImplementationGraphRAGStrategy,
FactorImplementationKnowledgeBaseV1,
)
from rdagent.components.task_implementation.factor_implementation.share_modules.factor_implementation_config import (
FACTOR_IMPLEMENT_SETTINGS,
2024-06-14 12:59:44 +08:00
)
from rdagent.core.evolving_agent import RAGEvoAgent
from rdagent.core.implementation import TaskGenerator
from rdagent.core.task import TaskImplementation
2024-06-14 12:59:44 +08:00
2024-06-14 12:59:44 +08:00
class CoSTEERFG(TaskGenerator):
def __init__(
self,
with_knowledge: bool = True,
with_feedback: bool = True,
knowledge_self_gen: bool = True,
) -> None:
self.max_loop = FACTOR_IMPLEMENT_SETTINGS.max_loop
self.knowledge_base_path = (
Path(FACTOR_IMPLEMENT_SETTINGS.knowledge_base_path)
if FACTOR_IMPLEMENT_SETTINGS.knowledge_base_path is not None
else None
)
self.new_knowledge_base_path = (
Path(FACTOR_IMPLEMENT_SETTINGS.new_knowledge_base_path)
if FACTOR_IMPLEMENT_SETTINGS.new_knowledge_base_path is not None
else None
)
2024-06-14 12:59:44 +08:00
self.with_knowledge = with_knowledge
self.with_feedback = with_feedback
self.knowledge_self_gen = knowledge_self_gen
self.evolving_strategy = FactorEvolvingStrategyWithGraph()
# declare the factor evaluator
self.factor_evaluator = FactorImplementationsMultiEvaluator(FactorImplementationEvaluatorV1())
self.evolving_version = 2
def load_or_init_knowledge_base(self, former_knowledge_base_path: Path = None, component_init_list: list = []):
if former_knowledge_base_path is not None and former_knowledge_base_path.exists():
factor_knowledge_base = pickle.load(open(former_knowledge_base_path, "rb"))
if self.evolving_version == 1 and not isinstance(
factor_knowledge_base, FactorImplementationKnowledgeBaseV1
):
raise ValueError("The former knowledge base is not compatible with the current version")
elif self.evolving_version == 2 and not isinstance(
factor_knowledge_base,
FactorImplementationGraphKnowledgeBase,
):
raise ValueError("The former knowledge base is not compatible with the current version")
else:
factor_knowledge_base = (
FactorImplementationGraphKnowledgeBase(
init_component_list=component_init_list,
)
if self.evolving_version == 2
else FactorImplementationKnowledgeBaseV1()
)
return factor_knowledge_base
2024-06-14 12:59:44 +08:00
def generate(self, tasks: List[FactorImplementTask]) -> List[TaskImplementation]:
# init knowledge base
factor_knowledge_base = self.load_or_init_knowledge_base(
former_knowledge_base_path=self.knowledge_base_path,
component_init_list=[],
)
# init rag method
self.rag = FactorImplementationGraphRAGStrategy(factor_knowledge_base)
2024-06-14 12:59:44 +08:00
# init indermediate items
factor_implementations = FactorEvovlingItem(target_factor_tasks=tasks)
self.evolve_agent = RAGEvoAgent(max_loop=self.max_loop, evolving_strategy=self.evolving_strategy, rag=self.rag)
factor_implementations = self.evolve_agent.multistep_evolve(
factor_implementations,
self.factor_evaluator,
with_knowledge=self.with_knowledge,
with_feedback=self.with_feedback,
knowledge_self_gen=self.knowledge_self_gen,
)
# save new knowledge base
if self.new_knowledge_base_path is not None:
pickle.dump(factor_knowledge_base, open(self.new_knowledge_base_path, "wb"))
self.knowledge_base = factor_knowledge_base
self.latest_factor_implementations = tasks
return factor_implementations