Skip to content

Commit 575f45d

Browse files
committed
Fix: Added a new regex to make sure the restoring logs are caught correclty.
Ref: Refactor list_log_entries function to be more flexible.
1 parent f76e189 commit 575f45d

9 files changed

Lines changed: 67 additions & 38 deletions

dags/orbax/maxtext_emc_restore_gcs.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -170,7 +170,8 @@
170170
project_id=test_config.cluster.project,
171171
location=zone_to_region(test_config.cluster.zone),
172172
cluster_name=test_config.cluster.name,
173-
text_filter=" AND ".join(log_filters),
173+
text_filters=log_filters,
174+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
174175
start_time=start_time,
175176
end_time=end_time,
176177
)

dags/orbax/maxtext_emc_restore_local.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -166,7 +166,8 @@
166166
project_id=test_config.cluster.project,
167167
location=zone_to_region(test_config.cluster.zone),
168168
cluster_name=test_config.cluster.name,
169-
text_filter=" AND ".join(log_filters),
169+
text_filters=log_filters,
170+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
170171
start_time=start_time,
171172
end_time=end_time,
172173
)

dags/orbax/maxtext_emc_resume_gcs.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -235,7 +235,8 @@
235235
project_id=test_config.cluster.project,
236236
location=zone_to_region(test_config.cluster.zone),
237237
cluster_name=test_config.cluster.name,
238-
text_filter=" AND ".join(log_filters),
238+
text_filters=log_filters,
239+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
239240
start_time=start_time,
240241
end_time=end_time,
241242
)

dags/orbax/maxtext_mtc_restore_local.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -169,7 +169,8 @@
169169
project_id=test_config.cluster.project,
170170
location=zone_to_region(test_config.cluster.zone),
171171
cluster_name=test_config.cluster.name,
172-
text_filter=" AND ".join(log_filters),
172+
text_filters=log_filters,
173+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
173174
start_time=start_time,
174175
end_time=end_time,
175176
)

dags/orbax/maxtext_mtc_save_gcs.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -147,7 +147,8 @@
147147
project_id=test_config.cluster.project,
148148
location=zone_to_region(test_config.cluster.zone),
149149
cluster_name=test_config.cluster.name,
150-
text_filter="Successful: backup for step",
150+
text_filters=["Successful: backup for step"],
151+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
151152
namespace="gke-managed-checkpointing",
152153
container_name="replication-worker",
153154
pod_pattern="multitier-driver",

dags/orbax/maxtext_reg_restore_gcs_with_node_disruption.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -169,7 +169,8 @@ def generate_workload_checkpoints_location(gcs_ckpt_location: str) -> str:
169169
project_id=test_config.cluster.project,
170170
location=zone_to_region(test_config.cluster.zone),
171171
cluster_name=test_config.cluster.name,
172-
text_filter="\"'event_type': 'restore'\"",
172+
text_filters=[r"'event_type': 'restore'"],
173+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
173174
start_time=start_time,
174175
end_time=end_time,
175176
)

dags/orbax/util/validation_util.py

Lines changed: 46 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,17 @@
44
from typing import Optional
55
from absl import logging
66
import re
7+
from enum import Flag, auto
78

89
from airflow.decorators import task
910
from airflow.exceptions import AirflowFailException
1011
from google.cloud import logging as logging_api
1112
from xlml.apis import gcs
1213

14+
class FilterMode(Flag):
15+
textPayload = auto()
16+
jsonPayload_message = auto()
17+
1318

