diff --git a/.github/workflows/dag-check.yml b/.github/workflows/dag-check.yml index 96fcb76bf..ac3982d0e 100644 --- a/.github/workflows/dag-check.yml +++ b/.github/workflows/dag-check.yml @@ -3,7 +3,6 @@ name: DAG Check on: pull_request: - branches: [master] types: [opened, synchronize, edited] push: diff --git a/.github/workflows/pyink-check.yml b/.github/workflows/pyink-check.yml index ea7a8f250..c5cb1cd83 100644 --- a/.github/workflows/pyink-check.yml +++ b/.github/workflows/pyink-check.yml @@ -2,11 +2,12 @@ name: Formatter on: pull_request: - branches: [master] types: [opened, synchronize, edited] push: branches: [master] + workflow_dispatch: {} + jobs: format_check: runs-on: ubuntu-latest diff --git a/.github/workflows/pylint-check.yml b/.github/workflows/pylint-check.yml index 5e23d2812..02f4cb31e 100644 --- a/.github/workflows/pylint-check.yml +++ b/.github/workflows/pylint-check.yml @@ -2,12 +2,13 @@ name: Linter on: pull_request: - branches: [master] types: [opened, synchronize, edited] push: branches: [master] + workflow_dispatch: {} + jobs: linting_check: runs-on: ubuntu-latest diff --git a/.github/workflows/require-checklist.yml b/.github/workflows/require-checklist.yml index d15d19d99..4da288575 100644 --- a/.github/workflows/require-checklist.yml +++ b/.github/workflows/require-checklist.yml @@ -2,10 +2,13 @@ name: Require Checklist on: pull_request: types: [opened, edited, synchronize] + + workflow_dispatch: {} + jobs: check_pr_body: runs-on: ubuntu-latest steps: - uses: mheap/require-checklist-action@v2 with: - requireChecklist: false # If this is true and there are no checklists detected, the action will fail \ No newline at end of file + requireChecklist: false # If this is true and there are no checklists detected, the action will fail diff --git a/.github/workflows/unit-test.yml b/.github/workflows/unit-test.yml index 4f1723cb2..b771ca7e9 100644 --- a/.github/workflows/unit-test.yml +++ b/.github/workflows/unit-test.yml @@ -3,7 +3,6 @@ name: Unit Test on: pull_request: - branches: [master] types: [opened, synchronize, edited] push: diff --git a/dags/orbax/maxtext_emc_restore_gcs.py b/dags/orbax/maxtext_emc_restore_gcs.py index 28420171b..aa04862fc 100644 --- a/dags/orbax/maxtext_emc_restore_gcs.py +++ b/dags/orbax/maxtext_emc_restore_gcs.py @@ -170,7 +170,8 @@ project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter=" AND ".join(log_filters), + text_filters=log_filters, + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/orbax/maxtext_emc_restore_local.py b/dags/orbax/maxtext_emc_restore_local.py index 7e3c04b3d..c9cba17e2 100644 --- a/dags/orbax/maxtext_emc_restore_local.py +++ b/dags/orbax/maxtext_emc_restore_local.py @@ -166,7 +166,8 @@ project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter=" AND ".join(log_filters), + text_filters=log_filters, + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/orbax/maxtext_emc_resume_gcs.py b/dags/orbax/maxtext_emc_resume_gcs.py index 0c2e62327..15cdc5b15 100644 --- a/dags/orbax/maxtext_emc_resume_gcs.py +++ b/dags/orbax/maxtext_emc_resume_gcs.py @@ -235,7 +235,8 @@ project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter=" AND ".join(log_filters), + text_filters=log_filters, + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/orbax/maxtext_mtc_restore_local.py b/dags/orbax/maxtext_mtc_restore_local.py index 66e8f445f..b2f2b0ed9 100644 --- a/dags/orbax/maxtext_mtc_restore_local.py +++ b/dags/orbax/maxtext_mtc_restore_local.py @@ -169,7 +169,8 @@ project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter=" AND ".join(log_filters), + text_filters=log_filters, + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/orbax/maxtext_mtc_save_gcs.py b/dags/orbax/maxtext_mtc_save_gcs.py index 2ecf1797d..3e1537922 100644 --- a/dags/orbax/maxtext_mtc_save_gcs.py +++ b/dags/orbax/maxtext_mtc_save_gcs.py @@ -147,7 +147,8 @@ project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter="Successful: backup for step", + text_filters=["Successful: backup for step"], + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, namespace="gke-managed-checkpointing", container_name="replication-worker", pod_pattern="multitier-driver", diff --git a/dags/orbax/maxtext_reg_restore_gcs_with_node_disruption.py b/dags/orbax/maxtext_reg_restore_gcs_with_node_disruption.py index e94cbdbed..68bb0b788 100644 --- a/dags/orbax/maxtext_reg_restore_gcs_with_node_disruption.py +++ b/dags/orbax/maxtext_reg_restore_gcs_with_node_disruption.py @@ -169,7 +169,8 @@ def generate_workload_checkpoints_location(gcs_ckpt_location: str) -> str: project_id=test_config.cluster.project, location=zone_to_region(test_config.cluster.zone), cluster_name=test_config.cluster.name, - text_filter="\"'event_type': 'restore'\"", + text_filters=[r"'event_type': 'restore'"], + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/orbax/util/validation_util.py b/dags/orbax/util/validation_util.py index 6c72b5769..cb9d15ae7 100644 --- a/dags/orbax/util/validation_util.py +++ b/dags/orbax/util/validation_util.py @@ -4,12 +4,17 @@ from typing import Optional from absl import logging import re +from enum import Flag, auto from airflow.decorators import task from airflow.exceptions import AirflowFailException from google.cloud import logging as logging_api from xlml.apis import gcs +class FilterMode(Flag): + textPayload = auto() + jsonPayload_message = auto() + @task def generate_timestamp(): @@ -60,11 +65,8 @@ def validate_checkpoint_at_steps_are_saved( location=location, cluster_name=cluster_name, pod_pattern=pod_pattern, - text_filter=( - "\"'event_type': 'save'\" AND " - f'(textPayload=~"{log_pattern}" OR ' - f'jsonPayload.message=~"{log_pattern}")' - ), + text_filters=[r"'event_type': 'save'"], + filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) @@ -137,7 +139,8 @@ def validate_log_with_gcs( namespace: str = "default", pod_pattern: str = ".*", container_name: Optional[str] = None, - text_filter: Optional[str] = None, + text_filters: Optional[list[str]] = None, + filter_mode: Optional["FilterMode"] = None, start_time: Optional[datetime] = None, end_time: Optional[datetime] = None, ) -> None: @@ -158,8 +161,8 @@ def validate_log_with_gcs( namespace: The Kubernetes namespace. Defaults to "default". pod_pattern: A glob pattern to match pod names. Defaults to "*". container_name: An optional container name to filter logs by. - text_filter: An optional string to filter log entries by their - `textPayload`. + text_filters: An optional list of strings to filter log entries. + filter_mode: A FilterMode flag to specify which payload fields to search. start_time: The start time for log retrieval. end_time: The end time for log retrieval. @@ -181,7 +184,8 @@ def validate_log_with_gcs( namespace=namespace, pod_pattern=pod_pattern, container_name=container_name, - text_filter=f'textPayload=~"{text_filter}"', + text_filters=text_filters, + filter_mode=filter_mode, start_time=start_time, end_time=end_time, ) @@ -352,7 +356,8 @@ def list_log_entries( namespace: str = "default", pod_pattern: str = ".*", container_name: Optional[str] = None, - text_filter: Optional[str] = None, + text_filters: Optional[list[str]] = None, + filter_mode: Optional["FilterMode"] = None, start_time: Optional[datetime] = None, end_time: Optional[datetime] = None, ) -> list[logging_api.LogEntry]: @@ -373,8 +378,8 @@ def list_log_entries( namespace: Kubernetes namespace (defaults to "default") pod_pattern: Pattern to match pod names (defaults to "*") container_name: Optional container name to filter logs - text_filter: Optional comma-separated string to - filter log entries by textPayload content + text_filters: Optional list of strings to filter log entries + filter_mode: Optional FilterMode to specify which payload fields to search start_time: Optional start time for log retrieval (defaults to 12 hours ago) end_time: Optional end time for log retrieval (defaults to now) @@ -408,8 +413,15 @@ def list_log_entries( if container_name: conditions.append(f'resource.labels.container_name="{container_name}"') - if text_filter: - conditions.append(f"{text_filter}") + if text_filters and filter_mode is not None: + filters = [] + for txt in text_filters: + if FilterMode.textPayload in filter_mode: + filters.append(f'textPayload=~"{txt}"') + if FilterMode.jsonPayload_message in filter_mode: + filters.append(f'jsonPayload.message=~"{txt}"') + if filters: + conditions.append(f"({' OR '.join(filters)})") log_filter = " AND ".join(conditions) @@ -425,7 +437,8 @@ def validate_log_exist( namespace: str = "default", pod_pattern: str = ".*", container_name: Optional[str] = None, - text_filter: Optional[str] = None, + text_filters: Optional[list[str]] = None, + filter_mode: Optional["FilterMode"] = None, start_time: Optional[datetime] = None, end_time: Optional[datetime] = None, ) -> None: @@ -438,7 +451,8 @@ def validate_log_exist( namespace=namespace, pod_pattern=pod_pattern, container_name=container_name, - text_filter=text_filter, + text_filters=text_filters, + filter_mode=filter_mode, start_time=start_time, end_time=end_time, ) @@ -460,21 +474,24 @@ def validate_restored_correct_checkpoint( ) -> None: """Validate the restored step is in the expected range.""" + reg_save_event = r"'event_type': 'save'" + reg_restor_event = r"'event_type': '(emergency_)?restore'" + reg_restoring = r"restoring from this run's directory step (\d+)" + entries = list_log_entries( project_id=project_id, location=location, cluster_name=cluster_name, namespace="default", pod_pattern=pod_pattern, - text_filter=( - "(textPayload:\"'event_type'\" OR jsonPayload.message:\"'event_type'\")" - ), + text_filters=[reg_save_event, reg_restor_event, reg_restoring], + filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) if not entries: - raise AirflowFailException("No event_type found in the log.") + raise AirflowFailException("No event_type or restore log found in the log.") local_saved_steps_before_restore = [] for entry in entries: @@ -488,7 +505,7 @@ def validate_restored_correct_checkpoint( logging.warning(f"Could not extract message from log entry: {entry}") continue - if re.search(r"'event_type': 'save'", message): + if re.search(reg_save_event, message): saved_step_match = re.search(r"'step': (\d+)", message) if not saved_step_match: raise AirflowFailException( @@ -497,7 +514,7 @@ def validate_restored_correct_checkpoint( local_saved_steps_before_restore.append(int(saved_step_match.group(1))) - elif re.search(r"'event_type': '(emergency_)?restore'", message): + elif re.search(reg_restor_event, message) or re.search(reg_restoring, message): logging.info("Found restore event: %s", message) logging.info( "Saved steps before restore: %s", local_saved_steps_before_restore @@ -505,7 +522,7 @@ def validate_restored_correct_checkpoint( restored_step_match = re.search( r"'step':\s*(?:np\.int32\()?(\d+)", message - ) + ) or re.search(reg_restoring, message) restored_step = ( int(restored_step_match.group(1)) if restored_step_match else None ) @@ -585,7 +602,8 @@ def validate_replicator_gcs_restore_log( namespace=namespace, pod_pattern=pod_pattern, container_name=container_name, - text_filter="Restoring from backup", + text_filters=["Restoring from backup"], + filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) @@ -680,7 +698,8 @@ def validate_replicator_gcs_backup_log( namespace=namespace, pod_pattern=pod_pattern, container_name=container_name, - text_filter=f'textPayload=~"{step_regex}"', + text_filters=[step_regex], + filter_mode=FilterMode.textPayload, start_time=start_time, end_time=end_time, ) @@ -737,7 +756,8 @@ def validate_checkpoints_save_regular_axlearn( location=location, cluster_name=cluster_name, pod_pattern=pod_pattern, - text_filter=f'jsonPayload.message=~"{log_pattern}"', + text_filters=[log_pattern], + filter_mode=FilterMode.jsonPayload_message, start_time=start_time, end_time=end_time, ) diff --git a/dags/post_training/maxtext_rl.py b/dags/post_training/maxtext_rl.py index 160d539ed..c85ef5cfc 100644 --- a/dags/post_training/maxtext_rl.py +++ b/dags/post_training/maxtext_rl.py @@ -150,7 +150,8 @@ project_id=training_config.cluster.project, location=zone_to_region(training_config.cluster.zone), cluster_name=training_config.cluster.name, - text_filter=f"\"'loss_algo': '{loss_algo.loss_name}'\"", + text_filters=[f"'loss_algo': '{loss_algo.loss_name}'"], + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, namespace="default", container_name="jax-tpu", pod_pattern=f"{loss_algo.value}.*", @@ -165,7 +166,8 @@ project_id=training_config.cluster.project, location=zone_to_region(training_config.cluster.zone), cluster_name=training_config.cluster.name, - text_filter='"Post RL Training"', + text_filters=["Post RL Training"], + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, namespace="default", container_name="jax-tpu", pod_pattern=f"{loss_algo.value}.*", diff --git a/dags/post_training/maxtext_sft.py b/dags/post_training/maxtext_sft.py index 906c92716..7d319379b 100644 --- a/dags/post_training/maxtext_sft.py +++ b/dags/post_training/maxtext_sft.py @@ -75,10 +75,11 @@ def validate_training( project_id=config.cluster.project, location=zone_to_region(config.cluster.zone), cluster_name=config.cluster.name, - text_filter=( - f"(jsonPayload.message: \"'event_type': 'save'\" " - f"AND jsonPayload.message: \"'step': {steps}\")" - ), + text_filters=[ + f"'event_type': 'save'.*'step': {steps}", + f"'step': {steps}.*'event_type': 'save'" + ], + filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message, namespace="default", container_name="jax-tpu", pod_pattern=f"{config.short_id}.*", diff --git a/dags/post_training/util/validation_util.py b/dags/post_training/util/validation_util.py index 873cd03f7..8def83fb9 100644 --- a/dags/post_training/util/validation_util.py +++ b/dags/post_training/util/validation_util.py @@ -16,6 +16,7 @@ from dags.orbax.util.validation_util import ( generate_timestamp, validate_log_exist, + FilterMode, ) @@ -78,4 +79,5 @@ def upload_to_vertex_ai( "generate_timestamp", "validate_log_exist", "upload_to_vertex_ai", + "FilterMode", ] diff --git a/scripts/code-style.sh b/scripts/code-style.sh index 36cfa13e9..ea50709ac 100755 --- a/scripts/code-style.sh +++ b/scripts/code-style.sh @@ -19,14 +19,37 @@ set -e FOLDERS_TO_FORMAT=("dags" "xlml") -for folder in "${FOLDERS_TO_FORMAT[@]}" -do - pyink "$folder" --pyink-indentation=2 --pyink-use-majority-quotes --line-length=80 --check --diff -done - -for folder in "${FOLDERS_TO_FORMAT[@]}" -do - pylint "./$folder" --fail-under=9.6 -done +HEAD_SHA="$(git rev-parse HEAD)" +BASE_BRANCH="dev" + +if ! git rev-parse --verify "$BASE_BRANCH" >/dev/null 2>&1; then + git fetch origin "$BASE_BRANCH":"$BASE_BRANCH" || { + echo "[code-style] base branch '$BASE_BRANCH' not found, skip diff-based check." + exit 0 + } +fi + +CHANGED_PY_FILES="$( + git diff --name-only --diff-filter=ACM "${BASE_BRANCH}" "${HEAD_SHA}" \ + | grep '\.py$' \ + | while read -r f; do + for folder in "${FOLDERS_TO_FORMAT[@]}"; do + if [[ "$f" == "$folder/"* ]]; then + echo "$f" + break + fi + done + done \ + | sort -u +)" + +if [[ -z "${CHANGED_PY_FILES}" ]]; then + echo "[pre-push hook] no changed files detected between ${HEAD_SHA} and ${BASE_BRANCH}" + exit 1 +fi + +pyink ${CHANGED_PY_FILES} --pyink-indentation=2 --pyink-use-majority-quotes --line-length=80 --check --diff + +pylint ${CHANGED_PY_FILES} --fail-under=9.6 --disable=E1123 echo "Successfully clean up all codes."