ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
This commit is contained in:
+4
@@ -0,0 +1,4 @@
|
||||
from swift.trainers import TrainerFactory
|
||||
|
||||
TrainerFactory.TRAINER_MAPPING["aworld_grpo"] = 'train.examples.train_gaia_with_aworld_swift.AworldTrainer'
|
||||
TrainerFactory.TRAINING_ARGS_MAPPING["aworld_grpo"] = 'train_gaia_with_aworld_swift.trainers.GRPOConfig'
|
||||
+19
@@ -0,0 +1,19 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
from typing import Union
|
||||
|
||||
from aworld.agents.llm_agent import Agent
|
||||
from aworld.core.agent.swarm import Swarm
|
||||
from train.adapter.swift.aworld_agent_trainer import AworldTrainer
|
||||
|
||||
GAIA_SYSTEM_PROMPT = """
|
||||
You are an all-capable AI assistant, aimed at solving any task presented by the user.
|
||||
"""
|
||||
|
||||
|
||||
class GaiaTrainer(AworldTrainer):
|
||||
def build_agents(self) -> Union[Agent, Swarm]:
|
||||
return Agent(
|
||||
name="gaia_super_agent",
|
||||
system_prompt=GAIA_SYSTEM_PROMPT,
|
||||
)
|
||||
+135
@@ -0,0 +1,135 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import re
|
||||
import string
|
||||
from typing import List
|
||||
|
||||
from swift.plugin import ORM, orms, rm_plugins
|
||||
from swift.utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
"""
|
||||
Step 1: Define a Reward Class
|
||||
Implement your custom reward calculation logic within the __call__ method.
|
||||
The method accepts the model's output completions and dataset columns (passed as kwargs) as input parameters.
|
||||
|
||||
Step 2: Register the Reward Class in orms
|
||||
For example:
|
||||
python orms['external_math_acc'] = MathAccuracy
|
||||
|
||||
Step 3: Configure the Arguments
|
||||
Use the following arguments when running the script:
|
||||
bash --plugin /path/to/plugin.py --reward_funcs external_math_acc
|
||||
"""
|
||||
|
||||
|
||||
class GaiaAnswerMatch(ORM):
|
||||
def __call__(self, completions, solution, **kwargs) -> List[float]:
|
||||
pattern = r'<answer>(.*?)</answer>'
|
||||
rewards = []
|
||||
logger.info(f"GaiaAnswerMatch|completions:{completions}, comp_match:{solution}")
|
||||
for content, sol in zip(completions, solution):
|
||||
comp_match = re.search(pattern, content, re.DOTALL | re.MULTILINE)
|
||||
logger.info(f"GaiaAnswerMatch|content:{content}, comp_match:{comp_match}, sol:{sol}")
|
||||
if not comp_match:
|
||||
rewards.append(0.0)
|
||||
continue
|
||||
comp_answer = comp_match.group(1).strip()
|
||||
|
||||
if question_scorer(comp_answer, sol):
|
||||
rewards.append(1.0)
|
||||
else:
|
||||
rewards.append(0.0)
|
||||
|
||||
return rewards
|
||||
|
||||
|
||||
class GaiaFormat(ORM):
|
||||
def __call__(self, completions, **kwargs) -> List[float]:
|
||||
"""Reward function that checks if the completion has a specific format."""
|
||||
pattern = r'<answer>[\s\S]*?</answer>'
|
||||
matches = [re.search(pattern, content, re.DOTALL | re.MULTILINE) for content in completions]
|
||||
reward = [0.1 if match else 0.0 for match in matches]
|
||||
return reward
|
||||
|
||||
|
||||
orms['external_gaia_answer_reward'] = GaiaAnswerMatch
|
||||
orms['external_gaia_format_reward'] = GaiaFormat
|
||||
|
||||
|
||||
def split_string(
|
||||
s: str,
|
||||
char_list: list[str] = [",", ";"],
|
||||
) -> list[str]:
|
||||
pattern = f"[{''.join(char_list)}]"
|
||||
return re.split(pattern, s)
|
||||
|
||||
|
||||
def normalize_str(input_str, remove_punct=True) -> str:
|
||||
no_spaces = re.sub(r"\s", "", input_str)
|
||||
|
||||
# Remove punctuation, if specified.
|
||||
if remove_punct:
|
||||
translator = str.maketrans("", "", string.punctuation)
|
||||
return no_spaces.lower().translate(translator)
|
||||
else:
|
||||
return no_spaces.lower()
|
||||
|
||||
|
||||
def normalize_number_str(number_str: str) -> float:
|
||||
# we replace these common units and commas to allow
|
||||
# conversion to float
|
||||
for char in ["$", "%", ","]:
|
||||
number_str = number_str.replace(char, "")
|
||||
try:
|
||||
return float(number_str)
|
||||
except ValueError:
|
||||
# print(f"String {number_str} cannot be normalized to number str.")
|
||||
return float("inf")
|
||||
|
||||
|
||||
def question_scorer(
|
||||
model_answer: str,
|
||||
ground_truth: str,
|
||||
) -> bool:
|
||||
def is_float(element: any) -> bool:
|
||||
try:
|
||||
float(element)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
if model_answer is None:
|
||||
model_answer = "None"
|
||||
|
||||
# if gt is a number
|
||||
if is_float(ground_truth):
|
||||
# print(f"Evaluating {model_answer} as a number.")
|
||||
normalized_answer = normalize_number_str(model_answer)
|
||||
return normalized_answer == float(ground_truth)
|
||||
# if gt is a list
|
||||
elif any(char in ground_truth for char in [",", ";"]):
|
||||
# question with the fish: normalization removes punct
|
||||
gt_elems = split_string(ground_truth)
|
||||
ma_elems = split_string(model_answer)
|
||||
|
||||
# check length is the same
|
||||
if len(gt_elems) != len(ma_elems):
|
||||
return False
|
||||
|
||||
# compare each element as float or str
|
||||
comparisons = []
|
||||
for ma_elem, gt_elem in zip(ma_elems, gt_elems):
|
||||
if is_float(gt_elem):
|
||||
normalized_ma_elem = normalize_number_str(ma_elem)
|
||||
comparisons.append(normalized_ma_elem == float(gt_elem))
|
||||
else:
|
||||
# we do not remove punct since comparisons can include punct
|
||||
comparisons.append(
|
||||
normalize_str(ma_elem, remove_punct=False)
|
||||
== normalize_str(gt_elem, remove_punct=False)
|
||||
)
|
||||
return all(comparisons)
|
||||
# if gt is a str
|
||||
else:
|
||||
return normalize_str(model_answer) == normalize_str(ground_truth)
|
||||
@@ -0,0 +1,2 @@
|
||||
- name: gaia_agent
|
||||
_target_: train.examples.train_gaia_with_aworld_verl.custom_agent_loop.GaiaAgentLoop
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
from typing import Union
|
||||
|
||||
from aworld.agents.llm_agent import Agent
|
||||
from aworld.config import AgentConfig
|
||||
from aworld.core.agent.swarm import Swarm
|
||||
|
||||
from train.adapter.verl.aworld_agent_loop import AworldAgentLoop
|
||||
from train.adapter.verl.common import get_agent_tool_env_and_servers
|
||||
from env.train_env import TranEnv
|
||||
|
||||
GAIA_SYSTEM_PROMPT = """
|
||||
You are an all-capable AI assistant, aimed at solving any task presented by the user.
|
||||
"""
|
||||
|
||||
|
||||
class GaiaAgentLoop(AworldAgentLoop):
|
||||
async def build_agents(self) -> Union[Agent, Swarm]:
|
||||
gaia_env_config, gaia_env_servers = get_agent_tool_env_and_servers()
|
||||
|
||||
return Agent(
|
||||
conf=AgentConfig(
|
||||
llm_model_name=await self.get_llm_server_model_name(),
|
||||
llm_base_url=await self.get_llm_server_address(),
|
||||
llm_api_key="",
|
||||
),
|
||||
name="gaia_super_agent",
|
||||
system_prompt=GAIA_SYSTEM_PROMPT,
|
||||
|
||||
# MCP tool configuration for the agent
|
||||
mcp_config=gaia_env_config,
|
||||
mcp_servers=gaia_env_servers,
|
||||
)
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def load_gaia_dataset(path: str, split: str = "validation", total_num_dataset: int = 300):
|
||||
data_dir = Path(path) / split
|
||||
|
||||
split_dataset = []
|
||||
rl_dataset = {
|
||||
"prompt": [],
|
||||
"data_source": [],
|
||||
"ability": [],
|
||||
"reward_model": [],
|
||||
"extra_info": [],
|
||||
"agent_name": [],
|
||||
}
|
||||
cnt = 0
|
||||
with open(data_dir / "metadata.jsonl", "r", encoding="utf-8") as metaf:
|
||||
lines = metaf.readlines()
|
||||
for line in lines:
|
||||
data = json.loads(line)
|
||||
if data["task_id"] == "0-0-0-0-0":
|
||||
continue
|
||||
if data["file_name"]:
|
||||
data["file_name"] = data_dir / data["file_name"]
|
||||
split_dataset.append(data)
|
||||
rl_dataset["prompt"].append(data["Question"])
|
||||
rl_dataset["extra_info"].append(
|
||||
{"task_id": data["task_id"], "split": split, "level": data["Level"], "answer": data["Final answer"]}
|
||||
)
|
||||
rl_dataset["agent_name"].append("gaia_agent")
|
||||
rl_dataset["data_source"].append("gaia")
|
||||
rl_dataset["ability"].append("agi")
|
||||
rl_dataset["reward_model"].append({"style": "GAIA", "ground_truth": data['Final answer']})
|
||||
|
||||
cnt += 1
|
||||
if cnt >= total_num_dataset:
|
||||
break
|
||||
|
||||
rl_dataset = pd.DataFrame(data=rl_dataset)
|
||||
return rl_dataset
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="GAIA Dataset Generator")
|
||||
parser.add_argument("--train_size", type=int, default=300, help="Number of training samples")
|
||||
parser.add_argument("--test_size", type=int, default=100, help="Number of testing samples")
|
||||
parser.add_argument("--output_dir", default="gaia_data/", help="Directory to save the dataset")
|
||||
parser.add_argument("--dataset_path", default="./gaia_dataset", help="GAIA dataset path")
|
||||
args = parser.parse_args()
|
||||
|
||||
gaia_dataset_path = args.dataset_path
|
||||
|
||||
train_dataset = load_gaia_dataset(path=gaia_dataset_path, split="validation", total_num_dataset=args.train_size)
|
||||
test_dataset = load_gaia_dataset(path=gaia_dataset_path, split="test", total_num_dataset=args.test_size)
|
||||
|
||||
# Make sure the dataset directory exists
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# Save the datasets to parquet files
|
||||
train_dataset.to_parquet(os.path.join(args.output_dir, "train.parquet"))
|
||||
test_dataset.to_parquet(os.path.join(args.output_dir, "test.parquet"))
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
import re
|
||||
import string
|
||||
from aworld.logs.util import logger
|
||||
|
||||
|
||||
def normalize_number_str(number_str: str) -> float:
|
||||
# we replace these common units and commas to allow
|
||||
# conversion to float
|
||||
for char in ["$", "%", ","]:
|
||||
number_str = number_str.replace(char, "")
|
||||
try:
|
||||
return float(number_str)
|
||||
except ValueError:
|
||||
# print(f"String {number_str} cannot be normalized to number str.")
|
||||
return float("inf")
|
||||
|
||||
def split_string(
|
||||
s: str,
|
||||
char_list: list[str] = [",", ";"],
|
||||
) -> list[str]:
|
||||
pattern = f"[{''.join(char_list)}]"
|
||||
return re.split(pattern, s)
|
||||
|
||||
def normalize_str(input_str, remove_punct=True) -> str:
|
||||
"""
|
||||
Normalize a string by:
|
||||
- Removing all white spaces
|
||||
- Optionally removing punctuation (if remove_punct is True)
|
||||
- Converting to lowercase
|
||||
Parameters:
|
||||
- input_str: str, the string to normalize
|
||||
- remove_punct: bool, whether to remove punctuation (default: True)
|
||||
Returns:
|
||||
- str, the normalized string
|
||||
"""
|
||||
# Remove all white spaces. Required e.g for seagull vs. sea gull
|
||||
no_spaces = re.sub(r"\s", "", input_str)
|
||||
|
||||
# Remove punctuation, if specified.
|
||||
if remove_punct:
|
||||
translator = str.maketrans("", "", string.punctuation)
|
||||
return no_spaces.lower().translate(translator)
|
||||
else:
|
||||
return no_spaces.lower()
|
||||
|
||||
def question_scorer(
|
||||
model_answer: str,
|
||||
ground_truth: str,
|
||||
) -> bool:
|
||||
def is_float(element: any) -> bool:
|
||||
try:
|
||||
float(element)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
if model_answer is None:
|
||||
model_answer = "None"
|
||||
|
||||
# if gt is a number
|
||||
if is_float(ground_truth):
|
||||
# print(f"Evaluating {model_answer} as a number.")
|
||||
normalized_answer = normalize_number_str(model_answer)
|
||||
return normalized_answer == float(ground_truth)
|
||||
|
||||
# if gt is a list
|
||||
elif any(char in ground_truth for char in [",", ";"]):
|
||||
# print(f"Evaluating {model_answer} as a comma separated list.")
|
||||
# question with the fish: normalization removes punct
|
||||
|
||||
gt_elems = split_string(ground_truth)
|
||||
ma_elems = split_string(model_answer)
|
||||
|
||||
# check length is the same
|
||||
if len(gt_elems) != len(ma_elems):
|
||||
# warnings.warn(
|
||||
# "Answer lists have different lengths, returning False.", UserWarning
|
||||
# )
|
||||
return False
|
||||
|
||||
# compare each element as float or str
|
||||
comparisons = []
|
||||
for ma_elem, gt_elem in zip(ma_elems, gt_elems):
|
||||
if is_float(gt_elem):
|
||||
normalized_ma_elem = normalize_number_str(ma_elem)
|
||||
comparisons.append(normalized_ma_elem == float(gt_elem))
|
||||
else:
|
||||
# we do not remove punct since comparisons can include punct
|
||||
comparisons.append(
|
||||
normalize_str(ma_elem, remove_punct=False)
|
||||
== normalize_str(gt_elem, remove_punct=False)
|
||||
)
|
||||
return all(comparisons)
|
||||
|
||||
# if gt is a str
|
||||
else:
|
||||
# print(f"Evaluating {model_answer} as a string.")
|
||||
return normalize_str(model_answer) == normalize_str(ground_truth)
|
||||
|
||||
|
||||
def gaia_reward_func(data_source, solution_str, ground_truth, extra_info=None):
|
||||
pattern = r'<answer>(.*?)</answer>'
|
||||
comp_match = re.search(pattern, solution_str, re.DOTALL | re.MULTILINE)
|
||||
|
||||
if not comp_match:
|
||||
return 0.0
|
||||
else:
|
||||
comp_answer = comp_match.group(1).strip()
|
||||
logger.info(f"comp_answer: {comp_answer}, ground_truth: {ground_truth}")
|
||||
if question_scorer(comp_answer, ground_truth):
|
||||
return 1.0
|
||||
else:
|
||||
return 0.0
|
||||
@@ -0,0 +1,139 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
set -xeuo pipefail
|
||||
|
||||
# ================= cluster topology =================
|
||||
export GPUS_PER_NODE=${SLURM_GPUS_ON_NODE:-${GPUS_PER_NODE:-1}} # GPUs on this node
|
||||
NNODES=${SLURM_JOB_NUM_NODES:-${NNODES:-1}}
|
||||
export NNODES
|
||||
export RAY_NUM_NODES=$NNODES
|
||||
|
||||
echo "Using $NNODES nodes and $GPUS_PER_NODE GPUs per node..."
|
||||
|
||||
# ================= data/model/tool =================
|
||||
HDFS_ROOT=${HDFS_ROOT:-$PWD}
|
||||
DATA_ROOT=${DATA_ROOT:-$PWD}
|
||||
|
||||
# Prefer local model if present, otherwise fall back to HF hub path
|
||||
model_path=${model_path:-$DATA_ROOT/Qwen/Qwen3-4B}
|
||||
if [ ! -d "$model_path" ]; then
|
||||
model_path=Qwen/Qwen3-4B
|
||||
fi
|
||||
|
||||
# Use the default output directory produced by create_dataset.py
|
||||
train_files=$DATA_ROOT/datasets/train.parquet
|
||||
test_files=$DATA_ROOT/datasets/test.parquet
|
||||
|
||||
# =================== custom ===================
|
||||
path_to_train="/your/path/to/train"
|
||||
reward_fn_name=gaia_reward_func
|
||||
reward_fn_file_path=${path_to_train}/examples/train_gaia_with_aworld_verl/metrics/gaia_reward_function.py
|
||||
|
||||
# Agent config
|
||||
agent_loop_config_path=${path_to_train}/examples/train_gaia_with_aworld_verl/agent.yaml
|
||||
|
||||
# set dummy_tool_config_path to enable auto_tool_choice
|
||||
dummy_tool_config_path=${path_to_train}/examples/verl/configs/dummy_tool_config.yaml
|
||||
|
||||
# =================== wandb ===================
|
||||
project_name=gaia
|
||||
experiment_name=qwe3
|
||||
default_local_dir=$DATA_ROOT/checkpoint/$experiment_name
|
||||
|
||||
# ================= algorithm =================
|
||||
adv_estimator=grpo
|
||||
|
||||
use_kl_in_reward=false
|
||||
kl_coef=0.0
|
||||
use_kl_loss=false
|
||||
kl_loss_coef=0.0
|
||||
|
||||
clip_ratio_low=0.2
|
||||
clip_ratio_high=0.28
|
||||
|
||||
max_turns=8
|
||||
max_prompt_length=1024
|
||||
max_response_length=2048
|
||||
actor_lr=1e-6
|
||||
|
||||
train_batch_size=1
|
||||
ppo_mini_batch_size=1
|
||||
n_resp_per_prompt=1
|
||||
n_resp_per_prompt_val=1
|
||||
|
||||
# =================== logging ===================
|
||||
export RAY_LOGGING_LEVEL=DEBUG
|
||||
export HYDRA_FULL_ERROR=1
|
||||
|
||||
# ================= performance =================
|
||||
export NCCL_IBEXT_DISABLE=1
|
||||
export NCCL_NVLS_ENABLE=1
|
||||
export NCCL_IB_HCA=mlx5
|
||||
export UCX_NET_DEVICES=mlx5_0:1,mlx5_1:1,mlx5_2:1,mlx5_3:1,mlx5_4:1,mlx5_5:1,mlx5_6:1,mlx5_7:1
|
||||
export VLLM_USE_V1=1
|
||||
export VLLM_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
infer_tp=1 # vLLM tensor parallel size
|
||||
train_sp=1 # Ulysses sequence parallel size for actor
|
||||
offload=true
|
||||
|
||||
actor_max_token_len_per_gpu=$(( (max_prompt_length + max_response_length) * 4 ))
|
||||
log_prob_max_token_len_per_gpu=$(( actor_max_token_len_per_gpu * 2 ))
|
||||
|
||||
train_files="['$train_files']"
|
||||
test_files="['$test_files']"
|
||||
|
||||
python3 -m verl.trainer.main_ppo \
|
||||
algorithm.adv_estimator=$adv_estimator \
|
||||
algorithm.use_kl_in_reward=$use_kl_in_reward \
|
||||
algorithm.kl_ctrl.kl_coef=$kl_coef \
|
||||
data.train_files="$train_files" \
|
||||
data.val_files="$test_files" \
|
||||
data.return_raw_chat=true \
|
||||
data.train_batch_size=$train_batch_size \
|
||||
data.max_prompt_length=$max_prompt_length \
|
||||
data.max_response_length=$max_response_length \
|
||||
data.filter_overlong_prompts=true \
|
||||
data.truncation='error' \
|
||||
actor_rollout_ref.model.path="$model_path" \
|
||||
actor_rollout_ref.model.use_remove_padding=true \
|
||||
actor_rollout_ref.model.enable_gradient_checkpointing=true \
|
||||
actor_rollout_ref.actor.use_kl_loss=$use_kl_loss \
|
||||
actor_rollout_ref.actor.kl_loss_coef=$kl_loss_coef \
|
||||
actor_rollout_ref.actor.clip_ratio_low=$clip_ratio_low \
|
||||
actor_rollout_ref.actor.clip_ratio_high=$clip_ratio_high \
|
||||
actor_rollout_ref.actor.clip_ratio_c=10.0 \
|
||||
actor_rollout_ref.actor.optim.lr=$actor_lr \
|
||||
actor_rollout_ref.actor.use_dynamic_bsz=true \
|
||||
actor_rollout_ref.actor.ppo_mini_batch_size=$ppo_mini_batch_size \
|
||||
actor_rollout_ref.actor.ppo_max_token_len_per_gpu=$actor_max_token_len_per_gpu \
|
||||
actor_rollout_ref.actor.ulysses_sequence_parallel_size=$train_sp \
|
||||
actor_rollout_ref.actor.fsdp_config.param_offload=$offload \
|
||||
actor_rollout_ref.actor.fsdp_config.optimizer_offload=$offload \
|
||||
actor_rollout_ref.ref.log_prob_max_token_len_per_gpu=$log_prob_max_token_len_per_gpu \
|
||||
actor_rollout_ref.rollout.name=vllm \
|
||||
actor_rollout_ref.rollout.mode=async \
|
||||
actor_rollout_ref.rollout.tensor_model_parallel_size=$infer_tp \
|
||||
actor_rollout_ref.rollout.multi_turn.max_user_turns=$max_turns \
|
||||
actor_rollout_ref.rollout.multi_turn.max_assistant_turns=$max_turns \
|
||||
actor_rollout_ref.rollout.multi_turn.format=hermes \
|
||||
actor_rollout_ref.rollout.agent.agent_loop_config_path=$agent_loop_config_path \
|
||||
actor_rollout_ref.rollout.gpu_memory_utilization=0.75 \
|
||||
actor_rollout_ref.rollout.n=$n_resp_per_prompt \
|
||||
actor_rollout_ref.rollout.val_kwargs.top_p=0.6 \
|
||||
actor_rollout_ref.rollout.val_kwargs.temperature=1.0 \
|
||||
actor_rollout_ref.rollout.val_kwargs.n=$n_resp_per_prompt_val \
|
||||
actor_rollout_ref.rollout.multi_turn.tool_config_path=$dummy_tool_config_path \
|
||||
custom_reward_function.path="${reward_fn_file_path}"\
|
||||
custom_reward_function.name="${reward_fn_name}"\
|
||||
trainer.logger=console \
|
||||
trainer.project_name=$project_name \
|
||||
trainer.experiment_name=$experiment_name \
|
||||
trainer.n_gpus_per_node="$GPUS_PER_NODE" \
|
||||
trainer.val_before_train=true \
|
||||
trainer.log_val_generations=50 \
|
||||
trainer.nnodes="$NNODES" \
|
||||
trainer.save_freq=-1 \
|
||||
trainer.default_local_dir="$default_local_dir" \
|
||||
trainer.test_freq=5 \
|
||||
trainer.total_epochs=1 "$@"
|
||||
Reference in New Issue
Block a user