Skip to content

Commit f86d5bc

Browse files
wangkuiyichanglan
authored andcommitted
Resolve circular dep in goodput monitor code with minimal change
GitOrigin-RevId: 7ee0b01
1 parent f770379 commit f86d5bc

4 files changed

Lines changed: 238 additions & 189 deletions

File tree

axlearn/cloud/gcp/measurement.py

Lines changed: 18 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -30,16 +30,16 @@
3030
from ml_goodput_measurement import monitoring as goodput_monitoring
3131

3232
from axlearn.cloud.common.utils import parse_kv_flags, to_bool
33-
from axlearn.common import measurement
33+
from axlearn.common import measurement_base
3434
from axlearn.common.config import REQUIRED, Required, config_class, maybe_set_config
3535

3636

37-
@measurement.register_recorder("goodput")
38-
class GoodputRecorder(measurement.Recorder):
37+
@measurement_base.register_recorder("goodput")
38+
class GoodputRecorder(measurement_base.Recorder):
3939
"""Records overall training goodput."""
4040

4141
@config_class
42-
class Config(measurement.Recorder.Config):
42+
class Config(measurement_base.Recorder.Config):
4343
"""Configures GoodputRecorder.
4444
4545
Attributes:
@@ -77,7 +77,7 @@ def from_flags(cls, fv: flags.FlagValues) -> "GoodputRecorder":
7777
- jax_backend: The type of jax backend.
7878
- enable_monitoring: Boolean to enable/disable goodput monitoring (default: true).
7979
"""
80-
cfg: measurement.Recorder.Config = cls.default_config()
80+
cfg: measurement_base.Recorder.Config = cls.default_config()
8181
parsed_flags = parse_kv_flags(fv.recorder_spec, delimiter="=")
8282
if "upload_interval" in parsed_flags:
8383
parsed_flags["upload_interval"] = int(parsed_flags["upload_interval"])
@@ -100,7 +100,7 @@ def __init__(self, cfg):
100100
self._logger_name = f"goodput_logger_{cfg.name}"
101101

102102
@contextlib.contextmanager
103-
def record_event(self, event: measurement.EventType, *args, **kwargs):
103+
def record_event(self, event: measurement_base.EventType, *args, **kwargs):
104104
"""Records a goodput event using a context manager."""
105105
# Lazily instantiate the recorder if it hasn't been already.
106106
if self._recorder is None:
@@ -225,7 +225,7 @@ def monitor_goodput():
225225

226226
return monitor_goodput()
227227

228-
def record(self, event: measurement.Event, *args, **kwargs):
228+
def record(self, event: measurement_base.Event, *args, **kwargs):
229229
"""Deprecated: `record()` is not used in GoodputRecorder.
230230
Use the record_event context manager instead.
231231
"""
@@ -238,27 +238,27 @@ def record(self, event: measurement.Event, *args, **kwargs):
238238
logging_enabled=(jax.process_index() == 0),
239239
)
240240

