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:
Tom
2026-10-01 15:43:33 +08:00
parent 558b660361
commit 109977da42
12 changed files with 414 additions and 348 deletions
+19 -16
View File
@@ -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
+10 -7
View File
@@ -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
+35 -32
View File
@@ -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
+18 -15
View File
@@ -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
+9 -6
View File
@@ -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
View File
@@ -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
+17 -14
View File
@@ -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
+12 -9
View File
@@ -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
+7 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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)