|
64 | 64 | 4. All training steps execute within expected parameters |
65 | 65 | 5. Metrics are successfully synced to Vertex AI TensorBoard |
66 | 66 | """, |
67 | | - concurrency=2, |
| 67 | + concurrency=1, |
68 | 68 | ) as dag: |
69 | 69 | training_config = test_config_util.RLTestConfig( |
70 | 70 | cluster=XpkClusters.TPU_V5P_128_CLUSTER, |
|
88 | 88 | # HF token retrieved from Airflow Variables for secure credential management |
89 | 89 | HF_TOKEN_LLAMA3_1 = models.Variable.get("HF_TOKEN_CIENET", None) |
90 | 90 |
|
| 91 | + tests = [] |
91 | 92 | for mode, image in test_config_util.POST_TRAINING_DOCKER_IMAGES: |
92 | 93 | # TODO: Enable stable mode once a new version of MaxText is available |
93 | 94 | if mode == test_config_util.SetupMode.STABLE: |
94 | 95 | continue # Skip stable for RL training tests |
95 | 96 |
|
96 | 97 | for loss_algo in training_config.loss_algos: |
97 | 98 | 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 | | - |
111 | 99 | with TaskGroup( |
112 | 100 | group_id=( |
113 | 101 | f"{loss_algo.value}-{mode.value}-" |
114 | 102 | f"{slice_num}x{training_config.accelerator}" |
115 | 103 | ) |
116 | 104 | ) as group: |
117 | 105 | 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 | + ) |
118 | 118 | start_time = validation_util.generate_timestamp.override( |
119 | 119 | task_id="generate_start_time" |
120 | 120 | )() |
121 | 121 |
|
122 | 122 | training_task = gke_config.get_gke_config( |
| 123 | + time_out_in_min=30, |
123 | 124 | num_slices=slice_num, |
124 | 125 | cluster=training_config.cluster, |
125 | | - time_out_in_min=30, |
126 | 126 | test_name=loss_algo.value, |
127 | 127 | run_model_cmds=rl_training_command, |
128 | 128 | docker_image=image.value, |
|
137 | 137 | )() |
138 | 138 |
|
139 | 139 | chain( |
| 140 | + run_name, |
140 | 141 | start_time, |
141 | 142 | training_task, |
142 | 143 | end_time, |
|
149 | 150 | project_id=training_config.cluster.project, |
150 | 151 | location=zone_to_region(training_config.cluster.zone), |
151 | 152 | 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}'\"", |
153 | 154 | namespace="default", |
154 | 155 | container_name="jax-tpu", |
155 | 156 | pod_pattern=f"{loss_algo.value}.*", |
|
183 | 184 | ) |
184 | 185 |
|
185 | 186 | chain(training_group, [validation_group, upload_to_vertex_ai]) |
| 187 | + tests.append(group) |
| 188 | + |
| 189 | + chain(*tests) |
0 commit comments