Skip to content

Commit e96aa55

Browse files
committed
Fix linting errors
1 parent 7565253 commit e96aa55

2 files changed

Lines changed: 49 additions & 21 deletions

File tree

dags/examples/maxtext_profile_namegen_example_dag.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,8 @@
1515
"""An example DAG to extract profile metrics.
1616
1717
Pretraining mixtral-8x7b model on 1xv4-128.
18-
Profile extraction can be easily integrated with gke_config + (to_name_gen_and_quarantine_task + run).
18+
Profile extraction can be easily integrated with gke_config
19+
+ (to_name_gen_and_quarantine_task + run).
1920
"""
2021

2122
import datetime

dags/multipod/configs/gke_config.py

Lines changed: 47 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -14,13 +14,13 @@
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
1822
from xlml.apis import gcp_config, metric_config, task, test_config
1923
from 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

2626
def 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

Comments
 (0)