|
11 | 11 |
|
12 | 12 | from airflow import models |
13 | 13 | from airflow.models.baseoperator import chain |
| 14 | +from airflow.utils.task_group import TaskGroup |
14 | 15 |
|
15 | 16 | from dags import composer_env |
16 | 17 | from dags.common import test_owner |
|
69 | 70 | accelerator="v5p-128", |
70 | 71 | slices=[1], # Single slice for RL training |
71 | 72 | model_name="llama3.1-70b", |
72 | | - short_id="mrl", |
73 | 73 | base_dir=( |
74 | 74 | f"{test_config_util.DEFAULT_BUCKET}/llama3.1-70b-Instruct/outputs" |
75 | 75 | ), |
|
94 | 94 |
|
95 | 95 | for loss_algo in training_config.loss_algos: |
96 | 96 | 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}" |
106 | 99 |
|
107 | 100 | rl_training_command = training_config.generate_rl_training_command( |
108 | 101 | loss_algo=loss_algo, |
109 | 102 | run_name=run_name, |
110 | 103 | hf_token=HF_TOKEN_LLAMA3_1, |
111 | 104 | ) |
112 | 105 |
|
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 | + ) |
0 commit comments