mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Split the fields only the rollout executor reads into rollout-only traits
Every trait whose fields are read only by the rollout executor process and its rollout functions gets a sibling RolloutOnlyConfig class next to it that carries those fields, and only RolloutConfig and AllConfig mix the siblings in, so a change to one of them restarts the rollout executor without touching the trainer, inference controller or multi-LoRA controller configs.
This commit is contained in:
@@ -150,22 +150,6 @@ class AlgoConfig(BaseConfig):
|
||||
gamma: A[float, Arg(help="PPO GAE gamma")] = 1.0
|
||||
lambd: A[float, Arg(help="PPO GAE lambd")] = 1.0
|
||||
normalize_advantages: A[bool, Arg()] = False
|
||||
grpo_std_normalization: A[
|
||||
bool,
|
||||
Arg(
|
||||
cli_name="--disable-grpo-std-normalization",
|
||||
action="store_false",
|
||||
help="from Dr.GRPO https://arxiv.org/pdf/2503.20783",
|
||||
),
|
||||
] = True
|
||||
rewards_normalization: A[
|
||||
bool,
|
||||
Arg(
|
||||
cli_name="--disable-rewards-normalization",
|
||||
action="store_false",
|
||||
help="Disable rewards normalization",
|
||||
),
|
||||
] = True
|
||||
use_rollout_entropy: A[
|
||||
bool,
|
||||
Arg(
|
||||
@@ -273,3 +257,22 @@ class AlgoConfig(BaseConfig):
|
||||
float,
|
||||
Arg(help="The threshold for Off-Policy Sequence Masking (OPSM)."),
|
||||
] = 1e-4
|
||||
|
||||
|
||||
class AlgoRolloutOnlyConfig(BaseConfig):
|
||||
grpo_std_normalization: A[
|
||||
bool,
|
||||
Arg(
|
||||
cli_name="--disable-grpo-std-normalization",
|
||||
action="store_false",
|
||||
help="from Dr.GRPO https://arxiv.org/pdf/2503.20783",
|
||||
),
|
||||
] = True
|
||||
rewards_normalization: A[
|
||||
bool,
|
||||
Arg(
|
||||
cli_name="--disable-rewards-normalization",
|
||||
action="store_false",
|
||||
help="Disable rewards normalization",
|
||||
),
|
||||
] = True
|
||||
|
||||
@@ -6,9 +6,6 @@ class CiConfig(BaseConfig):
|
||||
config_snapshot_name: A[str | None, Arg()] = None
|
||||
ci_enable_metrics_capture: bool
|
||||
|
||||
ci_inject_missing_prefetched_batch_bug: A[
|
||||
bool, Arg(help="Discard the restored prefetched batch to test sample ownership failure detection.")
|
||||
] = False
|
||||
ci_test: A[bool, Arg()] = False
|
||||
ci_tito_special_token_count_threshold: A[
|
||||
float,
|
||||
@@ -26,11 +23,17 @@ class CiConfig(BaseConfig):
|
||||
ci_metric_checker_expect_num: A[
|
||||
int | None, Arg(help="Require exactly this many eval checks, all meeting the CI threshold.")
|
||||
] = None
|
||||
ci_assert_prefill_lag_max: A[
|
||||
int | None,
|
||||
Arg(help="Require every rollout's prompt KV to lag its decode weight version by at most this much."),
|
||||
] = None
|
||||
ci_save_grad_norm: A[str | None, Arg()] = None
|
||||
ci_load_grad_norm: A[str | None, Arg()] = None
|
||||
ci_save_model_hash: A[bool, Arg()] = False
|
||||
ci_check_model_hash: A[bool, Arg()] = False
|
||||
|
||||
|
||||
class CiRolloutOnlyConfig(BaseConfig):
|
||||
ci_inject_missing_prefetched_batch_bug: A[
|
||||
bool, Arg(help="Discard the restored prefetched batch to test sample ownership failure detection.")
|
||||
] = False
|
||||
ci_assert_prefill_lag_max: A[
|
||||
int | None,
|
||||
Arg(help="Require every rollout's prompt KV to lag its decode weight version by at most this much."),
|
||||
] = None
|
||||
|
||||
@@ -9,16 +9,6 @@ from miles.utils.env_report.launcher_report import LAUNCHER_REPORT_ENV_VAR
|
||||
|
||||
# debug
|
||||
class DebugConfig(BaseConfig):
|
||||
save_debug_rollout_data: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Save the rollout data to this path for debugging. "
|
||||
"The file will be saved to `save_debug_rollout_data.format(rollout_id)`, "
|
||||
"so the template must contain the `{rollout_id}` placeholder."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
save_debug_trajectory_data: A[
|
||||
str | None,
|
||||
Arg(
|
||||
@@ -39,10 +29,6 @@ class DebugConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
load_debug_rollout_data_subsample: A[
|
||||
float | None,
|
||||
Arg(help="Subsample a portion of the debug rollout data for faster debugging."),
|
||||
] = None
|
||||
debug_rollout_only: A[
|
||||
bool,
|
||||
Arg(
|
||||
@@ -257,24 +243,6 @@ class DebugConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
ci_inject_rollout_data_start_rollout_id: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"First rollout_id whose training data is replaced by the " "--ci-inject-rollout-data-path recordings."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
ci_inject_rollout_data_min_match_ratio: A[
|
||||
float,
|
||||
Arg(
|
||||
help=(
|
||||
"Minimum mean response-token match ratio between the discarded generated "
|
||||
"data and the injected recording. Below this the engine weights are considered "
|
||||
"wrong (legitimate ulp-level drift only flips occasional sampled tokens)."
|
||||
)
|
||||
),
|
||||
] = 0.9
|
||||
env_report: A[
|
||||
str,
|
||||
Arg(
|
||||
@@ -311,3 +279,38 @@ class DebugConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = False
|
||||
|
||||
|
||||
class DebugRolloutOnlyConfig(BaseConfig):
|
||||
save_debug_rollout_data: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Save the rollout data to this path for debugging. "
|
||||
"The file will be saved to `save_debug_rollout_data.format(rollout_id)`, "
|
||||
"so the template must contain the `{rollout_id}` placeholder."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
load_debug_rollout_data_subsample: A[
|
||||
float | None,
|
||||
Arg(help="Subsample a portion of the debug rollout data for faster debugging."),
|
||||
] = None
|
||||
ci_inject_rollout_data_start_rollout_id: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"First rollout_id whose training data is replaced by the " "--ci-inject-rollout-data-path recordings."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
ci_inject_rollout_data_min_match_ratio: A[
|
||||
float,
|
||||
Arg(
|
||||
help=(
|
||||
"Minimum mean response-token match ratio between the discarded generated "
|
||||
"data and the injected recording. Below this the engine weights are considered "
|
||||
"wrong (legitimate ulp-level drift only flips occasional sampled tokens)."
|
||||
)
|
||||
),
|
||||
] = 0.9
|
||||
|
||||
@@ -4,22 +4,8 @@ from miles.utils.eval_config import EvalDatasetConfig
|
||||
|
||||
|
||||
class EvalConfig(BaseConfig):
|
||||
eval_datasets: list[EvalDatasetConfig]
|
||||
eval_uses_snapshots: bool
|
||||
|
||||
eval_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to the eval fn. Two kinds fit here. A rollout fn generates against the "
|
||||
"engines the framework hands it: the training engines, or the dedicated fleet "
|
||||
"when --eval-num-gpus is set. A CheckpointEvalFn subclass gets the snapshot "
|
||||
"path instead and owns the rest itself — weight delivery, endpoint, generation. "
|
||||
"If not set, defaults to --rollout-function-path."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
eval_prompt_data: A[
|
||||
list[str] | None,
|
||||
Arg(
|
||||
@@ -53,7 +39,6 @@ class EvalConfig(BaseConfig):
|
||||
eval_top_p: A[float | None, Arg()] = None
|
||||
eval_top_k: A[int | None, Arg()] = None
|
||||
eval_max_response_len: A[int | None, Arg()] = None
|
||||
eval_max_prompt_len: A[int | None, Arg()] = None
|
||||
eval_min_new_tokens: A[int | None, Arg()] = None
|
||||
eval_max_context_len: A[int | None, Arg()] = None
|
||||
eval_hf_dir: A[
|
||||
@@ -89,3 +74,21 @@ class EvalConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = 2
|
||||
|
||||
|
||||
class EvalRolloutOnlyConfig(BaseConfig):
|
||||
eval_datasets: list[EvalDatasetConfig]
|
||||
eval_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to the eval fn. Two kinds fit here. A rollout fn generates against the "
|
||||
"engines the framework hands it: the training engines, or the dedicated fleet "
|
||||
"when --eval-num-gpus is set. A CheckpointEvalFn subclass gets the snapshot "
|
||||
"path instead and owns the rest itself — weight delivery, endpoint, generation. "
|
||||
"If not set, defaults to --rollout-function-path."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
eval_max_prompt_len: A[int | None, Arg()] = None
|
||||
|
||||
@@ -27,6 +27,19 @@ class OnPolicyDistillationConfig(BaseConfig):
|
||||
float,
|
||||
Arg(help="On-policy distillation KL penalty coefficient. Default is 1.0."),
|
||||
] = 1.0
|
||||
opd_teacher_load: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The checkpoint for OPD teacher model. Required when --opd-type=megatron. "
|
||||
"The teacher model should have the same architecture as policy/ref model."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
opd_teacher_ckpt_step: A[int | None, Arg(help="The checkpoint step for OPD teacher model.")] = None
|
||||
|
||||
|
||||
class OnPolicyDistillationRolloutOnlyConfig(BaseConfig):
|
||||
opd_log_prob_top_k: A[
|
||||
int,
|
||||
Arg(
|
||||
@@ -87,13 +100,3 @@ class OnPolicyDistillationConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = "opd_teacher"
|
||||
opd_teacher_load: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The checkpoint for OPD teacher model. Required when --opd-type=megatron. "
|
||||
"The teacher model should have the same architecture as policy/ref model."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
opd_teacher_ckpt_step: A[int | None, Arg(help="The checkpoint step for OPD teacher model.")] = None
|
||||
|
||||
@@ -13,12 +13,6 @@ class RewardModelConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
eval_reward_key: A[str | None, Arg(help="The eval variant for --reward-key")] = None
|
||||
group_rm: A[bool, Arg(help="Whether to do rm on a whole group.")] = False
|
||||
rm_url: A[
|
||||
str | None,
|
||||
Arg(help="URL for the reward model service for --rm-type remote_rm, e.g. http://localhost:8000"),
|
||||
] = None
|
||||
custom_rm_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
@@ -29,6 +23,15 @@ class RewardModelConfig(BaseConfig):
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
|
||||
class RewardModelRolloutOnlyConfig(BaseConfig):
|
||||
eval_reward_key: A[str | None, Arg(help="The eval variant for --reward-key")] = None
|
||||
group_rm: A[bool, Arg(help="Whether to do rm on a whole group.")] = False
|
||||
rm_url: A[
|
||||
str | None,
|
||||
Arg(help="URL for the reward model service for --rm-type remote_rm, e.g. http://localhost:8000"),
|
||||
] = None
|
||||
custom_reward_post_process_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
|
||||
+218
-215
@@ -65,20 +65,6 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
# Sampling values reach the engine per request only: the built-in generate path sends them
|
||||
# itself and the session server fills fields an agent omits from its session's defaults.
|
||||
# They are never engine launch arguments: an engine shared by rollout and eval has no single default.
|
||||
namespaced_radix_cache: A[
|
||||
bool | None,
|
||||
Arg(
|
||||
action=argparse.BooleanOptionalAction,
|
||||
help=(
|
||||
"Whether every generation request carries a radix cache key naming the rollout call "
|
||||
"the sample started under, so prefix KV computed under old weights cannot serve "
|
||||
"samples of a later call. Defaults to true when --fully-async is combined with "
|
||||
"--pause-generation-mode in_place, where the engine never flushes the cache and the "
|
||||
"staleness of a shared prompt is otherwise unbounded; an explicit "
|
||||
"--no-namespaced-radix-cache is respected."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_temperature: A[
|
||||
float,
|
||||
Arg(help="the temperature for the inference engine during rollout."),
|
||||
@@ -102,25 +88,6 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = -1
|
||||
rollout_max_context_len: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The maximum context size for the inference engine during rollout."
|
||||
"It should no exceed the `max_position_embeddinds` in Huggingface model's `config.json`"
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_max_prompt_len: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The maximum length of the prompt for the inference engine during rollout. "
|
||||
"If set, we will filter out the long prompts during initialization of the global dataset. "
|
||||
"This is not recommended if the dataset is large."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_max_response_len: A[
|
||||
int | None,
|
||||
Arg(
|
||||
@@ -130,51 +97,6 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_skip_special_tokens: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to skip special tokens in the response during rollout. "
|
||||
"This is useful when you want to use the response as a prompt for the next rollout."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
rollout_stop: A[
|
||||
list[str] | None,
|
||||
Arg(
|
||||
type_parser=str,
|
||||
nargs="+",
|
||||
help=(
|
||||
"The stop words for the inference engine during rollout. "
|
||||
"It can be a list of strings or a single string. "
|
||||
"It may be hard to pass special tokens in command line, in that case rollout_stop_token_ids can be used."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_stop_token_ids: A[
|
||||
list[int] | None,
|
||||
Arg(
|
||||
type_parser=int,
|
||||
nargs="+",
|
||||
help=(
|
||||
"The stop token ids for the inference engine during rollout. "
|
||||
"It can be a list of integers or a single integer."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_shuffle: A[
|
||||
bool,
|
||||
Arg(help="Whether to shuffle the prompts during rollout."),
|
||||
] = False
|
||||
rollout_seed: A[
|
||||
int,
|
||||
Arg(
|
||||
help=(
|
||||
"The seed for the random number generator during rollout. "
|
||||
"This is used to shuffle the prompts and also for the random sampling of the prompts."
|
||||
)
|
||||
),
|
||||
] = 42
|
||||
object_store_backend: A[
|
||||
str,
|
||||
Arg(
|
||||
@@ -194,103 +116,7 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
Arg(help="Number of Mooncake memory replicas for each stored object."),
|
||||
] = 1
|
||||
|
||||
# sampling
|
||||
over_sampling_batch_size: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"This defines the granularity of the sampling batch in the rollout function. "
|
||||
"When the number of available samples falls below the target, a sampling "
|
||||
"operation of size over_sampling_batch_size will be triggered."
|
||||
"Regardless of whether partial rollout is used or filters are applied, "
|
||||
"the sampling granularity is always determined by this value. "
|
||||
"If this value is None, rollout_batch_size will be used as the default over_sampling_batch_size."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
dynamic_sampling_filter_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"This is the filter function for dynamic sampling. "
|
||||
"It should be able to judge whether the result of a prompt should be selected or not."
|
||||
"We will do dynamic filter for sampling as in DAPO. e.g. not all correct or all wrong samples."
|
||||
"You could use `miles.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std` as an example."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_submission_granularity: A[
|
||||
str | None,
|
||||
Arg(
|
||||
choices=["group", "sample"],
|
||||
help=(
|
||||
"Granularity at which a completed unit frees rollout submission capacity. "
|
||||
"`group` holds a slot until the whole prompt group returns; `sample` frees each "
|
||||
"slot as its own sample finishes, so a replacement group goes out once "
|
||||
"n_samples_per_prompt samples complete, whichever groups they came from. "
|
||||
"Prompt groups are submitted whole either way. Unset picks the driver default: "
|
||||
"`sample` under --fully-async, where groups completed beyond the batch are queued "
|
||||
"for later steps; `group` otherwise, where they are aborted at the end of the step "
|
||||
"and, without --partial-rollout, discarded."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
# partial rollout
|
||||
partial_rollout: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to use partial rollout. "
|
||||
"If set, the unfinished samples during dynamic sampling will be recycled back to data buffer. "
|
||||
"This is useful for long responses."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
mask_offpolicy_in_partial_rollout: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to mask previous generation in partial rollout. "
|
||||
"If set, only on-policy generated tokens will be used in training"
|
||||
)
|
||||
),
|
||||
] = False
|
||||
max_weight_staleness: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Maximum allowed gap between a group's oldest generated-token weight version and the current "
|
||||
"engine weight version. Prompt KV weight staleness is not considered. "
|
||||
"Groups exceeding this threshold are recycled back to "
|
||||
"the data buffer instead of being sent to training. Only effective in fully "
|
||||
"async mode. None (default) disables staleness filtering."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
async_max_concurrent_samples: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Maximum number of concurrently generating trajectories in fully async mode, "
|
||||
"decoupling generation concurrency from the training batch size. None (default) "
|
||||
"keeps the legacy bound of one training batch worth of trajectories "
|
||||
"(rollout_batch_size groups, i.e. rollout_batch_size * n_samples_per_prompt)."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
async_data_buffer_capacity_factor: A[
|
||||
float,
|
||||
Arg(
|
||||
help=(
|
||||
"Capacity of the finished-group data buffer between rollout production and "
|
||||
"training consumption in fully async mode, as a multiple of rollout_batch_size "
|
||||
"(floor(factor * rollout_batch_size) groups). When the buffer is full the "
|
||||
"producer blocks until training consumes, so generation cannot run "
|
||||
"unboundedly ahead of training."
|
||||
)
|
||||
),
|
||||
] = 2.0
|
||||
async_unused_samples_handler: A[
|
||||
str,
|
||||
Arg(
|
||||
@@ -304,17 +130,6 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
),
|
||||
),
|
||||
] = "drop"
|
||||
custom_async_data_buffer_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to a custom DataBuffer subclass replacing the fully async finished-group "
|
||||
"data buffer (see miles/rollout/fully_async_data_buffer.py). Constructed with "
|
||||
"DataBufferConstructorInput; it takes over dataflow/staleness control, so the "
|
||||
"--async-data-buffer-* args apply only if the custom class reads them."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
custom_generate_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
@@ -324,37 +139,7 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
),
|
||||
),
|
||||
] = None
|
||||
custom_rollout_log_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The custom function for logging rollout data. The signature of the functions is: "
|
||||
"def log_rollout_data(rollout_id, args, samples, rollout_extra_metrics, rollout_time) -> bool. "
|
||||
"The return value indicates whether to skip the default logging. "
|
||||
),
|
||||
),
|
||||
] = None
|
||||
custom_eval_rollout_log_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The custom function for logging eval rollout data. "
|
||||
"def log_eval_rollout_data(rollout_id, args, data, extra_metrics) -> bool. "
|
||||
"The return value indicates whether to skip the default logging. "
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
buffer_filter_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to the buffer filter function. "
|
||||
"It should be able to select the samples in the buffer. "
|
||||
"The function should take list[list[Sample]] and return list[list[Sample]]."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
# update weight
|
||||
update_weight_buffer_size: A[
|
||||
int,
|
||||
@@ -523,3 +308,221 @@ class RolloutRelatedConfig(BaseConfig):
|
||||
float,
|
||||
Arg(help="Seconds the trainer controller waits for one trainer cell's update_weights before giving it up."),
|
||||
] = 600.0
|
||||
|
||||
|
||||
class RolloutRelatedRolloutOnlyConfig(BaseConfig):
|
||||
namespaced_radix_cache: A[
|
||||
bool | None,
|
||||
Arg(
|
||||
action=argparse.BooleanOptionalAction,
|
||||
help=(
|
||||
"Whether every generation request carries a radix cache key naming the rollout call "
|
||||
"the sample started under, so prefix KV computed under old weights cannot serve "
|
||||
"samples of a later call. Defaults to true when --fully-async is combined with "
|
||||
"--pause-generation-mode in_place, where the engine never flushes the cache and the "
|
||||
"staleness of a shared prompt is otherwise unbounded; an explicit "
|
||||
"--no-namespaced-radix-cache is respected."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_max_context_len: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The maximum context size for the inference engine during rollout."
|
||||
"It should no exceed the `max_position_embeddinds` in Huggingface model's `config.json`"
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_max_prompt_len: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The maximum length of the prompt for the inference engine during rollout. "
|
||||
"If set, we will filter out the long prompts during initialization of the global dataset. "
|
||||
"This is not recommended if the dataset is large."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_skip_special_tokens: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to skip special tokens in the response during rollout. "
|
||||
"This is useful when you want to use the response as a prompt for the next rollout."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
rollout_stop: A[
|
||||
list[str] | None,
|
||||
Arg(
|
||||
type_parser=str,
|
||||
nargs="+",
|
||||
help=(
|
||||
"The stop words for the inference engine during rollout. "
|
||||
"It can be a list of strings or a single string. "
|
||||
"It may be hard to pass special tokens in command line, in that case rollout_stop_token_ids can be used."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_stop_token_ids: A[
|
||||
list[int] | None,
|
||||
Arg(
|
||||
type_parser=int,
|
||||
nargs="+",
|
||||
help=(
|
||||
"The stop token ids for the inference engine during rollout. "
|
||||
"It can be a list of integers or a single integer."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
rollout_shuffle: A[
|
||||
bool,
|
||||
Arg(help="Whether to shuffle the prompts during rollout."),
|
||||
] = False
|
||||
rollout_seed: A[
|
||||
int,
|
||||
Arg(
|
||||
help=(
|
||||
"The seed for the random number generator during rollout. "
|
||||
"This is used to shuffle the prompts and also for the random sampling of the prompts."
|
||||
)
|
||||
),
|
||||
] = 42
|
||||
# sampling
|
||||
over_sampling_batch_size: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"This defines the granularity of the sampling batch in the rollout function. "
|
||||
"When the number of available samples falls below the target, a sampling "
|
||||
"operation of size over_sampling_batch_size will be triggered."
|
||||
"Regardless of whether partial rollout is used or filters are applied, "
|
||||
"the sampling granularity is always determined by this value. "
|
||||
"If this value is None, rollout_batch_size will be used as the default over_sampling_batch_size."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
dynamic_sampling_filter_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"This is the filter function for dynamic sampling. "
|
||||
"It should be able to judge whether the result of a prompt should be selected or not."
|
||||
"We will do dynamic filter for sampling as in DAPO. e.g. not all correct or all wrong samples."
|
||||
"You could use `miles.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std` as an example."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
rollout_submission_granularity: A[
|
||||
str | None,
|
||||
Arg(
|
||||
choices=["group", "sample"],
|
||||
help=(
|
||||
"Granularity at which a completed unit frees rollout submission capacity. "
|
||||
"`group` holds a slot until the whole prompt group returns; `sample` frees each "
|
||||
"slot as its own sample finishes, so a replacement group goes out once "
|
||||
"n_samples_per_prompt samples complete, whichever groups they came from. "
|
||||
"Prompt groups are submitted whole either way. Unset picks the driver default: "
|
||||
"`sample` under --fully-async, where groups completed beyond the batch are queued "
|
||||
"for later steps; `group` otherwise, where they are aborted at the end of the step "
|
||||
"and, without --partial-rollout, discarded."
|
||||
),
|
||||
),
|
||||
] = None
|
||||
|
||||
partial_rollout: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to use partial rollout. "
|
||||
"If set, the unfinished samples during dynamic sampling will be recycled back to data buffer. "
|
||||
"This is useful for long responses."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
mask_offpolicy_in_partial_rollout: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Whether to mask previous generation in partial rollout. "
|
||||
"If set, only on-policy generated tokens will be used in training"
|
||||
)
|
||||
),
|
||||
] = False
|
||||
max_weight_staleness: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Maximum allowed gap between a group's oldest generated-token weight version and the current "
|
||||
"engine weight version. Prompt KV weight staleness is not considered. "
|
||||
"Groups exceeding this threshold are recycled back to "
|
||||
"the data buffer instead of being sent to training. Only effective in fully "
|
||||
"async mode. None (default) disables staleness filtering."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
async_max_concurrent_samples: A[
|
||||
int | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Maximum number of concurrently generating trajectories in fully async mode, "
|
||||
"decoupling generation concurrency from the training batch size. None (default) "
|
||||
"keeps the legacy bound of one training batch worth of trajectories "
|
||||
"(rollout_batch_size groups, i.e. rollout_batch_size * n_samples_per_prompt)."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
async_data_buffer_capacity_factor: A[
|
||||
float,
|
||||
Arg(
|
||||
help=(
|
||||
"Capacity of the finished-group data buffer between rollout production and "
|
||||
"training consumption in fully async mode, as a multiple of rollout_batch_size "
|
||||
"(floor(factor * rollout_batch_size) groups). When the buffer is full the "
|
||||
"producer blocks until training consumes, so generation cannot run "
|
||||
"unboundedly ahead of training."
|
||||
)
|
||||
),
|
||||
] = 2.0
|
||||
custom_async_data_buffer_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to a custom DataBuffer subclass replacing the fully async finished-group "
|
||||
"data buffer (see miles/rollout/fully_async_data_buffer.py). Constructed with "
|
||||
"DataBufferConstructorInput; it takes over dataflow/staleness control, so the "
|
||||
"--async-data-buffer-* args apply only if the custom class reads them."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
custom_rollout_log_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The custom function for logging rollout data. The signature of the functions is: "
|
||||
"def log_rollout_data(rollout_id, args, samples, rollout_extra_metrics, rollout_time) -> bool. "
|
||||
"The return value indicates whether to skip the default logging. "
|
||||
),
|
||||
),
|
||||
] = None
|
||||
custom_eval_rollout_log_function_path: A[
|
||||
CustomFunctionConfig | None,
|
||||
Arg(
|
||||
help=(
|
||||
"The custom function for logging eval rollout data. "
|
||||
"def log_eval_rollout_data(rollout_id, args, data, extra_metrics) -> bool. "
|
||||
"The return value indicates whether to skip the default logging. "
|
||||
),
|
||||
),
|
||||
] = None
|
||||
buffer_filter_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
help=(
|
||||
"Path to the buffer filter function. "
|
||||
"It should be able to select the samples in the buffer. "
|
||||
"The function should take list[list[Sample]] and return list[list[Sample]]."
|
||||
)
|
||||
),
|
||||
] = None
|
||||
|
||||
@@ -9,6 +9,23 @@ class RolloutBufferConfig(BaseConfig):
|
||||
] = -1
|
||||
min_batch_collection_ratio: A[float, Arg(help="Minimum batch collection ratio")] = 1
|
||||
rollout_task_type: A[str, Arg()] = "math"
|
||||
data_pad_size_multiplier: A[
|
||||
int,
|
||||
Arg(help="Multiplier for data padding size in data processing."),
|
||||
] = 128
|
||||
disable_rollout_trim_samples: A[
|
||||
bool,
|
||||
Arg(help="disable trim samples in rollout buffer when converting samples to train data"),
|
||||
] = False
|
||||
use_dynamic_global_batch_size: A[
|
||||
bool,
|
||||
Arg(
|
||||
help="enable dynamic global batch size, disable trim samples in rollout buffer when converting samples to train data"
|
||||
),
|
||||
] = False
|
||||
|
||||
|
||||
class RolloutBufferRolloutOnlyConfig(BaseConfig):
|
||||
loss_mask_type: A[
|
||||
str,
|
||||
Arg(
|
||||
@@ -16,10 +33,6 @@ class RolloutBufferConfig(BaseConfig):
|
||||
help="Loss mask type",
|
||||
),
|
||||
] = "qwen"
|
||||
data_pad_size_multiplier: A[
|
||||
int,
|
||||
Arg(help="Multiplier for data padding size in data processing."),
|
||||
] = 128
|
||||
rollout_sample_filter_path: A[
|
||||
str | None,
|
||||
Arg(
|
||||
@@ -42,13 +55,3 @@ class RolloutBufferConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
disable_rollout_trim_samples: A[
|
||||
bool,
|
||||
Arg(help="disable trim samples in rollout buffer when converting samples to train data"),
|
||||
] = False
|
||||
use_dynamic_global_batch_size: A[
|
||||
bool,
|
||||
Arg(
|
||||
help="enable dynamic global batch size, disable trim samples in rollout buffer when converting samples to train data"
|
||||
),
|
||||
] = False
|
||||
|
||||
@@ -28,15 +28,6 @@ class TrainConfig(BaseConfig):
|
||||
Arg(choices=["torch", "flashinfer"], help="Top-k backend for Miles DSA indexer."),
|
||||
] = "torch"
|
||||
true_on_policy_mode: A[bool, Arg(help="Whether to enable true-on-policy mode.")] = False
|
||||
recompute_logprobs_via_prefill: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Recompute rollout logprobs via SGLang prefill instead of decode kernels. "
|
||||
"Only needed for models whose prefill and decode paths are not numerically identical."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
train_env_vars: A[
|
||||
Any,
|
||||
Arg(
|
||||
@@ -172,3 +163,15 @@ class TrainConfig(BaseConfig):
|
||||
)
|
||||
),
|
||||
] = None
|
||||
|
||||
|
||||
class TrainRolloutOnlyConfig(BaseConfig):
|
||||
recompute_logprobs_via_prefill: A[
|
||||
bool,
|
||||
Arg(
|
||||
help=(
|
||||
"Recompute rollout logprobs via SGLang prefill instead of decode kernels. "
|
||||
"Only needed for models whose prefill and decode paths are not numerically identical."
|
||||
)
|
||||
),
|
||||
] = False
|
||||
|
||||
@@ -49,10 +49,6 @@ class WandbConfig(BaseConfig):
|
||||
bool,
|
||||
Arg(help="Whether to log information for multi-turn rollout."),
|
||||
] = False
|
||||
log_passrate: A[
|
||||
bool,
|
||||
Arg(help="Whether to turn on passrate logging, which will log the pass@n of the responses in the rollout."),
|
||||
] = False
|
||||
log_reward_category: A[
|
||||
str | None,
|
||||
Arg(
|
||||
@@ -67,3 +63,10 @@ class WandbConfig(BaseConfig):
|
||||
Arg(help="Explicitly log metrics for correct samples."),
|
||||
] = False
|
||||
wandb_run_id: A[str | None, Arg()] = None
|
||||
|
||||
|
||||
class WandbRolloutOnlyConfig(BaseConfig):
|
||||
log_passrate: A[
|
||||
bool,
|
||||
Arg(help="Whether to turn on passrate logging, which will log the pass@n of the responses in the rollout."),
|
||||
] = False
|
||||
|
||||
+33
-10
@@ -7,34 +7,37 @@ from miles.utils.args.component_orchestrator import OrchestratorOnlyConfig
|
||||
from miles.utils.args.component_rollout import InferenceControllerOnlyConfig, RolloutOnlyConfig
|
||||
from miles.utils.args.component_shared import SglangFieldsConfig
|
||||
from miles.utils.args.component_trainer import TrainerOnlyConfig
|
||||
from miles.utils.args.configs.algo import AlgoConfig
|
||||
from miles.utils.args.configs.algo import AlgoConfig, AlgoRolloutOnlyConfig
|
||||
from miles.utils.args.configs.backend_fields import RawTrainerBackendConfig, TrainerBackendTraitConfig
|
||||
from miles.utils.args.configs.ci import CiConfig
|
||||
from miles.utils.args.configs.ci import CiConfig, CiRolloutOnlyConfig
|
||||
from miles.utils.args.configs.cluster import ClusterConfig
|
||||
from miles.utils.args.configs.custom_megatron_plugins import CustomMegatronPluginsConfig, Dsv4MegatronPluginsConfig
|
||||
from miles.utils.args.configs.dashboard import DashboardConfig
|
||||
from miles.utils.args.configs.data import DataConfig
|
||||
from miles.utils.args.configs.debug import DebugConfig
|
||||
from miles.utils.args.configs.eval import EvalConfig
|
||||
from miles.utils.args.configs.debug import DebugConfig, DebugRolloutOnlyConfig
|
||||
from miles.utils.args.configs.eval import EvalConfig, EvalRolloutOnlyConfig
|
||||
from miles.utils.args.configs.fault_tolerance import FaultToleranceConfig
|
||||
from miles.utils.args.configs.lora import LoraConfig
|
||||
from miles.utils.args.configs.mlflow import MlflowConfig
|
||||
from miles.utils.args.configs.mtp_training import MtpTrainingConfig
|
||||
from miles.utils.args.configs.network import NetworkConfig
|
||||
from miles.utils.args.configs.on_policy_distillation import OnPolicyDistillationConfig
|
||||
from miles.utils.args.configs.on_policy_distillation import (
|
||||
OnPolicyDistillationConfig,
|
||||
OnPolicyDistillationRolloutOnlyConfig,
|
||||
)
|
||||
from miles.utils.args.configs.prefill_decode_disaggregation import PrefillDecodeDisaggregationConfig
|
||||
from miles.utils.args.configs.prometheus import PrometheusConfig
|
||||
from miles.utils.args.configs.reward_model import RewardModelConfig
|
||||
from miles.utils.args.configs.rollout import RolloutRelatedConfig
|
||||
from miles.utils.args.configs.rollout_buffer import RolloutBufferConfig
|
||||
from miles.utils.args.configs.reward_model import RewardModelConfig, RewardModelRolloutOnlyConfig
|
||||
from miles.utils.args.configs.rollout import RolloutRelatedConfig, RolloutRelatedRolloutOnlyConfig
|
||||
from miles.utils.args.configs.rollout_buffer import RolloutBufferConfig, RolloutBufferRolloutOnlyConfig
|
||||
from miles.utils.args.configs.router import RouterConfig
|
||||
from miles.utils.args.configs.run_uuid import RunUuidConfig
|
||||
from miles.utils.args.configs.scaling import ScalingConfig
|
||||
from miles.utils.args.configs.session import SessionConfig
|
||||
from miles.utils.args.configs.tensorboard import TensorboardConfig
|
||||
from miles.utils.args.configs.tinker import TinkerConfig
|
||||
from miles.utils.args.configs.train import TrainConfig
|
||||
from miles.utils.args.configs.wandb import WandbConfig
|
||||
from miles.utils.args.configs.train import TrainConfig, TrainRolloutOnlyConfig
|
||||
from miles.utils.args.configs.wandb import WandbConfig, WandbRolloutOnlyConfig
|
||||
from miles.utils.args.runtime_base import BaseLeafConfig
|
||||
|
||||
|
||||
@@ -181,6 +184,16 @@ class RolloutConfig(
|
||||
RawTrainerBackendConfig,
|
||||
TrainerBackendTraitConfig,
|
||||
RolloutOnlyConfig,
|
||||
AlgoRolloutOnlyConfig,
|
||||
CiRolloutOnlyConfig,
|
||||
DebugRolloutOnlyConfig,
|
||||
EvalRolloutOnlyConfig,
|
||||
OnPolicyDistillationRolloutOnlyConfig,
|
||||
RewardModelRolloutOnlyConfig,
|
||||
RolloutRelatedRolloutOnlyConfig,
|
||||
RolloutBufferRolloutOnlyConfig,
|
||||
TrainRolloutOnlyConfig,
|
||||
WandbRolloutOnlyConfig,
|
||||
RunUuidConfig,
|
||||
ClusterConfig,
|
||||
TrainConfig,
|
||||
@@ -281,6 +294,16 @@ class AllConfig(
|
||||
SglangFieldsConfig,
|
||||
OrchestratorOnlyConfig,
|
||||
RolloutOnlyConfig,
|
||||
AlgoRolloutOnlyConfig,
|
||||
CiRolloutOnlyConfig,
|
||||
DebugRolloutOnlyConfig,
|
||||
EvalRolloutOnlyConfig,
|
||||
OnPolicyDistillationRolloutOnlyConfig,
|
||||
RewardModelRolloutOnlyConfig,
|
||||
RolloutRelatedRolloutOnlyConfig,
|
||||
RolloutBufferRolloutOnlyConfig,
|
||||
TrainRolloutOnlyConfig,
|
||||
WandbRolloutOnlyConfig,
|
||||
InferenceControllerOnlyConfig,
|
||||
MultiLoraOnlyConfig,
|
||||
TinkerConfig,
|
||||
|
||||
+23
-10
@@ -24,34 +24,37 @@ from miles.backends.sglang_utils.sglang_scaling_config import SglangScalingConfi
|
||||
from miles.dashboard.args import validate_dashboard_args
|
||||
from miles.ray.specs.train import external_trainer_controller_addrs
|
||||
from miles.rollout.checkpoint_eval import is_checkpoint_eval_fn
|
||||
from miles.utils.args.configs.algo import AlgoConfig
|
||||
from miles.utils.args.configs.algo import AlgoConfig, AlgoRolloutOnlyConfig
|
||||
from miles.utils.args.configs.backend_fields import TrainerBackendTraitConfig
|
||||
from miles.utils.args.configs.ci import CiConfig
|
||||
from miles.utils.args.configs.ci import CiConfig, CiRolloutOnlyConfig
|
||||
from miles.utils.args.configs.cluster import ClusterConfig
|
||||
from miles.utils.args.configs.custom_megatron_plugins import CustomMegatronPluginsConfig, Dsv4MegatronPluginsConfig
|
||||
from miles.utils.args.configs.dashboard import DashboardConfig
|
||||
from miles.utils.args.configs.data import DataConfig
|
||||
from miles.utils.args.configs.debug import DebugConfig
|
||||
from miles.utils.args.configs.eval import EvalConfig
|
||||
from miles.utils.args.configs.debug import DebugConfig, DebugRolloutOnlyConfig
|
||||
from miles.utils.args.configs.eval import EvalConfig, EvalRolloutOnlyConfig
|
||||
from miles.utils.args.configs.fault_tolerance import _DEFAULT_FT_API_SERVER_PORT, FaultToleranceConfig
|
||||
from miles.utils.args.configs.lora import LoraConfig
|
||||
from miles.utils.args.configs.mlflow import MlflowConfig
|
||||
from miles.utils.args.configs.mtp_training import MtpTrainingConfig
|
||||
from miles.utils.args.configs.network import NetworkConfig
|
||||
from miles.utils.args.configs.on_policy_distillation import OnPolicyDistillationConfig
|
||||
from miles.utils.args.configs.on_policy_distillation import (
|
||||
OnPolicyDistillationConfig,
|
||||
OnPolicyDistillationRolloutOnlyConfig,
|
||||
)
|
||||
from miles.utils.args.configs.prefill_decode_disaggregation import PrefillDecodeDisaggregationConfig
|
||||
from miles.utils.args.configs.prometheus import PrometheusConfig
|
||||
from miles.utils.args.configs.reward_model import RewardModelConfig
|
||||
from miles.utils.args.configs.rollout import RolloutRelatedConfig
|
||||
from miles.utils.args.configs.rollout_buffer import RolloutBufferConfig
|
||||
from miles.utils.args.configs.reward_model import RewardModelConfig, RewardModelRolloutOnlyConfig
|
||||
from miles.utils.args.configs.rollout import RolloutRelatedConfig, RolloutRelatedRolloutOnlyConfig
|
||||
from miles.utils.args.configs.rollout_buffer import RolloutBufferConfig, RolloutBufferRolloutOnlyConfig
|
||||
from miles.utils.args.configs.router import RouterConfig
|
||||
from miles.utils.args.configs.run_uuid import RunUuidConfig
|
||||
from miles.utils.args.configs.scaling import ScalingConfig
|
||||
from miles.utils.args.configs.session import SessionConfig
|
||||
from miles.utils.args.configs.tensorboard import TensorboardConfig
|
||||
from miles.utils.args.configs.tinker import TinkerConfig
|
||||
from miles.utils.args.configs.train import TrainConfig
|
||||
from miles.utils.args.configs.wandb import WandbConfig
|
||||
from miles.utils.args.configs.train import TrainConfig, TrainRolloutOnlyConfig
|
||||
from miles.utils.args.configs.wandb import WandbConfig, WandbRolloutOnlyConfig
|
||||
from miles.utils.args.custom_function import add_user_provided_function_arguments, resolve_custom_function_configs
|
||||
from miles.utils.args.runtime import AllConfig
|
||||
from miles.utils.audit_utils.event_logger.logger import EVENTS_DIRNAME
|
||||
@@ -213,11 +216,15 @@ def get_miles_extra_args_provider(
|
||||
ClusterConfig.add_arguments(parser=parser)
|
||||
ScalingConfig.add_arguments(parser=parser)
|
||||
TrainConfig.add_arguments(parser=parser)
|
||||
TrainRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
RolloutRelatedConfig.add_arguments(parser=parser)
|
||||
RolloutRelatedRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
FaultToleranceConfig.add_arguments(parser=parser)
|
||||
DataConfig.add_arguments(parser=parser)
|
||||
EvalConfig.add_arguments(parser=parser)
|
||||
EvalRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
AlgoConfig.add_arguments(parser=parser)
|
||||
AlgoRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
TrainerBackendTraitConfig.add_arguments(parser=parser)
|
||||
reset_arg(parser=parser, name="--lr", type=float, default=1e-6)
|
||||
reset_arg(parser=parser, name="--clip-grad", type=float, default=1.0)
|
||||
@@ -233,14 +240,17 @@ def get_miles_extra_args_provider(
|
||||
),
|
||||
)
|
||||
OnPolicyDistillationConfig.add_arguments(parser=parser)
|
||||
OnPolicyDistillationRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
LoraConfig.add_arguments(parser=parser)
|
||||
WandbConfig.add_arguments(parser=parser)
|
||||
WandbRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
MlflowConfig.add_arguments(parser=parser)
|
||||
TensorboardConfig.add_arguments(parser=parser)
|
||||
PrometheusConfig.add_arguments(parser=parser)
|
||||
DashboardConfig.add_arguments(parser=parser.add_argument_group("miles dashboard"))
|
||||
RouterConfig.add_arguments(parser=parser)
|
||||
DebugConfig.add_arguments(parser=parser)
|
||||
DebugRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
SglangConfig.add_arguments(parser)
|
||||
# required whenever expert projections are LoRA targets, inert otherwise
|
||||
# (sglang's own default is False)
|
||||
@@ -248,12 +258,15 @@ def get_miles_extra_args_provider(
|
||||
SessionConfig.add_arguments(parser=parser)
|
||||
NetworkConfig.add_arguments(parser=parser)
|
||||
RewardModelConfig.add_arguments(parser=parser)
|
||||
RewardModelRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
RolloutBufferConfig.add_arguments(parser=parser)
|
||||
RolloutBufferRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
MtpTrainingConfig.add_arguments(parser=parser)
|
||||
reset_arg(parser=parser, name="--mtp-num-layers", type=int, default=None)
|
||||
reset_arg(parser=parser, name="--mtp-loss-scaling-factor", type=float, default=0.2)
|
||||
PrefillDecodeDisaggregationConfig.add_arguments(parser=parser)
|
||||
CiConfig.add_arguments(parser=parser)
|
||||
CiRolloutOnlyConfig.add_arguments(parser=parser)
|
||||
CustomMegatronPluginsConfig.add_arguments(parser=parser)
|
||||
Dsv4MegatronPluginsConfig.add_arguments(parser=parser)
|
||||
TinkerConfig.add_arguments(parser=parser)
|
||||
|
||||
Reference in New Issue
Block a user