Skip to content

Commit 8f2e872

Browse files
Refactor maxtext_rl DAG and organize tasks in a group (GoogleCloudPlatform#1125)
1 parent fae1429 commit 8f2e872

2 files changed

Lines changed: 79 additions & 79 deletions

File tree

dags/post_training/maxtext_rl.py

Lines changed: 69 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111

1212
from airflow import models
1313
from airflow.models.baseoperator import chain
14+
from airflow.utils.task_group import TaskGroup
1415

1516
from dags import composer_env
1617
from dags.common import test_owner
@@ -69,7 +70,6 @@
6970
accelerator="v5p-128",
7071
slices=[1], # Single slice for RL training
7172
model_name="llama3.1-70b",
72-
short_id="mrl",
7373
base_dir=(
7474
f"{test_config_util.DEFAULT_BUCKET}/llama3.1-70b-Instruct/outputs"
7575
),
@@ -94,79 +94,78 @@
9494

9595
for loss_algo in training_config.loss_algos:
9696
for slice_num in training_config.slices:
97-
task_suffix = f"{mode.value}_{loss_algo.value}_{slice_num}"
98-
run_name = validation_util.generate_posttraining_run_name.override(
99-
task_id=f"run_name_{task_suffix}"
100-
)(
101-
short_id=training_config.short_id,
102-
checkpointing_type=loss_algo.value,
103-
slice_number=slice_num,
104-
mode=mode.value,
105-
)
97+
current_datetime = datetime.datetime.now().strftime("%Y-%m-%d-%H-%M-%S")
98+
run_name = f"{loss_algo.value}-{mode.value}-{current_datetime}"
10699

107100
rl_training_command = training_config.generate_rl_training_command(
108101
loss_algo=loss_algo,
109102
run_name=run_name,
110103
hf_token=HF_TOKEN_LLAMA3_1,
111104
)
112105

113-
start_time = validation_util.generate_timestamp.override(
114-
task_id=f"start_time_{task_suffix}"
115-
)()
116-
117-
test_name = f"{training_config.short_id[:3]}{loss_algo.value[:3]}"
118-
119-
training_task = gke_config.get_gke_config(
120-
num_slices=slice_num,
121-
cluster=training_config.cluster,
122-
time_out_in_min=30,
123-
test_name=test_name,
124-
run_model_cmds=rl_training_command,
125-
docker_image=image.value,
126-
test_owner=test_owner.JACKY_F,
127-
).run(
128-
use_pathways=True,
129-
xpk_branch=MAIN_BRANCH,
130-
skip_post_process=True,
131-
)
132-
133-
end_time = validation_util.generate_timestamp.override(
134-
task_id=f"end_time_{task_suffix}"
135-
)()
136-
137-
validate_loss_algo = validation_util.validate_log_exist.override(
138-
task_id=f"validate_loss_{task_suffix}"
139-
)(
140-
project_id=training_config.cluster.project,
141-
location=zone_to_region(training_config.cluster.zone),
142-
cluster_name=training_config.cluster.name,
143-
text_filter=f'"Config param loss_algo: {loss_algo.value}"',
144-
namespace="default",
145-
container_name="jax-tpu",
146-
pod_pattern=f"{test_name}.*",
147-
start_time=start_time,
148-
end_time=end_time,
149-
)
150-
151-
validate_training = validation_util.validate_log_exist.override(
152-
task_id=f"validate_training_{task_suffix}"
153-
)(
154-
project_id=training_config.cluster.project,
155-
location=zone_to_region(training_config.cluster.zone),
156-
cluster_name=training_config.cluster.name,
157-
text_filter='"Post RL Training"',
158-
namespace="default",
159-
container_name="jax-tpu",
160-
pod_pattern=f"{test_name}.*",
161-
start_time=start_time,
162-
end_time=end_time,
163-
)
164-
165-
chain(
166-
run_name,
167-
start_time,
168-
training_task,
169-
end_time,
170-
validate_loss_algo,
171-
validate_training,
172-
)
106+
with TaskGroup(
107+
group_id=f"{loss_algo.value}-{mode.value}-{slice_num}x{training_config.accelerator}"
108+
) as group:
109+
with TaskGroup(group_id="run_training") as training_group:
110+
start_time = validation_util.generate_timestamp.override(
111+
task_id="generate_start_time"
112+
)()
113+
114+
training_task = gke_config.get_gke_config(
115+
num_slices=slice_num,
116+
cluster=training_config.cluster,
117+
time_out_in_min=30,
118+
test_name=loss_algo.value,
119+
run_model_cmds=rl_training_command,
120+
docker_image=image.value,
121+
test_owner=test_owner.JACKY_F,
122+
).run_model(
123+
use_pathways=True,
124+
xpk_branch=MAIN_BRANCH,
125+
)
126+
127+
end_time = validation_util.generate_timestamp.override(
128+
task_id="generate_end_time"
129+
)()
130+
131+
chain(
132+
start_time,
133+
training_task,
134+
end_time,
135+
)
136+
137+
with TaskGroup(group_id="validate_training") as validation_group:
138+
validate_loss_algo = validation_util.validate_log_exist.override(
139+
task_id="validate_loss_algo"
140+
)(
141+
project_id=training_config.cluster.project,
142+
location=zone_to_region(training_config.cluster.zone),
143+
cluster_name=training_config.cluster.name,
144+
text_filter=f'"Config param loss_algo: {loss_algo.loss_name}"',
145+
namespace="default",
146+
container_name="jax-tpu",
147+
pod_pattern=f"{loss_algo.value}.*",
148+
start_time=start_time,
149+
end_time=end_time,
150+
)
151+
152+
validate_training_logs = (
153+
validation_util.validate_log_exist.override(
154+
task_id="validate_training_logs"
155+
)(
156+
project_id=training_config.cluster.project,
157+
location=zone_to_region(training_config.cluster.zone),
158+
cluster_name=training_config.cluster.name,
159+
text_filter="Post RL Training",
160+
namespace="default",
161+
container_name="jax-tpu",
162+
pod_pattern=f"{loss_algo.value}.*",
163+
start_time=start_time,
164+
end_time=end_time,
165+
)
166+
)
167+
168+
chain(
169+
training_group,
170+
validation_group,
171+
)

dags/post_training/util/test_config_util.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,14 @@ class LossAlgo(Enum):
2323
GRPO = "grpo"
2424
GSPO = "gspo"
2525

26+
@property
27+
def loss_name(self) -> str:
28+
"""Returns the specific loss algorithm string used for computation."""
29+
return {
30+
LossAlgo.GRPO: "grpo",
31+
LossAlgo.GSPO: "gspo-token",
32+
}[self]
33+
2634

2735
@dataclass
2836
class RLTestConfig:
@@ -38,7 +46,6 @@ class RLTestConfig:
3846
slices: List of slice numbers to test with.
3947
model_name: The name of the model being trained
4048
(e.g., llama3.1-70b).
41-
short_id: A short identifier for the test run.
4249
base_dir: Base GCS directory for outputs.
4350
tokenizer_path: Path to the tokenizer (HuggingFace model
4451
path or local path).
@@ -53,7 +60,6 @@ class RLTestConfig:
5360
accelerator: str
5461
slices: list[int]
5562
model_name: str
56-
short_id: str
5763
base_dir: str
5864
tokenizer_path: str
5965
load_parameters_path: str
@@ -66,7 +72,6 @@ def __init__(
6672
accelerator: str,
6773
slices: list[int],
6874
model_name: str,
69-
short_id: str,
7075
base_dir: str,
7176
tokenizer_path: str,
7277
load_parameters_path: str,
@@ -82,7 +87,6 @@ def __init__(
8287
slices: The number of slices to be used.
8388
model_name: The name of the base model being tested
8489
(e.g., llama3.1-70b).
85-
short_id: A short identifier for the test run.
8690
base_dir: The base GCS directory for storing outputs.
8791
tokenizer_path: Path to the tokenizer (HuggingFace
8892
model path).
@@ -96,7 +100,6 @@ def __init__(
96100
self.accelerator = accelerator
97101
self.slices = slices
98102
self.model_name = model_name
99-
self.short_id = short_id
100103
self.base_dir = base_dir
101104
self.tokenizer_path = tokenizer_path
102105
self.load_parameters_path = load_parameters_path
@@ -129,11 +132,9 @@ def generate_rl_training_command(
129132
f"model_name={self.model_name} "
130133
f"tokenizer_path={self.tokenizer_path} "
131134
f"load_parameters_path={self.load_parameters_path} "
132-
f"base_output_directory={self.base_dir}"
135+
f"base_output_directory={self.base_dir} "
136+
f"loss_algo={loss_algo.loss_name}"
133137
)
134138

135-
if loss_algo == LossAlgo.GSPO:
136-
command += f" loss_algo={loss_algo.value}-token"
137-
138139
# Return as tuple for k8s yaml compatibility.
139140
return (command,)

0 commit comments

Comments
 (0)