diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 269339e056..bad74eeba0 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -913,6 +913,7 @@ skip_first_n_steps_for_profiler: 1 # Profile for a small number of steps to avoid a large profile file size. profiler_steps: 5 hide_profiler_step_metric: false +enable_continuous_profiling: false profile_cleanly: true # If set to true, adds a block_until_ready on train state which aligns the profile for each step. profile_periodically_period: -1 # If set to a positive integer, profile every profile_periodically_period steps. # This is useful to debug scenarios where performance is changing. diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 67eb7c7fd6..afa9d802c5 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -2203,6 +2203,7 @@ class Profiling(BaseModel): profile_cleanly: bool = Field(True, description="Add block_until_ready to align profile for each step.") profile_periodically_period: int = Field(-1, description="If positive, profile every N steps.") hide_profiler_step_metric: bool = Field(False, description="Whether to enable profiler step metric.") + enable_continuous_profiling: bool = Field(False, description="Enable continuous profiling in tunix profiler. Once enabled, it will support saving profile > 2GB.") enable_jax_profiler: bool = Field(False, description="Enable the JAX live profiler.") jax_profiler_port: int = Field(9999, description="Port for the JAX profiler.") enable_tpu_profiling_options: bool = Field(False, description="Enable TPU advanced profiling options.") diff --git a/src/maxtext/trainers/post_train/rl/train_rl.py b/src/maxtext/trainers/post_train/rl/train_rl.py index 3296971007..5348e4741b 100644 --- a/src/maxtext/trainers/post_train/rl/train_rl.py +++ b/src/maxtext/trainers/post_train/rl/train_rl.py @@ -474,7 +474,9 @@ def create_rl_components( # pylint: disable=too-many-positional-arguments log_dir=trainer_config.tensorboard_dir, skip_first_n_steps=trainer_config.skip_first_n_steps_for_profiler, profiler_steps=trainer_config.profiler_steps, + # Skip setting tracer levels. set_profile_options=False, + enable_continuous_profiling=trainer_config.enable_continuous_profiling, ) # Parse vllm_additional_config