1419
@task
1520
def generate_timestamp():
@@ -60,11 +65,8 @@ def validate_checkpoint_at_steps_are_saved(
6065
location=location,
6166
cluster_name=cluster_name,
6267
pod_pattern=pod_pattern,
63-
text_filter=(
64-
"\"'event_type': 'save'\" AND "
65-
f'(textPayload=~"{log_pattern}" OR '
66-
f'jsonPayload.message=~"{log_pattern}")'
67-
),
68+
text_filters=[r"'event_type': 'save'"],
69+
filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message,
6870
start_time=start_time,
6971
end_time=end_time,
7072
)
@@ -137,7 +139,8 @@ def validate_log_with_gcs(
137139
namespace: str = "default",
138140
pod_pattern: str = ".*",
139141
container_name: Optional[str] = None,
140-
text_filter: Optional[str] = None,
142+
text_filters: Optional[list[str]] = None,
143+
filter_mode: Optional["FilterMode"] = None,
141144
start_time: Optional[datetime] = None,
142145
end_time: Optional[datetime] = None,
143146
) -> None:
@@ -158,8 +161,8 @@ def validate_log_with_gcs(
158161
namespace: The Kubernetes namespace. Defaults to "default".
159162
pod_pattern: A glob pattern to match pod names. Defaults to "*".
160163
container_name: An optional container name to filter logs by.
161-
text_filter: An optional string to filter log entries by their
162-
`textPayload`.
164+
text_filters: An optional list of strings to filter log entries.
165+
filter_mode: A FilterMode flag to specify which payload fields to search.
163166
start_time: The start time for log retrieval.
164167
end_time: The end time for log retrieval.
165168
@@ -181,7 +184,8 @@ def validate_log_with_gcs(
181184
namespace=namespace,
182185
pod_pattern=pod_pattern,
183186
container_name=container_name,
184-
text_filter=f'textPayload=~"{text_filter}"',
187+
text_filters=text_filters,
188+
filter_mode=filter_mode,
185189
start_time=start_time,
186190
end_time=end_time,
187191
)
@@ -352,7 +356,8 @@ def list_log_entries(
352356
namespace: str = "default",
353357
pod_pattern: str = ".*",
354358
container_name: Optional[str] = None,
355-
text_filter: Optional[str] = None,
359+
text_filters: Optional[list[str]] = None,
360+
filter_mode: Optional["FilterMode"] = None,
356361
start_time: Optional[datetime] = None,
357362
end_time: Optional[datetime] = None,
358363
) -> list[logging_api.LogEntry]:
@@ -373,8 +378,8 @@ def list_log_entries(
373378
namespace: Kubernetes namespace (defaults to "default")
374379
pod_pattern: Pattern to match pod names (defaults to "*")
375380
container_name: Optional container name to filter logs
376-
text_filter: Optional comma-separated string to
377-
filter log entries by textPayload content
381+
text_filters: Optional list of strings to filter log entries
382+
filter_mode: Optional FilterMode to specify which payload fields to search
378383
start_time: Optional start time for log retrieval
379384
(defaults to 12 hours ago)
380385
end_time: Optional end time for log retrieval (defaults to now)
@@ -408,8 +413,15 @@ def list_log_entries(
408413

409414
if container_name:
410415
conditions.append(f'resource.labels.container_name="{container_name}"')
411-
if text_filter:
412-
conditions.append(f"{text_filter}")
416+
if text_filters and filter_mode is not None:
417+
filters = []
418+
for txt in text_filters:
419+
if FilterMode.textPayload in filter_mode:
420+
filters.append(f'textPayload=~"{txt}"')
421+
if FilterMode.jsonPayload_message in filter_mode:
422+
filters.append(f'jsonPayload.message=~"{txt}"')
423+
if filters:
424+
conditions.append(f"({' OR '.join(filters)})")
413425

414426
log_filter = " AND ".join(conditions)
415427

@@ -425,7 +437,8 @@ def validate_log_exist(
425437
namespace: str = "default",
426438
pod_pattern: str = ".*",
427439
container_name: Optional[str] = None,
428-
text_filter: Optional[str] = None,
440+
text_filters: Optional[list[str]] = None,
441+
filter_mode: Optional["FilterMode"] = None,
429442
start_time: Optional[datetime] = None,
430443
end_time: Optional[datetime] = None,
431444
) -> None:
@@ -438,7 +451,8 @@ def validate_log_exist(
438451
namespace=namespace,
439452
pod_pattern=pod_pattern,
440453
container_name=container_name,
441-
text_filter=text_filter,
454+
text_filters=text_filters,
455+
filter_mode=filter_mode,
442456
start_time=start_time,
443457
end_time=end_time,
444458
)
@@ -460,21 +474,24 @@ def validate_restored_correct_checkpoint(
460474
) -> None:
461475
"""Validate the restored step is in the expected range."""
462476

477+
reg_save_event = r"'event_type': 'save'"
478+
reg_restor_event = r"'event_type': '(emergency_)?restore'"
479+
reg_restoring = r"restoring from this run's directory step (\d+)"
480+
463481
entries = list_log_entries(
464482
project_id=project_id,
465483
location=location,
466484
cluster_name=cluster_name,
467485
namespace="default",
468486
pod_pattern=pod_pattern,
469-
text_filter=(
470-
"(textPayload:\"'event_type'\" OR jsonPayload.message:\"'event_type'\")"
471-
),
487+
text_filters=[reg_save_event, reg_restor_event, reg_restoring],
488+
filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message,
472489
start_time=start_time,
473490
end_time=end_time,
474491
)
475492

476493
if not entries:
477-
raise AirflowFailException("No event_type found in the log.")
494+
raise AirflowFailException("No event_type or restore log found in the log.")
478495

479496
local_saved_steps_before_restore = []
480497
for entry in entries:
@@ -488,7 +505,7 @@ def validate_restored_correct_checkpoint(
488505
logging.warning(f"Could not extract message from log entry: {entry}")
489506
continue
490507

491-
if re.search(r"'event_type': 'save'", message):
508+
if re.search(reg_save_event, message):
492509
saved_step_match = re.search(r"'step': (\d+)", message)
493510
if not saved_step_match:
494511
raise AirflowFailException(
@@ -497,15 +514,15 @@ def validate_restored_correct_checkpoint(
497514

498515
local_saved_steps_before_restore.append(int(saved_step_match.group(1)))
499516

500-
elif re.search(r"'event_type': '(emergency_)?restore'", message):
517+
elif re.search(reg_restor_event, message) or re.search(reg_restoring, message):
501518
logging.info("Found restore event: %s", message)
502519
logging.info(
503520
"Saved steps before restore: %s", local_saved_steps_before_restore
504521
)
505522

506523
restored_step_match = re.search(
507524
r"'step':\s*(?:np\.int32\()?(\d+)", message
508-
)
525+
) or re.search(reg_restoring, message)
509526
restored_step = (
510527
int(restored_step_match.group(1)) if restored_step_match else None
511528
)
@@ -585,7 +602,8 @@ def validate_replicator_gcs_restore_log(
585602
namespace=namespace,
586603
pod_pattern=pod_pattern,
587604
container_name=container_name,
588-
text_filter="Restoring from backup",
605+
text_filters=["Restoring from backup"],
606+
filter_mode=FilterMode.textPayload | FilterMode.jsonPayload_message,
589607
start_time=start_time,
590608
end_time=end_time,
591609
)
@@ -680,7 +698,8 @@ def validate_replicator_gcs_backup_log(
680698
namespace=namespace,
681699
pod_pattern=pod_pattern,
682700
container_name=container_name,
683-
text_filter=f'textPayload=~"{step_regex}"',
701+
text_filters=[step_regex],
702+
filter_mode=FilterMode.textPayload,
684703
start_time=start_time,
685704
end_time=end_time,
686705
)
@@ -737,7 +756,8 @@ def validate_checkpoints_save_regular_axlearn(
737756
location=location,
738757
cluster_name=cluster_name,
739758
pod_pattern=pod_pattern,
740-
text_filter=f'jsonPayload.message=~"{log_pattern}"',
759+
text_filters=[log_pattern],
760+
filter_mode=FilterMode.jsonPayload_message,
741761
start_time=start_time,
742762
end_time=end_time,
743763
)

dags/post_training/maxtext_rl.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -150,7 +150,8 @@
150150
project_id=training_config.cluster.project,
151151
location=zone_to_region(training_config.cluster.zone),
152152
cluster_name=training_config.cluster.name,
153-
text_filter=f"\"'loss_algo': '{loss_algo.loss_name}'\"",
153+
text_filters=[f"'loss_algo': '{loss_algo.loss_name}'"],
154+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
154155
namespace="default",
155156
container_name="jax-tpu",
156157
pod_pattern=f"{loss_algo.value}.*",
@@ -165,7 +166,8 @@
165166
project_id=training_config.cluster.project,
166167
location=zone_to_region(training_config.cluster.zone),
167168
cluster_name=training_config.cluster.name,
168-
text_filter='"Post RL Training"',
169+
text_filters=["Post RL Training"],
170+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
169171
namespace="default",
170172
container_name="jax-tpu",
171173
pod_pattern=f"{loss_algo.value}.*",

dags/post_training/maxtext_sft.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -75,10 +75,11 @@ def validate_training(
7575
project_id=config.cluster.project,
7676
location=zone_to_region(config.cluster.zone),
7777
cluster_name=config.cluster.name,
78-
text_filter=(
79-
f"(jsonPayload.message: \"'event_type': 'save'\" "
80-
f"AND jsonPayload.message: \"'step': {steps}\")"
81-
),
78+
text_filters=[
79+
f"'event_type': 'save'.*'step': {steps}",
80+
f"'step': {steps}.*'event_type': 'save'"
81+
],
82+
filter_mode=validation_util.FilterMode.textPayload | validation_util.FilterMode.jsonPayload_message,
8283
namespace="default",
8384
container_name="jax-tpu",
8485
pod_pattern=f"{config.short_id}.*",

0 commit comments

Comments
 (0)