3030from ml_goodput_measurement import monitoring as goodput_monitoring
3131
3232from axlearn .cloud .common .utils import parse_kv_flags , to_bool
33- from axlearn .common import measurement
33+ from axlearn .common import measurement_base
3434from 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 (
0 commit comments