Skip to content

Commit ec552cf

Browse files
authored
Fix maxtext_rl and maxtext_sft DAGs (GoogleCloudPlatform#1211)
1 parent 6e52f5c commit ec552cf

3 files changed

Lines changed: 34 additions & 30 deletions

File tree

dags/post_training/maxtext_rl.py

Lines changed: 20 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -64,7 +64,7 @@
6464
4. All training steps execute within expected parameters
6565
5. Metrics are successfully synced to Vertex AI TensorBoard
6666
""",
67-
concurrency=2,
67+
concurrency=1,
6868
) as dag:
6969
training_config = test_config_util.RLTestConfig(
7070
cluster=XpkClusters.TPU_V5P_128_CLUSTER,
@@ -88,41 +88,41 @@
8888
# HF token retrieved from Airflow Variables for secure credential management
8989
HF_TOKEN_LLAMA3_1 = models.Variable.get("HF_TOKEN_CIENET", None)
9090

91+
tests = []
9192
for mode, image in test_config_util.POST_TRAINING_DOCKER_IMAGES:
9293
# TODO: Enable stable mode once a new version of MaxText is available
9394
if mode == test_config_util.SetupMode.STABLE:
9495
continue # Skip stable for RL training tests
9596

9697
for loss_algo in training_config.loss_algos:
9798
for slice_num in training_config.slices:
98-
run_name = validation_util.generate_run_name(
99-
prefix=loss_algo.value,
100-
mode=mode.value,
101-
num_slices=slice_num,
102-
)
103-
104-
rl_training_command = training_config.generate_rl_training_command(
105-
loss_algo=loss_algo,
106-
run_name=run_name,
107-
hf_token=HF_TOKEN_LLAMA3_1,
108-
num_slices=slice_num,
109-
)
110-
11199
with TaskGroup(
112100
group_id=(
113101
f"{loss_algo.value}-{mode.value}-"
114102
f"{slice_num}x{training_config.accelerator}"
115103
)
116104
) as group:
117105
with TaskGroup(group_id="run_training") as training_group:
106+
run_name = validation_util.generate_run_name(
107+
prefix=loss_algo.value,
108+
mode=mode.value,
109+
num_slices=slice_num,
110+
)
111+
112+
rl_training_command = training_config.generate_rl_training_command(
113+
loss_algo=loss_algo,
114+
run_name=run_name,
115+
hf_token=HF_TOKEN_LLAMA3_1,
116+
num_slices=slice_num,
117+
)
118118
start_time = validation_util.generate_timestamp.override(
119119
task_id="generate_start_time"
120120
)()
121121

122122
training_task = gke_config.get_gke_config(
123+
time_out_in_min=30,
123124
num_slices=slice_num,
124125
cluster=training_config.cluster,
125-
time_out_in_min=30,
126126
test_name=loss_algo.value,
127127
run_model_cmds=rl_training_command,
128128
docker_image=image.value,
@@ -137,6 +137,7 @@
137137
)()
138138

139139
chain(
140+
run_name,
140141
start_time,
141142
training_task,
142143
end_time,
@@ -149,7 +150,7 @@
149150
project_id=training_config.cluster.project,
150151
location=zone_to_region(training_config.cluster.zone),
151152
cluster_name=training_config.cluster.name,
152-
text_filter=f'"Config param loss_algo: {loss_algo.loss_name}"',
153+
text_filter=f"\"'loss_algo': '{loss_algo.loss_name}'\"",
153154
namespace="default",
154155
container_name="jax-tpu",
155156
pod_pattern=f"{loss_algo.value}.*",
@@ -183,3 +184,6 @@
183184
)
184185

185186
chain(training_group, [validation_group, upload_to_vertex_ai])
187+
tests.append(group)
188+
189+
chain(*tests)

dags/post_training/maxtext_sft.py

Lines changed: 13 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -75,7 +75,7 @@ def validate_training(
7575
project_id=config.cluster.project,
7676
location=zone_to_region(config.cluster.zone),
7777
cluster_name=config.cluster.name,
78-
text_filter=f"\"'step': {steps}\"",
78+
text_filter=f'"Training: 100%" AND "{steps}/{steps}"',
7979
namespace="default",
8080
container_name="jax-tpu",
8181
pod_pattern=f"{config.short_id}.*",
@@ -122,7 +122,7 @@ def validate_training(
122122
3. No infrastructure failures or container launch issues occur
123123
4. All training steps execute within expected parameters
124124
""",
125-
concurrency=2,
125+
concurrency=1,
126126
) as dag:
127127
training_steps = 30
128128

@@ -152,19 +152,19 @@ def validate_training(
152152
continue # Skip stable for SFT training tests
153153

154154
for slice_num in training_config.slices:
155-
run_name = validation_util.generate_run_name(
156-
prefix="sft",
157-
mode=mode.value,
158-
num_slices=slice_num,
159-
)
160-
161-
sft_training_command = training_config.generate_sft_training_command(
162-
run_name=run_name,
163-
hf_token=HF_TOKEN_CIENET,
164-
)
165155
with TaskGroup(
166156
group_id=f"sft-{mode.value}-{slice_num}x{training_config.accelerator}"
167157
) as group:
158+
run_name = validation_util.generate_run_name(
159+
prefix="sft",
160+
mode=mode.value,
161+
num_slices=slice_num,
162+
)
163+
164+
sft_training_command = training_config.generate_sft_training_command(
165+
run_name=run_name,
166+
hf_token=HF_TOKEN_CIENET,
167+
)
168168
training_group, start_time, end_time = run_training(
169169
slice_num, training_config, sft_training_command, image
170170
)
@@ -181,4 +181,4 @@ def validate_training(
181181
run_name_prefix=run_name,
182182
)
183183

184-
chain(training_group, [validation_group, upload_to_vertex_ai])
184+
chain(run_name, training_group, [validation_group, upload_to_vertex_ai])

dags/post_training/util/test_config_util.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -158,7 +158,7 @@ def generate_rl_training_command(
158158
])
159159

160160
rl_command = (
161-
"python -m src.maxtext.rl.train_rl "
161+
"python -m src.MaxText.rl.train_rl "
162162
f"{self.rl_config_path} " + " ".join(command_params)
163163
)
164164

0 commit comments

Comments
 (0)