1414
1515"""Utilities to construct configs for maxtext DAG on GKE."""
1616
17- from dags .common import test_owner
17+ import datetime
18+ from typing import Any , Iterable
19+
20+ from dags import gcs_bucket
21+ from dags .common .vm_resource import Project , XpkClusters
1822from xlml .apis import gcp_config , metric_config , task , test_config
1923from xlml .apis .xpk_cluster_config import XpkClusterConfig
20- from dags import gcs_bucket
21- from dags .common .vm_resource import TpuVersion , Project , XpkClusters , GpuVersion , CpuVersion
22- from typing import Any , Iterable
23- import datetime
2424
2525
2626def get_gke_config (
@@ -31,7 +31,9 @@ def get_gke_config(
3131 run_model_cmds : Iterable [str ],
3232 cluster : XpkClusterConfig = XpkClusters .TPU_V4_8_MAXTEXT_CLUSTER ,
3333 num_slices : int = 1 ,
34- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
34+ dataset_name : metric_config .DatasetOption = (
35+ metric_config .DatasetOption .XLML_DATASET
36+ ),
3537 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
3638 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
3739 base_output_directory : str = None ,
@@ -90,7 +92,9 @@ def get_gke_config_with_interrupt(
9092 expect_reach_to_step : int ,
9193 cluster : XpkClusterConfig = XpkClusters .TPU_V4_8_MAXTEXT_CLUSTER ,
9294 num_slices : int = 1 ,
93- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
95+ dataset_name : metric_config .DatasetOption = (
96+ metric_config .DatasetOption .XLML_DATASET
97+ ),
9498 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
9599 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
96100 base_output_directory : str = None ,
@@ -153,7 +157,9 @@ def get_gke_config_with_name_gen_and_quarantine(
153157 run_model_cmds : Iterable [str ],
154158 cluster : XpkClusterConfig = XpkClusters .TPU_V4_8_MAXTEXT_CLUSTER ,
155159 num_slices : int = 1 ,
156- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
160+ dataset_name : metric_config .DatasetOption = (
161+ metric_config .DatasetOption .XLML_DATASET
162+ ),
157163 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
158164 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
159165 base_output_directory : str = None ,
@@ -216,7 +222,9 @@ def get_gke_maxtext_nightly_config(
216222 test_owner : str ,
217223 cluster : XpkClusterConfig = XpkClusters .TPU_V4_8_MAXTEXT_CLUSTER ,
218224 num_slices : int = 1 ,
219- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
225+ dataset_name : metric_config .DatasetOption = (
226+ metric_config .DatasetOption .XLML_DATASET
227+ ),
220228 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
221229 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
222230) -> task .XpkTask :
@@ -234,18 +242,25 @@ def get_gke_maxtext_nightly_config(
234242 base_output_directory = (
235243 f"{ gcs_bucket .BASE_OUTPUT_DIR } /maxtext/nightly/automated/{ current_date } "
236244 )
237- run_name = f"{ num_slices } slice-V{ cluster .device_version .value } _{ cluster .core_count } -maxtext-nightly-{ current_datetime } "
245+ run_name = (
246+ f"{ num_slices } slice-V{ cluster .device_version .value } _{ cluster .core_count } "
247+ f"-maxtext-nightly-{ current_datetime } "
248+ )
238249
239250 run_model_cmds = (
240251 "bash src/dependencies/scripts/preflight.sh PLATFORM=GKE" ,
241252 (
242253 "JAX_PLATFORM_NAME=TPU XLA_FLAGS='--xla_dump_to=/tmp/xla_dump/'"
243254 " ENABLE_PJRT_COMPATIBILITY=true"
244- f" python3 -m maxtext.trainers.pre_train.train src/maxtext/configs/base.yml run_name={ run_name } "
245- f" base_output_directory={ base_output_directory } "
255+ " python3 -m maxtext.trainers.pre_train.train"
256+ " src/maxtext/configs/base.yml"
257+ f" run_name={ run_name } base_output_directory={ base_output_directory } "
246258 " dataset_path=gs://max-datasets-rogue dataset_type=synthetic"
247- " model_name=llama3-8b per_device_batch_size=12 reuse_example_batch=1 metrics_file='metrics.txt'"
248- " steps=50 enable_checkpointing=false profiler=xplane upload_all_profiler_results=true skip_first_n_steps_for_profiler=10 profiler_steps=10 gcs_metrics=true"
259+ " model_name=llama3-8b per_device_batch_size=12 reuse_example_batch=1"
260+ " metrics_file='metrics.txt' steps=50 enable_checkpointing=false"
261+ " profiler=xplane upload_all_profiler_results=true"
262+ " skip_first_n_steps_for_profiler=10 profiler_steps=10"
263+ " gcs_metrics=true"
249264 ),
250265 )
251266
@@ -316,7 +331,9 @@ def get_gke_gpt3_6b_nightly_config(
316331 test_owner : str ,
317332 cluster : XpkClusterConfig = XpkClusters .TPU_V4_8_MAXTEXT_CLUSTER ,
318333 num_slices : int = 1 ,
319- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
334+ dataset_name : metric_config .DatasetOption = (
335+ metric_config .DatasetOption .XLML_DATASET
336+ ),
320337 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
321338 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
322339) -> task .XpkTask :
@@ -334,18 +351,26 @@ def get_gke_gpt3_6b_nightly_config(
334351 base_output_directory = (
335352 f"{ gcs_bucket .BASE_OUTPUT_DIR } /maxtext/nightly/automated/{ current_date } "
336353 )
337- run_name = f"{ num_slices } slice-V{ cluster .device_version .value } _{ cluster .core_count } -gpt3-6b-nightly-{ current_datetime } "
354+ run_name = (
355+ f"{ num_slices } slice-V{ cluster .device_version .value } _{ cluster .core_count } "
356+ f"-gpt3-6b-nightly-{ current_datetime } "
357+ )
338358
339359 run_model_cmds = (
340360 "bash src/dependencies/scripts/preflight.sh PLATFORM=GKE" ,
341361 (
342362 "JAX_PLATFORM_NAME=TPU XLA_FLAGS='--xla_dump_to=/tmp/xla_dump/'"
343363 " ENABLE_PJRT_COMPATIBILITY=true"
344- f" python3 -m maxtext.trainers.pre_train.train src/maxtext/configs/base.yml run_name={ run_name } model_name=gpt3-6b"
364+ " python3 -m maxtext.trainers.pre_train.train"
365+ " src/maxtext/configs/base.yml"
366+ f" run_name={ run_name } model_name=gpt3-6b"
345367 f" base_output_directory={ base_output_directory } "
346368 " dataset_path=gs://max-datasets-rogue dataset_type=synthetic"
347- " per_device_batch_size=12 reuse_example_batch=1 global_parameter_scale=1 metrics_file='metrics.txt'"
348- " steps=50 enable_checkpointing=false profiler=xplane upload_all_profiler_results=true skip_first_n_steps_for_profiler=10 profiler_steps=10 gcs_metrics=true"
369+ " per_device_batch_size=12 reuse_example_batch=1"
370+ " global_parameter_scale=1 metrics_file='metrics.txt' steps=50"
371+ " enable_checkpointing=false profiler=xplane"
372+ " upload_all_profiler_results=true skip_first_n_steps_for_profiler=10"
373+ " profiler_steps=10 gcs_metrics=true"
349374 ),
350375 )
351376
@@ -379,7 +404,9 @@ def get_maxtext_cpu_end_to_end_gke_config(
379404 cluster : XpkClusterConfig = XpkClusters .CPU_N2_STANDARD_64_CLUSTER ,
380405 machine_count : int = 1 ,
381406 num_slices : int = 1 ,
382- dataset_name : metric_config .DatasetOption = metric_config .DatasetOption .XLML_DATASET ,
407+ dataset_name : metric_config .DatasetOption = (
408+ metric_config .DatasetOption .XLML_DATASET
409+ ),
383410 dataset_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
384411 composer_project : str = Project .CLOUD_ML_AUTO_SOLUTIONS .value ,
385412 base_output_directory : str = None ,
0 commit comments