mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
bfcl cookbook
This commit is contained in:
parent
dec8b45186
commit
93354a277f
5 changed files with 1237 additions and 0 deletions
562
experiencemaker/cookbook/bfcl/bfcl_agent.py
Normal file
562
experiencemaker/cookbook/bfcl/bfcl_agent.py
Normal file
|
|
@ -0,0 +1,562 @@
|
|||
import os
|
||||
|
||||
os.environ["BFCL_DATA_PATH"] = "data/multiturn_data_base_val.jsonl"
|
||||
os.environ["BFCL_ANSWER_PATH"] = "data/possible_answer"
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
import re
|
||||
import time
|
||||
import json
|
||||
import ray
|
||||
import warnings
|
||||
import tempfile
|
||||
import requests
|
||||
|
||||
from tqdm import tqdm
|
||||
from pathlib import Path
|
||||
from loguru import logger
|
||||
from openai import OpenAI
|
||||
from typing import Dict, List, Any
|
||||
|
||||
from bfcl_utils import (
|
||||
load_test_case,
|
||||
handle_user_turn,
|
||||
handle_tool_calls,
|
||||
extract_tool_schema,
|
||||
extract_single_turn_response,
|
||||
extract_multi_turn_responses,
|
||||
capture_and_print_score_files,
|
||||
create_error_response
|
||||
)
|
||||
from bfcl_eval.model_handler.api_inference.qwen import QwenAPIHandler
|
||||
from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import (
|
||||
is_empty_execute_response,
|
||||
)
|
||||
from bfcl_eval.eval_checker.eval_runner import (
|
||||
multi_turn_runner,
|
||||
ast_file_runner,
|
||||
)
|
||||
from bfcl_eval.eval_checker.eval_runner_helper import record_cost_latency
|
||||
from bfcl_eval.utils import (
|
||||
is_multi_turn,
|
||||
is_relevance_or_irrelevance,
|
||||
find_file_with_suffix,
|
||||
load_file,
|
||||
)
|
||||
|
||||
@ray.remote
|
||||
class BFCLAgent:
|
||||
"""A minimal ReAct Agent for BFCL-v3(multi-turn) tasks."""
|
||||
|
||||
def __init__(self,
|
||||
index: int,
|
||||
task_ids: List[str],
|
||||
experiment_name: str,
|
||||
data_path: str = os.getenv("BFCL_DATA_PATH"),
|
||||
answer_path: Path = Path(os.getenv("BFCL_ANSWER_PATH")),
|
||||
model_name: str = "qwen3-8b",
|
||||
temperature: float = 0.9,
|
||||
max_interactions: int = 30,
|
||||
max_response_size: int = 2000,
|
||||
num_runs: int = 1,
|
||||
enable_thinking: bool = False,
|
||||
use_experience: bool = False,
|
||||
experience_base_url: str = "http://0.0.0.0:8001/",
|
||||
experience_workspace_id: str = "bfcl_8b_0725"):
|
||||
|
||||
self.index: int = index
|
||||
self.task_ids: List[str] = task_ids
|
||||
self.categories: List[str] = [task_id.rsplit("_", 1)[0] if "_" in task_id else task_id for task_id in task_ids]
|
||||
self.experiment_name: str = experiment_name
|
||||
self.data_path: str = data_path
|
||||
self.answer_path: Path = answer_path
|
||||
self.model_name: str = model_name
|
||||
self.temperature: float = temperature
|
||||
self.max_interactions: int = max_interactions
|
||||
self.max_response_size: int = max_response_size
|
||||
self.num_runs: int = num_runs
|
||||
self.enable_thinking: bool = enable_thinking
|
||||
self.use_experience: bool = use_experience
|
||||
self.experience_base_url: str = experience_base_url
|
||||
self.experience_workspace_id: str = experience_workspace_id
|
||||
|
||||
self.history: List[List[List[dict]]] = [[] for _ in range(num_runs)]
|
||||
self.test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)]
|
||||
self.original_test_entry: List[List[Dict[str, Any]]] = [[] for _ in range(num_runs)]
|
||||
self.tool_schema: List[List[List[dict]]] = [[] for _ in range(num_runs)]
|
||||
self.current_turn = [[0 for _ in range(len(task_ids))] for _ in range(num_runs)]
|
||||
|
||||
for run_id in range(num_runs):
|
||||
for task_index in range(len(task_ids)):
|
||||
self.init_state(run_id, task_index)
|
||||
|
||||
def init_state(self, run_id, i) -> Dict[str, Any]:
|
||||
"""载入测试用例并返回首条 user 消息"""
|
||||
self.test_entry[run_id].append(load_test_case(self.data_path, self.task_ids[i]))
|
||||
self.original_test_entry[run_id].append(self.test_entry[run_id][i].get("extra", {}))
|
||||
self.tool_schema[run_id].append(extract_tool_schema(self.test_entry[run_id][i].get("tools", [{}])))
|
||||
|
||||
# 初始历史
|
||||
msg = self.test_entry[run_id][i].get("messages", [])[0]
|
||||
if self.use_experience:
|
||||
query = msg["content"]
|
||||
exp = self.get_experience(query)
|
||||
self.history[run_id].append([self.get_query_with_experience(query, exp)])
|
||||
else:
|
||||
self.history[run_id].append([msg])
|
||||
self.current_turn[run_id][i] = 1
|
||||
|
||||
def get_query_with_experience(self, query: str, experience: str):
|
||||
return {
|
||||
"role": "user",
|
||||
"content": "Task:\n" + query + "\n\nSome Related Experience to help you to complete the task:\n" + experience
|
||||
}
|
||||
|
||||
def get_experience(self, query: str):
|
||||
response = requests.post(url=self.experience_base_url + "retriever", json={
|
||||
"workspace_id": self.experience_workspace_id,
|
||||
"query": query,
|
||||
"top_k": 5
|
||||
})
|
||||
logger.info(f"query:{query}")
|
||||
|
||||
if response.status_code != 200:
|
||||
print(response.text)
|
||||
return ""
|
||||
|
||||
response = response.json()
|
||||
print(response)
|
||||
experience_merged: str = response["experience_merged"]
|
||||
print(f"experience_merged={experience_merged}")
|
||||
return experience_merged
|
||||
|
||||
def update_experience(self, trajectories):
|
||||
response = requests.post(url=self.experience_base_url + "summarizer", json={
|
||||
"workspace_id": self.experience_workspace_id,
|
||||
"trajectories": trajectories,
|
||||
})
|
||||
response.raise_for_status()
|
||||
response = response.json()
|
||||
return response["experiences"]
|
||||
|
||||
def call_llm(self, messages: list, tool_schemas: list[dict]) -> str:
|
||||
for i in range(100):
|
||||
try:
|
||||
client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))
|
||||
# Change this function to modify the base llm
|
||||
response = client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=messages,
|
||||
tools=tool_schemas,
|
||||
temperature=self.temperature,
|
||||
seed=0,
|
||||
extra_body={"enable_thinking": self.enable_thinking},
|
||||
stream=self.enable_thinking,
|
||||
parallel_tool_calls=True,
|
||||
)
|
||||
if not self.enable_thinking:
|
||||
out_msg = response.choices[0].message
|
||||
return out_msg.model_dump(exclude_unset=True, exclude_none=True)
|
||||
else:
|
||||
reasoning_content = "" # Complete reasoning process
|
||||
answer_content = "" # Define complete response
|
||||
tool_info = [] # Store tool invocation information
|
||||
is_answering = False # Determine whether the reasoning process has finished and response has started
|
||||
|
||||
for chunk in response:
|
||||
if not chunk.choices:
|
||||
# Handle usage information
|
||||
continue
|
||||
else:
|
||||
delta = chunk.choices[0].delta
|
||||
# Handle AI's thought process (chain reasoning)
|
||||
if hasattr(delta, 'reasoning_content') and delta.reasoning_content is not None:
|
||||
reasoning_content += delta.reasoning_content
|
||||
|
||||
# Handle final response content
|
||||
else:
|
||||
if not is_answering: # Print title when entering the response phase for the first time
|
||||
is_answering = True
|
||||
if delta.content is not None:
|
||||
answer_content += delta.content
|
||||
|
||||
# Handle tool invocation information (support parallel tool calls)
|
||||
if delta.tool_calls is not None:
|
||||
for tool_call in delta.tool_calls:
|
||||
index = tool_call.index # Tool call index, used for parallel calls
|
||||
|
||||
# Dynamically expand tool information storage list
|
||||
while len(tool_info) <= index:
|
||||
tool_info.append({"id": "", "type": "function", "index": index, "function": { "name": "", "arguments": "" }})
|
||||
|
||||
# Collect tool call ID (used for subsequent function calls)
|
||||
if tool_call.id:
|
||||
tool_info[index]['id'] += tool_call.id
|
||||
|
||||
# Collect function name (used for subsequent routing to specific functions)
|
||||
if tool_call.function and tool_call.function.name:
|
||||
tool_info[index]['function']['name'] += tool_call.function.name
|
||||
|
||||
# Collect function parameters (in JSON string format, need subsequent parsing)
|
||||
if tool_call.function and tool_call.function.arguments:
|
||||
tool_info[index]['function']['arguments'] += tool_call.function.arguments
|
||||
msg = {
|
||||
"role": "assistant",
|
||||
"content": answer_content,
|
||||
"reasoning_content": reasoning_content,
|
||||
}
|
||||
if tool_info:
|
||||
msg["tool_calls"] = tool_info
|
||||
return msg
|
||||
except Exception as e:
|
||||
logger.exception(f"encounter error with {e.args}")
|
||||
time.sleep(1 + i * 10)
|
||||
|
||||
return "call llm error"
|
||||
|
||||
def env_step(self, run_id: int, index: int, messages: str) -> str:
|
||||
"""
|
||||
Process one step in the conversation.
|
||||
Both single turn and multi turn are supported.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages, with the last one being assistant response
|
||||
test_entry: Test entry containing initial_config, involved_classes, question etc.
|
||||
**kwargs: Additional arguments for compatibility
|
||||
|
||||
Returns:
|
||||
Dict containing next message and tools if applicable
|
||||
"""
|
||||
try:
|
||||
if not messages:
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
|
||||
if messages[-1]["role"] != "assistant":
|
||||
return create_error_response(
|
||||
"Last message must be from assistant"
|
||||
)
|
||||
|
||||
if "tool_calls" in messages[-1] and len(messages[-1]["tool_calls"]) > 0:
|
||||
try:
|
||||
tool_calls = messages[-1]["tool_calls"]
|
||||
decoded_calls = self._convert_tool_calls_to_execution_format(
|
||||
tool_calls
|
||||
)
|
||||
# decoded_calls:[function(param=xxx)]
|
||||
print(f"decoded_calls: {decoded_calls}")
|
||||
# todo 实现decode_execute,返回prm
|
||||
# if self.decode_execute(decoded_calls):
|
||||
if is_empty_execute_response(decoded_calls):
|
||||
warnings.warn(
|
||||
f"is_empty_execute_response: {is_empty_execute_response(decoded_calls)}"
|
||||
)
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
return handle_tool_calls(
|
||||
tool_calls, decoded_calls, self.original_test_entry[run_id][index], self.current_turn[run_id][index]
|
||||
)
|
||||
except Exception as e:
|
||||
warnings.warn(f"处理工具调用时发生错误: {str(e)}")
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
else:
|
||||
return handle_user_turn(self.original_test_entry[run_id][index], self.current_turn[run_id][index])
|
||||
|
||||
except Exception as e:
|
||||
return create_error_response(f"处理请求时发生错误: {str(e)}")
|
||||
|
||||
def _convert_tool_calls_to_execution_format(
|
||||
self, tool_calls: List[Dict[str, Any]]
|
||||
) -> List[str]:
|
||||
"""
|
||||
Convert OpenAI format tool calls to execution format.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls in OpenAI format
|
||||
|
||||
Returns:
|
||||
List of function calls in string format
|
||||
"""
|
||||
execution_list = []
|
||||
|
||||
for tool_call in tool_calls:
|
||||
function = tool_call.get("function", {})
|
||||
function_name = function.get("name", "")
|
||||
|
||||
try:
|
||||
arguments = function.get("arguments", "{}")
|
||||
if isinstance(arguments, str):
|
||||
args_dict = json.loads(arguments)
|
||||
else:
|
||||
args_dict = arguments
|
||||
|
||||
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
|
||||
execution_list.append(f"{function_name}({args_str})")
|
||||
|
||||
except Exception as e:
|
||||
execution_list.append(f"{function_name}()")
|
||||
|
||||
return execution_list
|
||||
|
||||
def get_reward(self, run_id, index) -> float:
|
||||
try:
|
||||
if not self.history[run_id][index] or not self.original_test_entry[run_id][index]:
|
||||
return 0.0
|
||||
|
||||
model_name = "env_handler"
|
||||
handler = QwenAPIHandler(
|
||||
model_name, temperature=1.0
|
||||
) # FIXME: magic number
|
||||
|
||||
model_result_data = self._convert_conversation_to_eval_format(run_id, index)
|
||||
|
||||
prompt_data = [self.original_test_entry[run_id][index]]
|
||||
|
||||
state = {"leaderboard_table": {}}
|
||||
record_cost_latency(
|
||||
state["leaderboard_table"], model_name, [model_result_data]
|
||||
)
|
||||
|
||||
if is_relevance_or_irrelevance(self.categories[index]):
|
||||
accuracy, _ = self._eval_relevance_test(
|
||||
handler, model_result_data, prompt_data, model_name, self.category
|
||||
)
|
||||
else:
|
||||
# Find the corresponding possible answer file
|
||||
|
||||
possible_answer_file = find_file_with_suffix(
|
||||
self.answer_path, self.categories[index]
|
||||
)
|
||||
possible_answer = load_file(possible_answer_file, sort_by_id=True)
|
||||
possible_answer = [
|
||||
item for item in possible_answer if item["id"] == self.task_ids[index]
|
||||
]
|
||||
if is_multi_turn(self.categories[index]):
|
||||
accuracy, _ = self._eval_multi_turn_test(
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
self.categories[index],
|
||||
)
|
||||
else:
|
||||
accuracy, _ = self._eval_single_turn_test(
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
self.categories[index],
|
||||
)
|
||||
print(f"model_result_data: {model_result_data}")
|
||||
print(f"possible_answer: {possible_answer}") if possible_answer else None
|
||||
|
||||
return accuracy
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
return 0
|
||||
|
||||
def _convert_conversation_to_eval_format(self, run_id, index) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert conversation history to evaluation format.
|
||||
|
||||
Args:
|
||||
conversation_result: Result from run_conversation
|
||||
original_test_entry: Original test entry data
|
||||
|
||||
Returns:
|
||||
Data in format expected by multi_turn_runner or other runners
|
||||
"""
|
||||
if is_multi_turn(self.categories[index]):
|
||||
turns_data = extract_multi_turn_responses(self.history[run_id][index])
|
||||
else:
|
||||
turns_data = extract_single_turn_response(self.history[run_id][index])
|
||||
|
||||
model_result_data = {
|
||||
"id": self.task_ids[index],
|
||||
"result": turns_data,
|
||||
"latency": 0,
|
||||
"input_token_count": 0,
|
||||
"output_token_count": 0,
|
||||
}
|
||||
|
||||
return model_result_data
|
||||
|
||||
def _eval_multi_turn_test(
|
||||
self,
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
test_category,
|
||||
):
|
||||
"""
|
||||
Evaluate multi-turn test.
|
||||
|
||||
Args:
|
||||
handler: Model handler instance
|
||||
model_result_data: Model result data
|
||||
prompt_data: Prompt data
|
||||
possible_answer: Possible answer data
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
|
||||
Returns:
|
||||
Tuple of (accuracy, total_count)
|
||||
"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
score_dir = Path(temp_dir)
|
||||
accuracy, total_count = multi_turn_runner(
|
||||
handler=handler,
|
||||
model_result=[model_result_data],
|
||||
prompt=prompt_data,
|
||||
possible_answer=possible_answer,
|
||||
model_name=model_name,
|
||||
test_category=test_category,
|
||||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir, model_name, test_category, "multi_turn"
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
def _eval_single_turn_test(
|
||||
self,
|
||||
handler,
|
||||
model_result_data,
|
||||
prompt_data,
|
||||
possible_answer,
|
||||
model_name,
|
||||
test_category,
|
||||
):
|
||||
"""
|
||||
Evaluate single-turn AST test.
|
||||
|
||||
Args:
|
||||
handler: Model handler instance
|
||||
model_result_data: Model result data
|
||||
prompt_data: Prompt data
|
||||
possible_answer: Possible answer data
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
|
||||
Returns:
|
||||
Tuple of (accuracy, total_count)
|
||||
"""
|
||||
language = "Python"
|
||||
if "java" in test_category.lower():
|
||||
language = "Java"
|
||||
elif "js" in test_category.lower() or "javascript" in test_category.lower():
|
||||
language = "JavaScript"
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
score_dir = Path(temp_dir)
|
||||
accuracy, total_count = ast_file_runner(
|
||||
handler=handler,
|
||||
model_result=[model_result_data],
|
||||
prompt=prompt_data,
|
||||
possible_answer=possible_answer,
|
||||
language=language,
|
||||
test_category=test_category,
|
||||
model_name=model_name,
|
||||
score_dir=score_dir,
|
||||
)
|
||||
capture_and_print_score_files(
|
||||
score_dir, model_name, test_category, "single_turn"
|
||||
)
|
||||
return accuracy, total_count
|
||||
|
||||
def execute(self):
|
||||
result = []
|
||||
for task_index, task_id in enumerate(tqdm(self.task_ids, desc=f"ray_index={self.index}")):
|
||||
for run_id in range(self.num_runs):
|
||||
try:
|
||||
for i in range(self.max_interactions):
|
||||
llm_output = self.call_llm(self.history[run_id][task_index], self.tool_schema[run_id][task_index])
|
||||
self.history[run_id][task_index].append(llm_output)
|
||||
|
||||
env_output = self.env_step(run_id, task_index, self.history[run_id][task_index])
|
||||
# 与环境交互后env_output有以下几种返回情况:
|
||||
# 1. 触发query, 附带着available tools列表, {"messages": [{"role": "user", "content": user_query}], "tools": tools}
|
||||
# 2. 返回工具调用结果, {"messages": [{"role": "tool", "content": {<execution_results>}, 'tool_call_id': 'chatcmpl-tool-xxx'}]}
|
||||
# <execution_results>: 正确执行时返回结果dict, e.g., {"travel_cost_list": [1140.0]}, 错误时返回error信息, e.g., {"error": "cd: temporary: No such directory. You cannot use path to change directory."}
|
||||
# 3. 回合结束, 返回{"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
|
||||
# 4. 程序出错, 返回{"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
|
||||
|
||||
# tool_list更新
|
||||
if "tools" in env_output:
|
||||
self.tool_schema[run_id][task_index] = extract_tool_schema(env_output["tools"])
|
||||
|
||||
new_tool_calls=[]
|
||||
new_tool_call_ids=[]
|
||||
next_user_msg = ""
|
||||
for idx, msg in enumerate(env_output.get("messages", [])):
|
||||
if msg["role"] == "tool" and len(msg["content"])>0:
|
||||
new_tool_calls.append(msg.get("content", ""))
|
||||
new_tool_call_ids.append(msg.get("tool_call_id", ""))
|
||||
elif msg["role"] == "user":
|
||||
next_user_msg = msg.get("content", "")
|
||||
self.current_turn[run_id][task_index] += 1
|
||||
else: # for env role messages
|
||||
next_user_msg = msg.get("content", "")
|
||||
|
||||
if new_tool_calls:
|
||||
for idx, call in enumerate(new_tool_calls):
|
||||
self.history[run_id][task_index].append({"role": "tool", "content": str(call), "tool_call_id": new_tool_call_ids[idx]})
|
||||
else:
|
||||
self.history[run_id][task_index].append({"role": "user", "content": next_user_msg})
|
||||
|
||||
logger.info(f"index={self.index} task_id={task_id} iteration={i}")
|
||||
|
||||
if self.task_completed(run_id, task_index):
|
||||
break
|
||||
|
||||
reward = self.get_reward(run_id, task_index)
|
||||
# if reward == 1:
|
||||
#
|
||||
# self.update_experience([process_msg_to_trajectory(task_id, msg, reward)]) # selectively add experiences when succeed
|
||||
|
||||
t_result = {
|
||||
"run_id": run_id,
|
||||
"task_id": self.task_ids[task_index],
|
||||
"experiment_name": self.experiment_name,
|
||||
"task_completed": self.task_completed(run_id, task_index),
|
||||
"reward": reward,
|
||||
"task_history": self.history[run_id][task_index],
|
||||
}
|
||||
result.append(t_result)
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"encounter error with {e.args}")
|
||||
result.append({})
|
||||
return result
|
||||
|
||||
def task_completed(self, run_id, index):
|
||||
"""
|
||||
Check if task is completed.
|
||||
|
||||
Returns:
|
||||
True if task is completed, False otherwise
|
||||
"""
|
||||
return self.history[run_id][index][-1]["content"] == "[CONVERSATION_COMPLETED]"
|
||||
|
||||
def main():
|
||||
with open(os.getenv("BFCL_DATA_PATH"), "r", encoding="utf-8") as f:
|
||||
task_ids = [json.loads(l)["id"] for l in f]
|
||||
dataset_name = "dev"
|
||||
agent = BFCLAgent(
|
||||
index=0,
|
||||
task_id=task_ids[0],
|
||||
experiment_name=f"zouying_{dataset_name}",
|
||||
)
|
||||
result = agent.execute()
|
||||
logger.info(f"result={json.dumps(result)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
386
experiencemaker/cookbook/bfcl/bfcl_utils.py
Normal file
386
experiencemaker/cookbook/bfcl/bfcl_utils.py
Normal file
|
|
@ -0,0 +1,386 @@
|
|||
import json
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Any
|
||||
|
||||
from bfcl_eval.constants.type_mappings import GORILLA_TO_OPENAPI
|
||||
from bfcl_eval.constants.default_prompts import (
|
||||
DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
|
||||
)
|
||||
from bfcl_eval.model_handler.model_style import ModelStyle
|
||||
from bfcl_eval.model_handler.utils import (
|
||||
convert_to_function_call,
|
||||
convert_to_tool,
|
||||
default_decode_ast_prompting,
|
||||
default_decode_execute_prompting,
|
||||
format_execution_results_prompting,
|
||||
func_doc_language_specific_pre_processing,
|
||||
retry_with_backoff,
|
||||
system_prompt_pre_processing_chat_model,
|
||||
)
|
||||
from bfcl_eval.eval_checker.multi_turn_eval.multi_turn_utils import (
|
||||
execute_multi_turn_func_call,
|
||||
)
|
||||
|
||||
|
||||
def load_test_case(data_path: str, test_id: str | None) -> Dict[str, Any]:
|
||||
"""按 ID / 行号加载单条 JSONL 测试用例。找不到就抛错。"""
|
||||
if not Path(data_path).exists():
|
||||
raise FileNotFoundError(f"BFCL data file '{data_path}' not found")
|
||||
|
||||
if test_id is None:
|
||||
raise ValueError("task_id is required")
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
if str(test_id).isdigit():
|
||||
idx = int(test_id)
|
||||
for line_no, line in enumerate(f):
|
||||
if line_no == idx:
|
||||
return json.loads(line)
|
||||
raise ValueError(f"Test case index {idx} not found in {data_path}")
|
||||
else:
|
||||
for line in f:
|
||||
data = json.loads(line)
|
||||
if data.get("id") == test_id:
|
||||
return data
|
||||
raise ValueError(f"Test case id '{test_id}' not found in {data_path}")
|
||||
|
||||
def handle_user_turn(
|
||||
test_entry: Dict[str, Any], current_turn: int
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle user turn by returning appropriate content from test_entry["question"].
|
||||
For non-first turns, processes user query and tools.
|
||||
|
||||
Args:
|
||||
test_entry: Test entry containing conversation data
|
||||
current_turn: Current turn number
|
||||
|
||||
Returns:
|
||||
Response containing next user message and tools
|
||||
"""
|
||||
try:
|
||||
current_turn_message = []
|
||||
tools = compile_tools(test_entry)
|
||||
questions = test_entry.get("question", [])
|
||||
holdout_function = test_entry.get("holdout_function", {})
|
||||
|
||||
if str(current_turn) in holdout_function:
|
||||
test_entry["function"].extend(holdout_function[str(current_turn)])
|
||||
tools = compile_tools(test_entry)
|
||||
assert (
|
||||
len(questions[current_turn]) == 0
|
||||
), "Holdout turn should not have user message."
|
||||
current_turn_message = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": DEFAULT_USER_PROMPT_FOR_ADDITIONAL_FUNCTION_FC,
|
||||
}
|
||||
]
|
||||
return create_user_response(current_turn_message, tools)
|
||||
if current_turn >= len(questions):
|
||||
return create_completion_response()
|
||||
|
||||
current_turn_message = questions[current_turn]
|
||||
|
||||
return create_user_response(current_turn_message, tools)
|
||||
|
||||
except Exception as e:
|
||||
return create_error_response(f"处理用户轮次时发生错误: {str(e)}")
|
||||
|
||||
def handle_tool_calls(
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
decoded_calls: list[str],
|
||||
test_entry: Dict[str, Any],
|
||||
current_turn: int,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle tool calls from assistant.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls in OpenAI format
|
||||
decoded_calls: List of decoded function calls
|
||||
test_entry: Test entry containing environment data
|
||||
current_turn: Current turn number
|
||||
|
||||
Returns:
|
||||
Response containing tool execution results
|
||||
"""
|
||||
execution_results, _ = execute_multi_turn_func_call(
|
||||
func_call_list=decoded_calls,
|
||||
initial_config=test_entry["initial_config"],
|
||||
involved_classes=test_entry["involved_classes"],
|
||||
model_name="env_handler",
|
||||
test_entry_id=test_entry["id"],
|
||||
long_context=(
|
||||
"long_context" in test_entry["id"] or "composite" in test_entry["id"]
|
||||
),
|
||||
is_evaL_run=False,
|
||||
)
|
||||
# print('execution_results in handler_tool_calls:', execution_results)
|
||||
|
||||
return create_tool_response(tool_calls, execution_results)
|
||||
|
||||
|
||||
def compile_tools(test_entry: dict) -> list:
|
||||
"""
|
||||
Compile functions into tools format.
|
||||
|
||||
Args:
|
||||
test_entry: Test entry containing functions
|
||||
|
||||
Returns:
|
||||
List of tools in OpenAI format
|
||||
"""
|
||||
functions: list = test_entry["function"]
|
||||
test_category: str = test_entry["id"].rsplit("_", 1)[0]
|
||||
|
||||
functions = func_doc_language_specific_pre_processing(functions, test_category)
|
||||
tools = convert_to_tool(functions, GORILLA_TO_OPENAPI, ModelStyle.OpenAI_Completions)
|
||||
|
||||
return tools
|
||||
|
||||
def create_tool_response(
|
||||
tool_calls: List[Dict[str, Any]], execution_results: List[str]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for tool calls.
|
||||
|
||||
Args:
|
||||
tool_calls: List of tool calls
|
||||
execution_results: List of execution results
|
||||
|
||||
Returns:
|
||||
Response containing tool execution results
|
||||
"""
|
||||
tool_messages = []
|
||||
for i, (tool_call, result) in enumerate(zip(tool_calls, execution_results)):
|
||||
tool_messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"content": result,
|
||||
"tool_call_id": tool_call.get("id", f"call_{i}"),
|
||||
}
|
||||
)
|
||||
|
||||
return {"messages": tool_messages}
|
||||
|
||||
def create_user_response(
|
||||
question_turn: List[Dict[str, Any]], tools: List[Dict[str, Any]]
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response containing user message.
|
||||
|
||||
Args:
|
||||
question_turn: List of messages for current turn
|
||||
tools: List of available tools
|
||||
|
||||
Returns:
|
||||
Response containing user message and tools
|
||||
"""
|
||||
user_content = ""
|
||||
for msg in question_turn:
|
||||
if msg["role"] == "user":
|
||||
user_content = msg["content"]
|
||||
break
|
||||
|
||||
return {"messages": [{"role": "user", "content": user_content}], "tools": tools}
|
||||
|
||||
def create_completion_response() -> Dict[str, Any]:
|
||||
"""
|
||||
Create response indicating conversation completion.
|
||||
|
||||
Returns:
|
||||
Response with completion message
|
||||
"""
|
||||
return {"messages": [{"role": "env", "content": "[CONVERSATION_COMPLETED]"}]}
|
||||
|
||||
def create_error_response(error_message: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Create response for error conditions.
|
||||
|
||||
Args:
|
||||
error_message: Error message to include
|
||||
|
||||
Returns:
|
||||
Response containing error message
|
||||
"""
|
||||
return {"messages": [{"role": "env", "content": f"[ERROR] {error_message}"}]}
|
||||
|
||||
def decode_execute(result):
|
||||
"""
|
||||
Decode execute results for compatibility with evaluation framework.
|
||||
|
||||
Args:
|
||||
result: Result to decode
|
||||
|
||||
Returns:
|
||||
List of decoded function calls
|
||||
"""
|
||||
return default_decode_execute_prompting(result)
|
||||
|
||||
def extract_single_turn_response(messages: List[Dict[str, Any]]) -> str:
|
||||
"""
|
||||
Extract single-turn response from conversation messages.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages
|
||||
|
||||
Returns:
|
||||
String representation of the response
|
||||
"""
|
||||
for message in reversed(messages):
|
||||
if message["role"] == "assistant":
|
||||
if "tool_calls" in message and message["tool_calls"]:
|
||||
formatted_calls = []
|
||||
for tool_call in message["tool_calls"]:
|
||||
formatted_call = format_single_tool_call_for_eval(
|
||||
tool_call
|
||||
)
|
||||
if formatted_call:
|
||||
formatted_calls.append(formatted_call)
|
||||
return "\n".join(formatted_calls) if formatted_calls else ""
|
||||
elif message.get("content"):
|
||||
return message["content"]
|
||||
|
||||
return ""
|
||||
|
||||
def extract_multi_turn_responses(
|
||||
messages: List[Dict[str, Any]]
|
||||
) -> List[List[str]]:
|
||||
"""
|
||||
Extract multi-turn responses from conversation messages.
|
||||
|
||||
Args:
|
||||
messages: List of conversation messages
|
||||
|
||||
Returns:
|
||||
List of turns, each turn is a list of function call strings
|
||||
"""
|
||||
turns_data = []
|
||||
current_turn_responses = []
|
||||
|
||||
i = 0
|
||||
while i < len(messages):
|
||||
message = messages[i]
|
||||
|
||||
if message["role"] == "user":
|
||||
if current_turn_responses:
|
||||
turns_data.append(current_turn_responses)
|
||||
current_turn_responses = []
|
||||
|
||||
i += 1
|
||||
while i < len(messages) and messages[i]["role"] == "assistant":
|
||||
assistant_msg = messages[i]
|
||||
|
||||
if "tool_calls" in assistant_msg and assistant_msg["tool_calls"]:
|
||||
for tool_call in assistant_msg["tool_calls"]:
|
||||
formatted_call = format_single_tool_call_for_eval(
|
||||
tool_call
|
||||
)
|
||||
if formatted_call:
|
||||
current_turn_responses.append(formatted_call)
|
||||
|
||||
i += 1
|
||||
|
||||
while i < len(messages) and messages[i]["role"] == "tool":
|
||||
i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
if current_turn_responses:
|
||||
turns_data.append(current_turn_responses)
|
||||
|
||||
return turns_data
|
||||
|
||||
def format_single_tool_call_for_eval(tool_call: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Format a single tool call into string representation for evaluation.
|
||||
|
||||
Args:
|
||||
tool_call: Single tool call in OpenAI format
|
||||
|
||||
Returns:
|
||||
Formatted string representation
|
||||
"""
|
||||
function = tool_call.get("function", {})
|
||||
function_name = function.get("name", "")
|
||||
|
||||
try:
|
||||
arguments = function.get("arguments", "{}")
|
||||
if isinstance(arguments, str):
|
||||
args_dict = json.loads(arguments)
|
||||
else:
|
||||
args_dict = arguments
|
||||
|
||||
args_str = ", ".join([f"{k}={repr(v)}" for k, v in args_dict.items()])
|
||||
return f"{function_name}({args_str})"
|
||||
|
||||
except Exception as e:
|
||||
return f"{function_name}()"
|
||||
|
||||
def capture_and_print_score_files(
|
||||
score_dir: Path, model_name: str, test_category: str, eval_type: str
|
||||
):
|
||||
"""
|
||||
Capture and print contents of score files written to score_dir.
|
||||
|
||||
Args:
|
||||
score_dir: Directory containing score files
|
||||
model_name: Name of the model
|
||||
test_category: Category of the test
|
||||
eval_type: Type of evaluation (relevance/multi_turn/single_turn)
|
||||
"""
|
||||
try:
|
||||
print(f"\n=== {eval_type.upper()} Evaluation Result Files ===")
|
||||
print(f"Model: {model_name}")
|
||||
print(f"Test Category: {test_category}")
|
||||
print(f"Evaluation Type: {eval_type}")
|
||||
|
||||
for file_path in score_dir.rglob("*"):
|
||||
if file_path.is_file():
|
||||
relative_path = file_path.relative_to(score_dir)
|
||||
print(f"\n--- File: {relative_path} ---")
|
||||
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
|
||||
if (
|
||||
file_path.suffix == ".json"
|
||||
or content.strip().startswith("{")
|
||||
or content.strip().startswith("[")
|
||||
):
|
||||
try:
|
||||
import json
|
||||
|
||||
lines = content.strip().split("\n")
|
||||
formatted_lines = []
|
||||
for line in lines:
|
||||
if line.strip():
|
||||
parsed = json.loads(line)
|
||||
formatted_lines.append(
|
||||
json.dumps(
|
||||
parsed, ensure_ascii=False, indent=2
|
||||
)
|
||||
)
|
||||
content = "\n".join(formatted_lines)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
print(content)
|
||||
|
||||
except UnicodeDecodeError:
|
||||
print(f"[Binary file, size: {file_path.stat().st_size} bytes]")
|
||||
except Exception as e:
|
||||
print(f"[Error reading file: {str(e)}]")
|
||||
|
||||
print(f"=== {eval_type.upper()} Evaluation Result Files End ===\n")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error capturing evaluation result files: {str(e)}")
|
||||
|
||||
def extract_tool_schema(tools):
|
||||
for i in range(len(tools)):
|
||||
tools[i]['function'].pop("response")
|
||||
return tools
|
||||
11
experiencemaker/cookbook/bfcl/requirements.md
Normal file
11
experiencemaker/cookbook/bfcl/requirements.md
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
## BFCL installation
|
||||
git clone https://github.com/ShishirPatil/gorilla.git
|
||||
|
||||
#### Change directory to the `berkeley-function-call-leaderboard`
|
||||
cd gorilla/berkeley-function-call-leaderboard
|
||||
|
||||
### Install the package in editable mode
|
||||
pip install -e .
|
||||
|
||||
#### Move the dataset to the data folder under bfcl
|
||||
cp -r gorilla/berkeley-function-call-leaderboard/bfcl_eval/data {/path/to/bfcl/data}
|
||||
119
experiencemaker/cookbook/bfcl/run_bfcl.py
Normal file
119
experiencemaker/cookbook/bfcl/run_bfcl.py
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
import os
|
||||
import time
|
||||
import ray
|
||||
from ray import logger
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv("../../.env")
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
from bfcl_agent import BFCLAgent
|
||||
|
||||
|
||||
def run_agent(dataset_name: str,
|
||||
experiment_suffix: str,
|
||||
max_workers: int,
|
||||
num_runs: int = 4,
|
||||
model_name: str = "qwen3-8b",
|
||||
data_path: str = "data/multiturn_data_base_val.jsonl",
|
||||
answer_path: Path = Path("data/possible_answer"),
|
||||
use_experience: bool = False,
|
||||
enable_thinking: bool = False,
|
||||
experience_base_url: str = "http://0.0.0.0:8001/",
|
||||
experience_workspace_id: str = "bfcl_8b_0725"):
|
||||
experiment_name = dataset_name + "_" + experiment_suffix
|
||||
path: Path = Path(f"./exp_result/{model_name}")
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
task_ids = [json.loads(l)["id"] for l in f]
|
||||
|
||||
result: list = []
|
||||
|
||||
def dump_file():
|
||||
with open(path / f"{experiment_name}.jsonl", "a") as f:
|
||||
for x in result:
|
||||
f.write(json.dumps(x) + "\n")
|
||||
|
||||
if max_workers > 1:
|
||||
future_list: list = []
|
||||
for i in range(max_workers):
|
||||
actor = BFCLAgent.remote(
|
||||
index=i,
|
||||
task_ids=task_ids[i::max_workers],
|
||||
experiment_name=experiment_name,
|
||||
data_path=data_path,
|
||||
answer_path=answer_path,
|
||||
model_name=model_name,
|
||||
num_runs=num_runs,
|
||||
use_experience=use_experience,
|
||||
enable_thinking=enable_thinking,
|
||||
experience_base_url=experience_base_url,
|
||||
experience_workspace_id=experience_workspace_id
|
||||
)
|
||||
future = actor.execute.remote()
|
||||
future_list.append(future)
|
||||
time.sleep(1)
|
||||
logger.info("submit complete")
|
||||
|
||||
for i, future in enumerate(future_list):
|
||||
t_result = ray.get(future)
|
||||
if t_result:
|
||||
if isinstance(t_result, list):
|
||||
result.extend(t_result)
|
||||
else:
|
||||
result.append(t_result)
|
||||
|
||||
logger.info(f"{i + 1}/{len(task_ids)} complete")
|
||||
dump_file()
|
||||
|
||||
else:
|
||||
for index, task_id in enumerate(task_ids):
|
||||
agent = BFCLAgent(index=index,
|
||||
task_ids=[task_id],
|
||||
experiment_name=experiment_name,
|
||||
num_runs=num_runs,
|
||||
model_name=model_name,
|
||||
data_path=data_path,
|
||||
answer_path=answer_path,
|
||||
enable_thinking=enable_thinking,
|
||||
use_experience=use_experience,
|
||||
experience_base_url=experience_base_url,
|
||||
experience_workspace_id=experience_workspace_id)
|
||||
task_results = agent.execute()
|
||||
if isinstance(task_results, list):
|
||||
result.extend(task_results)
|
||||
else:
|
||||
result.append(task_results)
|
||||
dump_file()
|
||||
|
||||
def main():
|
||||
max_workers = 4
|
||||
num_runs = 4 # Run each task 4 times
|
||||
use_experience = True
|
||||
experience_base_url = "http://0.0.0.0:8002/"
|
||||
experience_workspace_id = "bfcl_v1_extract_compare"
|
||||
if max_workers > 1:
|
||||
ray.init(num_cpus=4)
|
||||
for run_id in range(num_runs):
|
||||
run_agent(
|
||||
dataset_name="bfcl-multi-turn-base-val",
|
||||
experiment_suffix=f"0812-w-exp-w-think-extract-compare-recall-rewrite",
|
||||
model_name="qwen3-8b",
|
||||
max_workers=max_workers,
|
||||
num_runs=1,
|
||||
data_path="data/multiturn_data_base_val.jsonl",
|
||||
answer_path=Path("data/possible_answer"),
|
||||
enable_thinking=True,
|
||||
use_experience=use_experience,
|
||||
experience_base_url=experience_base_url,
|
||||
experience_workspace_id=experience_workspace_id,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
159
experiencemaker/cookbook/bfcl/run_exp_statistic.py
Normal file
159
experiencemaker/cookbook/bfcl/run_exp_statistic.py
Normal file
|
|
@ -0,0 +1,159 @@
|
|||
import json
|
||||
from pathlib import Path
|
||||
from collections import defaultdict
|
||||
import pandas as pd
|
||||
|
||||
from loguru import logger
|
||||
|
||||
|
||||
def calculate_best_at_k(scores: list, k: int) -> float:
|
||||
"""
|
||||
Calculate best@k
|
||||
Divide scores into groups of size k, take the maximum value in each group,
|
||||
then average these maximum values
|
||||
|
||||
Args:
|
||||
scores: List of after_score values for all runs of a task
|
||||
k: Group size
|
||||
|
||||
Returns:
|
||||
best@k value
|
||||
"""
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
group_maxs.append(max(group))
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def calculate_pass_at_k(scores: list, k: int) -> float:
|
||||
if len(scores) % k != 0:
|
||||
raise ValueError(f"Length of scores ({len(scores)}) must be divisible by k ({k})")
|
||||
|
||||
group_maxs = []
|
||||
for i in range(0, len(scores), k):
|
||||
group = scores[i:i + k]
|
||||
is_pass = 1.0 if max(group) >=1.0 else 0.0
|
||||
group_maxs.append(is_pass)
|
||||
|
||||
return sum(group_maxs) / len(group_maxs)
|
||||
|
||||
|
||||
def get_possible_k_values(total_runs: int) -> list:
|
||||
"""
|
||||
Get all possible k values (factors of total_runs)
|
||||
|
||||
Args:
|
||||
total_runs: Total number of runs
|
||||
|
||||
Returns:
|
||||
List of k values in descending order
|
||||
"""
|
||||
k_values = []
|
||||
for k in range(1, total_runs + 1):
|
||||
if total_runs % k == 0:
|
||||
k_values.append(k)
|
||||
return sorted(k_values, reverse=True) # Sort from large to small
|
||||
|
||||
|
||||
def run_exp_statistic():
|
||||
path: Path = Path(f"./no_exp_result/qwen3-8b")
|
||||
|
||||
# Store results for all experiments
|
||||
all_results = {}
|
||||
for file in [f for f in path.glob("*.jsonl")]:
|
||||
# Group results by task_id
|
||||
task_results = defaultdict(list)
|
||||
print(file)
|
||||
with open(file, "r") as f:
|
||||
for line in f:
|
||||
if not line.strip():
|
||||
continue
|
||||
data = json.loads(line)
|
||||
|
||||
if isinstance(data, list):
|
||||
for part_data in data:
|
||||
task_id = part_data["task_id"]
|
||||
after_score = part_data["reward"]
|
||||
task_results[task_id].append(after_score)
|
||||
else:
|
||||
task_id = data["task_id"]
|
||||
after_score = data["reward"]
|
||||
task_results[task_id].append(after_score)
|
||||
|
||||
if not task_results:
|
||||
logger.warning(f"No valid data found in file {file}")
|
||||
continue
|
||||
|
||||
# Check if each task has consistent number of runs
|
||||
run_counts = [len(scores) for scores in task_results.values()]
|
||||
if len(set(run_counts)) > 1:
|
||||
logger.warning(f"Inconsistent number of runs for different tasks in file {file}: {set(run_counts)}")
|
||||
continue
|
||||
|
||||
num_runs = run_counts[0]
|
||||
logger.info(f"File {file}: {len(task_results)} tasks, {num_runs} runs per task")
|
||||
|
||||
# Get all possible k values
|
||||
k_values = get_possible_k_values(num_runs)
|
||||
logger.info(f"Calculable best@k values: {k_values}")
|
||||
|
||||
# Calculate various best@k values
|
||||
file_results = {"file": file.name}
|
||||
|
||||
for k in k_values:
|
||||
best_at_k_scores = []
|
||||
pass_at_k_scores = []
|
||||
for task_id, scores in task_results.items():
|
||||
try:
|
||||
best_k_score = calculate_best_at_k(scores, k)
|
||||
pass_at_k_score = calculate_pass_at_k(scores, k)
|
||||
pass_at_k_scores.append(pass_at_k_score)
|
||||
best_at_k_scores.append(best_k_score)
|
||||
except ValueError as e:
|
||||
logger.error(f"Error calculating best@{k} for task {task_id}: {e}")
|
||||
continue
|
||||
|
||||
if best_at_k_scores:
|
||||
avg_best_at_k = sum(best_at_k_scores) / len(best_at_k_scores)
|
||||
file_results[f"best@{k}"] = avg_best_at_k
|
||||
logger.info(f"file={file.name} best@{k}={avg_best_at_k:.4f}")
|
||||
|
||||
if pass_at_k_scores:
|
||||
avg_pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores)
|
||||
file_results[f"pass@{k}"] = avg_pass_at_k
|
||||
logger.info(f"file={file.name} pass@{k}={avg_pass_at_k:.4f}")
|
||||
|
||||
all_results[file.name] = file_results
|
||||
|
||||
# Create and display table
|
||||
if all_results:
|
||||
df = pd.DataFrame(list(all_results.values()))
|
||||
df = df.set_index('file')
|
||||
|
||||
# Sort columns by the number in column name (best@8, best@4, best@2, best@1)
|
||||
# best_columns = [col for col in df.columns if col.startswith('best@')]
|
||||
best_columns = [col for col in df.columns]
|
||||
best_columns.sort(key=lambda x: x, reverse=False)
|
||||
df = df[best_columns]
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Experiment Results Summary Table")
|
||||
print("=" * 80)
|
||||
print(df.round(4))
|
||||
print("=" * 80)
|
||||
|
||||
# Save table to CSV
|
||||
output_path = path / "experiment_summary.csv"
|
||||
df.to_csv(output_path)
|
||||
logger.info(f"Results table saved to: {output_path}")
|
||||
else:
|
||||
logger.warning("No valid experiment results found")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
run_exp_statistic()
|
||||
Loading…
Add table
Reference in a new issue