Sweeping hyperparameters with `TrainingGroup`
Sweep hyperparameters across runs with TrainingGroup
Tuning an RL run usually means launching the same recipe several times
with one or two knobs changed — learning rate, rollout temperature, KL
coefficient — and comparing the curves. Doing that by hand is tedious and
error-prone: you copy a TrainConfig, tweak a field, and hope you didn’t
fat-finger a name or leave two runs sharing a checkpoint.
TrainingGroup makes this a first-class operation. You give it one base
TrainConfig and a grid of field overrides, and it expands them
into the cross-product of concrete runs — one independent training job per
combination, each with its own training_run_id but a shared group_id so
the dashboard can show them side by side.
Three things make it safe to launch a big sweep:
get_train_configs()returns the derivedTrainConfigs so you can see exactly what will run before spending a single GPU-second.- Invalid overrides — a misspelled field, or a value of the wrong type — raise immediately, before any run starts. A sweep never dies three variants deep because of a typo.
launch()starts every variant as a detached Modal run withTrainingRunhandles you can inspect or wait on later.
from typing import Any
from modal_training_gym import ( HuggingFaceDataset, Qwen3_4B, SlimeRecipe, TrainConfig, TrainingGroup,)1. Define the base run
Section titled “1. Define the base run”Start with an ordinary TrainConfig — the recipe you’d launch if you
weren’t sweeping. Everything the sweep doesn’t override is inherited from
here, so set the fields you want held constant once. We use a short Qwen3-4B
DAPO-math run as the base.
class MathDataset(HuggingFaceDataset): hf_repo = "zhuzilin/dapo-math-17k" input_key = "prompt" label_key = "label" output_format = "jsonl" apply_chat_template = True always_prepare = True
def load(self, split: str = "all") -> Any: from datasets import load_dataset
ds = load_dataset(self.hf_repo, self.hf_config, split=self.hf_split) stop = len(ds) if not self.n_rows else min(self.n_rows, len(ds)) return ds.select(range(stop))
train_dataset = MathDataset(n_rows=2_000)base = TrainConfig( model=Qwen3_4B(), dataset=train_dataset, recipe=SlimeRecipe( rm_type="dapo", gpu_type="H100", colocate=True, actor_num_nodes=1, actor_num_gpus_per_node=8, tensor_model_parallel_size=2, sequence_parallel=True, rollout_num_gpus_per_engine=1, num_rollout=15, rollout_batch_size=16, n_samples_per_prompt=8, rollout_max_response_len=8192, rollout_temperature=1.0, global_batch_size=32, lr=1e-6, advantage_estimator="grpo", use_kl_loss=False, kl_coef=0.0, use_dynamic_batch_size=True, max_tokens_per_gpu=9216, sglang_mem_fraction_static=0.75, save_interval=10, ),)2. Wrap it in a TrainingGroup
Section titled “2. Wrap it in a TrainingGroup”The grid maps dotted field paths to the list of values to try. Paths
address the composed config: recipe.lr sets lr on the recipe,
model.model_name would set it on the model, and so on. The group is the
cross-product of every axis — here 2 × 2 = 4 runs.
The grid is validated the moment you construct the group: misspell
recipe.lr as recipe.learning_rate and you get an immediate error with a
“did you mean” hint, not a four-hour run that crashes at the end.
Pass a name to give the group a stable, human-readable id — it’s slugified
and stamped onto every run, so you can find and filter the whole sweep by it
in the dashboard’s Group filter. Omit it and one is autogenerated.
group = TrainingGroup( base=base, grid={ "recipe.lr": [1e-6, 5e-6], "recipe.rollout_temperature": [0.8, 1.0], }, name="qwen4b-lr-temp-sweep",)3. Preview the derived runs (no GPU)
Section titled “3. Preview the derived runs (no GPU)”get_train_configs() expands and fully validates the grid, returning the
concrete TrainConfigs. Inspect them before launching anything — confirm
the values are what you expect and that each run is distinct.
configs = group.get_train_configs()print(f"{len(configs)} runs in group {group.group_id}:")for cfg in configs: print( f" lr={cfg.recipe.lr:<8} " f"temp={cfg.recipe.rollout_temperature}" )4. Launch the sweep
Section titled “4. Launch the sweep”Blocking param sweep
Section titled “Blocking param sweep”train() runs every variant and returns the successful TrainResults.
With max_parallel > 1 the variants run concurrently — each is a detached,
independent Modal app — and a single failure is recorded in
group.failures instead of sinking the rest of the sweep.
Every result carries the shared group_id, so you can pull the whole sweep
back together afterwards (and the dashboard groups them automatically).
Background param sweep
Section titled “Background param sweep”launch() starts every variant as a detached Modal run and returns a list of
TrainingRun handles. Each handle has the training_run_id, Modal app URL,
function-call id, and shared group_id.
Pass prepare_inputs=True to run the model/download conversion steps before
spawning training, matching the one-shot TrainConfig.train() behavior. A
single launch failure is recorded in group.failures instead of sinking the
rest of the sweep.
Call launch.result() to wait for a handle’s trained TrainResult. If you
don’t need the handles, group.train(max_parallel=...) wraps this pattern
and returns the successful TrainResults directly.
launches = group.launch(prepare_inputs=True)print(f"group {group.group_id}: {len(launches)} runs launched")for launch in launches: print( f" {launch.training_run_id} " f"app={launch.modal_app_id} " f"group_id={launch.group_id}" )if group.failures: for overrides, err in group.failures: print(f" FAILED {overrides}: {err}")
results = []for launch in launches: result = launch.result() results.append(result) print(f"completed {result.training_run_id} (group_id={result.group_id})")
print(f"group {group.group_id}: {len(results)} runs completed")- One base
TrainConfig+ agridof dotted overrides → the cross-product of independent runs, each with a uniquetraining_run_idand a sharedgroup_id(set fromname, or autogenerated). get_train_configs()shows you the derived runs, and validation rejects bad fields/values before anything launches.train(max_parallel=...)fans the sweep out across Modal;group.failuresisolates any run that didn’t make it.launch(prepare_inputs=True)fans the sweep out across Modal and returnsTrainingRunhandles;group.failuresisolates any run that didn’t launch.- Use
launch.result()to wait on a specific launched run, ortrain(max_parallel=...)when you want one blocking call that returns the successfulTrainResults. - Every run is tagged with the
group_id, so the dashboard’s Group filter pulls the whole sweep together for side-by-side comparison.
Sweep any composed field, not just the recipe — model.* and dataset.*
paths work too, so you can vary the dataset size or a model setting the same
way. Keep the grid small to start: cost scales with the product of the axes.
Related API Reference
Section titled “Related API Reference”Source: tutorials/rl/007_param_sweep/007_param_sweep.py
| Open in Modal Notebook