Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion .github/workflows/dag-check.yml
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@ name: DAG Check

on:
pull_request:
branches: [master]
types: [opened, synchronize, edited]

push:
Expand Down
3 changes: 2 additions & 1 deletion .github/workflows/pyink-check.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 2 additions & 1 deletion .github/workflows/pylint-check.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion .github/workflows/require-checklist.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
requireChecklist: false # If this is true and there are no checklists detected, the action will fail
1 change: 0 additions & 1 deletion .github/workflows/unit-test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@ name: Unit Test

on:
pull_request:
branches: [master]
types: [opened, synchronize, edited]

push:
Expand Down
3 changes: 2 additions & 1 deletion dags/orbax/maxtext_emc_restore_gcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
3 changes: 2 additions & 1 deletion dags/orbax/maxtext_emc_restore_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
3 changes: 2 additions & 1 deletion dags/orbax/maxtext_emc_resume_gcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
3 changes: 2 additions & 1 deletion dags/orbax/maxtext_mtc_restore_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
3 changes: 2 additions & 1 deletion dags/orbax/maxtext_mtc_save_gcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
72 changes: 46 additions & 26 deletions dags/orbax/util/validation_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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:
Expand All @@ -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.

Expand All @@ -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,
)
Expand Down Expand Up @@ -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]:
Expand All @@ -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)
Expand Down Expand Up @@ -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)

Expand All @@ -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:
Expand All @@ -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,
)
Expand All @@ -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:
Expand All @@ -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(
Expand All @@ -497,15 +514,15 @@ 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
)

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
)
Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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,
)
Expand Down
6 changes: 4 additions & 2 deletions dags/post_training/maxtext_rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}.*",
Expand All @@ -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}.*",
Expand Down
9 changes: 5 additions & 4 deletions dags/post_training/maxtext_sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}.*",
Expand Down
2 changes: 2 additions & 0 deletions dags/post_training/util/validation_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from dags.orbax.util.validation_util import (
generate_timestamp,
validate_log_exist,
FilterMode,
)


Expand Down Expand Up @@ -78,4 +79,5 @@ def upload_to_vertex_ai(
"generate_timestamp",
"validate_log_exist",
"upload_to_vertex_ai",
"FilterMode",
]
Loading