241-
if event == measurement.Event.START_JOB:
241+
if event == measurement_base.Event.START_JOB:
242242
self._recorder.record_job_start_time(*args, **kwargs)
243-
elif event == measurement.Event.END_JOB:
243+
elif event == measurement_base.Event.END_JOB:
244244
self._recorder.record_job_end_time(*args, **kwargs)
245-
elif event == measurement.Event.START_STEP:
245+
elif event == measurement_base.Event.START_STEP:
246246
self._recorder.record_step_start_time(*args, **kwargs)
247-
elif event == measurement.Event.START_ACCELERATOR_INIT:
247+
elif event == measurement_base.Event.START_ACCELERATOR_INIT:
248248
self._recorder.record_tpu_init_start_time(*args, **kwargs)
249-
elif event == measurement.Event.END_ACCELERATOR_INIT:
249+
elif event == measurement_base.Event.END_ACCELERATOR_INIT:
250250
self._recorder.record_tpu_init_end_time(*args, **kwargs)
251-
elif event == measurement.Event.START_TRAINING_PREPARATION:
251+
elif event == measurement_base.Event.START_TRAINING_PREPARATION:
252252
self._recorder.record_training_preparation_start_time(*args, **kwargs)
253-
elif event == measurement.Event.END_TRAINING_PREPARATION:
253+
elif event == measurement_base.Event.END_TRAINING_PREPARATION:
254254
self._recorder.record_training_preparation_end_time(*args, **kwargs)
255-
elif event == measurement.Event.START_DATA_LOADING:
255+
elif event == measurement_base.Event.START_DATA_LOADING:
256256
self._recorder.record_data_loading_start_time(*args, **kwargs)
257-
elif event == measurement.Event.END_DATA_LOADING:
257+
elif event == measurement_base.Event.END_DATA_LOADING:
258258
self._recorder.record_data_loading_end_time(*args, **kwargs)
259-
elif event == measurement.Event.START_CUSTOM_BADPUT_EVENT:
259+
elif event == measurement_base.Event.START_CUSTOM_BADPUT_EVENT:
260260
self._recorder.record_custom_badput_event_start_time(*args, **kwargs)
261-
elif event == measurement.Event.END_CUSTOM_BADPUT_EVENT:
261+
elif event == measurement_base.Event.END_CUSTOM_BADPUT_EVENT:
262262
self._recorder.record_custom_badput_event_end_time(*args, **kwargs)
263263
else:
264264
logging.log_first_n(

axlearn/common/measurement.py

Lines changed: 21 additions & 159 deletions
Original file line numberDiff line numberDiff line change
@@ -2,169 +2,31 @@
22

33
"""A library to measure e2e metrics like goodput."""
44

5-
import contextlib
6-
import enum
75
import importlib
8-
from typing import Optional, TypeVar
6+
from typing import Optional
97

108
from absl import flags, logging
119

12-
from axlearn.common.config import REQUIRED, Configurable, Required, config_class
13-
14-
15-
class Event(enum.Enum):
16-
"""Event to be recorded (Legacy).
17-
18-
Attributes:
19-
START_JOB: Start of job.
20-
END_JOB: End of job.
21-
START_STEP: Start of a training step. Should be recorded with `step` as a positional arg.
22-
START_ACCELERATOR_INIT: Start of accelerator mesh initialization.
23-
END_ACCELERATOR_INIT: End of accelerator mesh initialization.
24-
START_TRAINING_PREPARATION: Start of training preparation.
25-
END_TRAINING_PREPARATION: End of training preparation.
26-
START_DATA_LOADING: Start of data loading.
27-
END_DATA_LOADING: End of data loading.
28-
START_CUSTOM_BADPUT_EVENT: Start of custom badput event.
29-
END_CUSTOM_BADPUT_EVENT: End of custom badput event.
30-
"""
31-
32-
START_JOB = "START_JOB"
33-
END_JOB = "END_JOB"
34-
START_STEP = "START_STEP"
35-
START_ACCELERATOR_INIT = "START_ACCELERATOR_INIT"
36-
END_ACCELERATOR_INIT = "END_ACCELERATOR_INIT"
37-
START_TRAINING_PREPARATION = "START_TRAINING_PREPARATION"
38-
END_TRAINING_PREPARATION = "END_TRAINING_PREPARATION"
39-
START_DATA_LOADING = "START_DATA_LOADING"
40-
END_DATA_LOADING = "END_DATA_LOADING"
41-
START_CUSTOM_BADPUT_EVENT = "START_CUSTOM_BADPUT_EVENT"
42-
END_CUSTOM_BADPUT_EVENT = "END_CUSTOM_BADPUT_EVENT"
43-
44-
45-
class EventType(enum.Enum):
46-
"""Event to be recorded.
47-
48-
Attributes:
49-
JOB: Start and end of the job.
50-
STEP: Start of a training step. Should be recorded with `step` as a positional arg.
51-
ACCELERATOR_INIT: Start and end of accelerator mesh initialization.
52-
TRAINING_PREPARATION: Start and end of training preparation.
53-
DATA_LOADING: Start and end of data loading.
54-
CUSTOM_BADPUT_EVENT: Start and end of custom badput events.
55-
"""
56-
57-
JOB = "job"
58-
STEP = "step"
59-
ACCELERATOR_INIT = "tpu_init"
60-
TRAINING_PREPARATION = "training_preparation"
61-
DATA_LOADING = "data_loading"
62-
CUSTOM_BADPUT_EVENT = "custom_badput_event"
63-
64-
65-
class Recorder(Configurable):
66-
"""The base interface for collecting e2e metrics."""
67-
68-
@config_class
69-
class Config(Configurable.Config):
70-
"""Configures Recorder.
71-
72-
Attributes:
73-
name: Name of the recorder.
74-
"""
75-
76-
name: Required[str] = REQUIRED
77-
78-
@classmethod
79-
def from_flags(cls, fv: Optional[flags.FlagValues]) -> "Recorder":
80-
"""Converts flags to a recorder."""
81-
raise NotImplementedError(cls)
82-
83-
def record(self, event: Event, *args, **kwargs):
84-
"""Records a single, instantaneous event.
85-
86-
Note:
87-
This method is maintained for backward compatibility.
88-
New child recorder implementations should prioritize implementing the
89-
`record_event` context manager instead.
90-
"""
91-
raise NotImplementedError(type(self))
92-
93-
def start_monitoring(self, **kwargs):
94-
"""Starts computing and uploading metrics in the background.
95-
96-
Note:
97-
This method is maintained for backward compatibility. New child
98-
recorders should prioritize implementing the `maybe_monitor_all`
99-
context manager, which provides a clearer lifecycle for monitoring.
100-
"""
101-
raise NotImplementedError(type(self))
102-
103-
@contextlib.contextmanager
104-
def record_event(self, event: EventType, *args, **kwargs):
105-
"""A context manager to record the start and end of an event.
106-
107-
This is the preferred method for recording events. Child classes
108-
should implement this context manager to handle timed operations.
109-
110-
Example:
111-
with recorder.record_event(EventType.ACCELERATOR_INIT):
112-
# Device initialization.
113-
"""
114-
# pylint: disable=unnecessary-pass
115-
# pylint: disable=unused-argument
116-
try:
117-
yield
118-
finally:
119-
pass
120-
121-
@contextlib.contextmanager
122-
def maybe_monitor_all(self):
123-
"""Context manager to start and stop computing and monitoring metrics.
124-
125-
This is the preferred method for monitoring events. Child classes
126-
should implement this context manager to manage all types of monitoring.
127-
128-
Example:
129-
with recorder.maybe_monitor_all():
130-
# Train
131-
"""
132-
yield
133-
134-
135-
_recorders: dict[str, type] = {}
136-
_T = TypeVar("_T")
137-
138-
139-
def register_recorder(name: str):
140-
def fn(cls: _T) -> _T:
141-
"""Registers a recorder class for `get_recorder_config`."""
142-
if name in _recorders:
143-
raise ValueError(f"Recorder {name} is already registered.")
144-
_recorders[name] = cls
145-
return cls
146-
147-
return fn
148-
149-
150-
def define_flags(**kwargs):
151-
"""Common measurement flags."""
152-
153-
flags.DEFINE_string(
154-
"recorder_type",
155-
None,
156-
"The recorder type. It can be a recorder name, e.g. `my_recorder`, or "
157-
"a module paired with a recorder name, e.g. `my.module:my_recorder`.",
158-
**kwargs,
159-
)
160-
flags.DEFINE_multi_string(
161-
"recorder_spec",
162-
[],
163-
"Recorder spec provided as key=value. "
164-
"Refer to each recorders's `from_flags` method docstring for details.",
165-
**kwargs,
166-
)
167-
10+
from axlearn.common.measurement_base import (
11+
Event,
12+
EventType,
13+
Recorder,
14+
_recorders,
15+
define_flags,
16+
register_recorder,
17+
)
18+
19+
__all__ = [
20+
"Event",
21+
"EventType",
22+
"Recorder",
23+
"define_flags",
24+
"register_recorder",
25+
"global_recorder",
26+
"initialize",
27+
"record_event",
28+
"start_monitoring",
29+
]
16830

16931
global_recorder: Optional[Recorder] = None
17032

0 commit comments

Comments
 (0)