Files
wehub-resource-sync e64161ec32
CI / ci (3.11) (push) Has been cancelled
CI / ci (3.10) (push) Has been cancelled
CI / dependabot (push) Has been cancelled
Release / release_and_publish (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:36:15 +08:00

62 lines
2.4 KiB
Python

import asyncio
from typing import Any
from rdagent.app.finetune.llm.conf import LLMFinetunePropSetting
from rdagent.components.coder.finetune.conf import get_ft_env
from rdagent.components.workflow.rd_loop import RDLoop
from rdagent.core.conf import RD_AGENT_SETTINGS
from rdagent.core.exception import CoderError
from rdagent.core.proposal import HypothesisFeedback
from rdagent.log import rdagent_logger as logger
from rdagent.scenarios.finetune.proposal.trace import FTTrace
class LLMFinetuneRDLoop(RDLoop):
"""LLM fine-tuning loop using standard RDLoop workflow"""
skip_loop_error = (CoderError,)
withdraw_loop_error = ()
def __init__(self, PROP_SETTING: LLMFinetunePropSetting):
# Store finetune-specific settings
self.ft_rd_setting = PROP_SETTING
self.dataset = PROP_SETTING.dataset
self.model = PROP_SETTING.base_model
# Initialize using base class
super().__init__(PROP_SETTING)
# Replace generic Trace with FTTrace for SOTA tracking
self.trace = FTTrace(scen=self.trace.scen)
async def direct_exp_gen(self, prev_out: dict[str, Any]):
"""Generate LLM fine-tuning experiment"""
exp = await self.hypothesis_gen.async_gen(self.trace, self)
logger.log_object(exp.hypothesis, tag="hypothesis")
logger.log_object(exp.sub_tasks, tag="experiment generation")
return exp
def coding(self, prev_out: dict[str, Any]):
"""Generate fine-tuning code"""
exp = prev_out["direct_exp_gen"]
exp = self.coder.develop(exp)
logger.log_object(exp.sub_workspace_list, tag="coder result")
return exp
def feedback(self, prev_out: dict[str, Any]):
"""Generate feedback for LLM fine-tuning experiment - always call LLM"""
# Get experiment from available sources
exp = prev_out.get("running") or prev_out.get("coding") or prev_out.get("direct_exp_gen")
e = prev_out.get(self.EXCEPTION_KEY, None)
feedback = self.summarizer.generate_feedback(exp, self.trace, exception=e)
logger.log_object(feedback, tag="feedback")
return feedback
def record(self, prev_out: dict[str, Any]):
"""Record the experiment and feedback into trace"""
feedback = prev_out["feedback"]
exp = prev_out.get("running") or prev_out.get("coding") or prev_out.get("direct_exp_gen")
self.trace.sync_dag_parent_and_hist((exp, feedback), prev_out[self.LOOP_IDX_KEY])