diff --git a/ml_goodput_measurement/src/checkpoint_badput_calculator.py b/ml_goodput_measurement/src/checkpoint_badput_calculator.py index 34a2110..eb5c2d3 100644 --- a/ml_goodput_measurement/src/checkpoint_badput_calculator.py +++ b/ml_goodput_measurement/src/checkpoint_badput_calculator.py @@ -382,6 +382,11 @@ def _save_statistics( ) return vertical_step_stats + @staticmethod + def _get_duration_stat(step_stats: Dict[str, str], key: str) -> float: + """Safely extracts a float duration in seconds, defaulting None to 0.0.""" + return float(step_stats.get(key) or 0.0) + def calculate_save_operation_checkpoint_manager_blocking_time( self, operation_type: Optional[str] = OPERATION_TYPE_PERSISTENT_AND_LOCAL, ) -> SaveCheckpointManagerVerticalStepStats: @@ -408,36 +413,42 @@ def calculate_save_operation_checkpoint_manager_blocking_time( step_already_processed[step_info].step = step_info step_already_processed[ step_info - ].total_checkpoint_manager_blocking_time = float( - step_stats[_CHECKPOINT_MANAGER_SAVE_DURATION_SECS] + ].total_checkpoint_manager_blocking_time = self._get_duration_stat( + step_stats, _CHECKPOINT_MANAGER_SAVE_DURATION_SECS ) step_already_processed[step_info].total_checkpointer_blocking_time = ( - float(step_stats[_CHECKPOINTER_SAVE_DURATION_SECS]) + self._get_duration_stat( + step_stats, _CHECKPOINTER_SAVE_DURATION_SECS + ) ) - step_already_processed[step_info].total_wait_for_prev_time = float( - step_stats[_WAIT_FOR_PREV_DURATION_SECS] + step_already_processed[step_info].total_wait_for_prev_time = ( + self._get_duration_stat(step_stats, _WAIT_FOR_PREV_DURATION_SECS) ) - step_already_processed[step_info].total_get_old_steps_time = float( - step_stats[_GET_OLD_STEPS_DURATION_SECS] + step_already_processed[step_info].total_get_old_steps_time = ( + self._get_duration_stat(step_stats, _GET_OLD_STEPS_DURATION_SECS) ) step_already_processed[step_info].occurrence = 1 else: step_already_processed[step_info].step = step_info step_already_processed[ step_info - ].total_checkpoint_manager_blocking_time += float( - step_stats[_CHECKPOINT_MANAGER_SAVE_DURATION_SECS] + ].total_checkpoint_manager_blocking_time += self._get_duration_stat( + step_stats, _CHECKPOINT_MANAGER_SAVE_DURATION_SECS ) step_already_processed[ step_info - ].total_checkpointer_blocking_time += float( - step_stats[_CHECKPOINTER_SAVE_DURATION_SECS] + ].total_checkpointer_blocking_time += self._get_duration_stat( + step_stats, _CHECKPOINTER_SAVE_DURATION_SECS ) - step_already_processed[step_info].total_wait_for_prev_time += float( - step_stats[_WAIT_FOR_PREV_DURATION_SECS] + step_already_processed[ + step_info + ].total_wait_for_prev_time += self._get_duration_stat( + step_stats, _WAIT_FOR_PREV_DURATION_SECS ) - step_already_processed[step_info].total_get_old_steps_time += float( - step_stats[_GET_OLD_STEPS_DURATION_SECS] + step_already_processed[ + step_info + ].total_get_old_steps_time += self._get_duration_stat( + step_stats, _GET_OLD_STEPS_DURATION_SECS ) step_already_processed[step_info].occurrence += 1 @@ -591,11 +602,15 @@ def calculate_restore_operation_checkpoint_manager_blocking_time( step_already_processed[step_info] = RestoreProcessedStep() step_already_processed[step_info].step = step_info - step_already_processed[step_info].total_checkpoint_manager_time = float( - step_stats[_CHECKPOINT_MANAGER_RESTORE_DURATION_SECS] + step_already_processed[step_info].total_checkpoint_manager_time = ( + self._get_duration_stat( + step_stats, _CHECKPOINT_MANAGER_RESTORE_DURATION_SECS + ) ) - step_already_processed[step_info].total_restore_time = float( - step_stats[_CHECKPOINTER_RESTORE_DURATION_SECS] + step_already_processed[step_info].total_restore_time = ( + self._get_duration_stat( + step_stats, _CHECKPOINTER_RESTORE_DURATION_SECS + ) ) if ( step_stats.get(_BROADCAST_DURATION_SECS) @@ -610,11 +625,13 @@ def calculate_restore_operation_checkpoint_manager_blocking_time( step_already_processed[step_info].step = step_info step_already_processed[ step_info - ].total_checkpoint_manager_time += float( - step_stats[_CHECKPOINT_MANAGER_RESTORE_DURATION_SECS] + ].total_checkpoint_manager_time += self._get_duration_stat( + step_stats, _CHECKPOINT_MANAGER_RESTORE_DURATION_SECS ) - step_already_processed[step_info].total_restore_time += float( - step_stats[_CHECKPOINTER_RESTORE_DURATION_SECS] + step_already_processed[ + step_info + ].total_restore_time += self._get_duration_stat( + step_stats, _CHECKPOINTER_RESTORE_DURATION_SECS ) if ( step_stats.get(_BROADCAST_DURATION_SECS) @@ -671,6 +688,3 @@ def calculate_restore_operation_checkpoint_manager_blocking_time( checkpoint_badput_calculator.calculate_save_operation_checkpoint_manager_blocking_time( args.operation_type ) - - - diff --git a/ml_goodput_measurement/src/goodput.py b/ml_goodput_measurement/src/goodput.py index 096417a..359053c 100644 --- a/ml_goodput_measurement/src/goodput.py +++ b/ml_goodput_measurement/src/goodput.py @@ -25,6 +25,7 @@ compute_baseline_step_time = goodput_utils.compute_baseline_step_time BadputType = goodput_utils.BadputType +RestartType = goodput_utils.RestartType CheckpointLoggerOptions = checkpoint_badput_calculator.CheckpointLoggerOptions CheckpointBadputCalculator = ( checkpoint_badput_calculator.CheckpointBadputCalculator @@ -661,8 +662,6 @@ def __init__( self._last_disruption_time = None # pyrefly: ignore[bad-assignment] self._last_disrupted_step = None # pyrefly: ignore[bad-assignment] self._historical_step_times = {} - self._pending_credit_step = None - self._pending_credit_duration = 0.0 def sync_to_gcs(self): """Syncs the underlying cache files to GCS.""" @@ -782,6 +781,297 @@ def _extract_custom_sync_intervals( return intervals + def _get_restart_event( + self, payload: dict[str, Any] + ) -> Optional[tuple[float, RestartType]]: + """Returns `(timestamp, RestartType)` if `payload` marks a job restart. + + Subclasses (such as `ElasticGoodputCalculator`) override this hook to + recognize additional in-process restart events (e.g. `elastic_wait_start` + or `elastic_reinit_start`) when `job_start_time` is not emitted. + + Args: + payload: Dictionary payload of a Cloud Logging entry. + + Returns: + A tuple of `(restart_timestamp, RestartType)` if the entry indicates a + restart, or None otherwise. + """ + if _JOB_START_TIME in payload: + return float(payload[_JOB_START_TIME]), RestartType.INFRA_RESTART + return None + + def _extract_checkpoint_save_events( + self, entries: list[Any] + ) -> list[tuple[int, bool, Optional[float], float]]: + """Extracts `(save_step, is_local, save_start_time, save_duration)` tuples. + + Args: + entries: A list of Cloud Logging entry payloads. + + Returns: + A list of tuples `(save_step, is_local, save_start_time, save_duration)` + for all checkpoint save operations with positive blocking duration. + """ + save_events = [] + for entry in entries: + if not isinstance(entry, dict): + continue + if ( + entry.get('event_type') == 'save' + and 'step' in entry + and entry.get('step') is not None + ): + save_step = int(entry['step']) + is_local = not str(entry.get('directory', '')).startswith('gs://') + raw_start_time = entry.get('checkpoint_manager_blocking_start_time') + save_start_time = ( + float(raw_start_time) if raw_start_time is not None else None + ) + save_duration = float( + entry.get('checkpoint_manager_blocking_duration_secs') or 0.0 + ) + if save_duration > 0.0: + save_events.append( + (save_step, is_local, save_start_time, save_duration) + ) + return save_events + + @staticmethod + def _get_checkpoint_save_by_type_in_window( + save_events: list[tuple[int, bool, Optional[float], float]], + step_num: int, + window_start_time: float, + window_end_time: Optional[float] = None, + ) -> dict[bool, float]: + """Computes average blocking save duration per locality within a window. + + When multiple workers log a checkpoint save for the same `(step_num, + is_local)` within the step execution window, their durations are averaged + (matching `CheckpointBadputCalculator`), while local and persistent saves + are tracked separately. + + Args: + save_events: Extracted `(save_step, is_local, save_start_time, + save_duration)` tuples. + step_num: The training step number to match. + window_start_time: Inclusive start timestamp of the step execution window. + window_end_time: Optional exclusive end timestamp of the step execution + window. + + Returns: + A dictionary mapping `is_local` (`bool`) to the average blocking save + duration (in seconds) for `step_num` within `[window_start_time, + window_end_time)`. + """ + durations_by_locality: dict[bool, list[float]] = {} + for save_step, is_local, save_start_time, save_duration in save_events: + if save_step != step_num: + continue + if save_start_time is None or ( + save_start_time >= window_start_time + and (window_end_time is None or save_start_time < window_end_time) + ): + durations_by_locality.setdefault(is_local, []).append(save_duration) + return { + is_local: sum(durations) / len(durations) + for is_local, durations in durations_by_locality.items() + if durations + } + + @classmethod + def _get_checkpoint_save_duration_in_window( + cls, + save_events: list[tuple[int, bool, Optional[float], float]], + step_num: int, + window_start_time: float, + window_end_time: Optional[float] = None, + ) -> float: + """Computes synchronous checkpoint save blocking duration within a window. + + Args: + save_events: Extracted `(save_step, is_local, save_start_time, + save_duration)` tuples. + step_num: The training step number to match. + window_start_time: Inclusive start timestamp of the step execution window. + window_end_time: Optional exclusive end timestamp of the step execution + window. + + Returns: + Total synchronous checkpoint save blocking duration (in seconds) for + `step_num` that started within `[window_start_time, window_end_time)`. + """ + return sum( + cls._get_checkpoint_save_by_type_in_window( + save_events, step_num, window_start_time, window_end_time + ).values() + ) + + def _compute_segment_startup_overhead( + self, step_times: list[float], min_step: int + ) -> float: + """Computes XLA compilation/startup extra time for a segment of steps. + + Uses the current segment's step times to establish a steady-state baseline + when `len(step_times) > 1`, or falls back to prior segments' historical step + times when a post-restart segment has only 1 completed step. Also updates + `self._historical_step_times` for the compilation step to exclude the + startup overhead. + + Args: + step_times: List of completed step durations (with custom sync and + checkpoint save blocking times already removed) in the current segment. + min_step: The starting step number of the current segment. + + Returns: + The excess startup/compilation duration (`PROGRAM_STARTUP` badput) in + seconds. + """ + steps_in_segment = len(step_times) + if steps_in_segment == 0: + return 0.0 + + # Attribute the maximum excess step time in the startup window to program + # startup. + startup_window = min(_STARTUP_COMPILATION_WINDOW, steps_in_segment) + prior_step_times = [ + duration + for step_idx, duration in self._historical_step_times.items() + if step_idx < min_step or step_idx >= min_step + steps_in_segment + ] + if steps_in_segment > 1: + baseline_step_time = compute_baseline_step_time(step_times) + elif prior_step_times: + # When a post-restart segment has only 1 completed step before another + # disruption, use prior segments' historical step times as the baseline so + # compilation overhead on that single step is not miscounted as goodput. + baseline_step_time = compute_baseline_step_time(prior_step_times) + else: + baseline_step_time = step_times[0] + + compilation_step_offset = max( + range(startup_window), key=lambda i: step_times[i] + ) + has_reference_baseline = steps_in_segment > 1 or bool(prior_step_times) + startup_extra_time = ( + max(0.0, step_times[compilation_step_offset] - baseline_step_time) + if has_reference_baseline + else 0.0 + ) + + # Adjust the compilation step's historical time to exclude startup overhead. + compilation_step = min_step + compilation_step_offset + if ( + startup_extra_time > 0.0 + and compilation_step in self._historical_step_times + ): + self._historical_step_times[compilation_step] -= startup_extra_time + + return startup_extra_time + + def _estimate_sequential_restart_step_time( + self, + prev_step: int, + segment_productive_time: float, + num_completed_steps: int, + ) -> float: + """Estimates the productive duration of the tail step before a restart. + + When a job checkpoints at `prev_step`, restarts, and resumes sequentially at + `prev_step + 1`, `prev_step` has no subsequent step in the closed segment. + We estimate its execution time from the segment's average productive step + time (or historical step times if `prev_step` was the only step in the + segment). + + Args: + prev_step: The step number of the last step before the sequential restart. + segment_productive_time: Total productive time of completed steps in the + closed segment. + num_completed_steps: Number of completed step intervals in the closed + segment. + + Returns: + Estimated productive execution duration (in seconds) for `prev_step`. + """ + if num_completed_steps > 0: + return segment_productive_time / num_completed_steps + if prev_step in self._historical_step_times: + return self._historical_step_times[prev_step] + if self._historical_step_times: + return compute_baseline_step_time( + list(self._historical_step_times.values()) + ) + return 0.0 + + def _compute_salvaged_historical_time( + self, + disrupted_segment_min_step: int, + resume_step: int, + completed_productive_steps: set[int], + ) -> float: + """Computes salvaged productive time across a jump-forward rollback. + + When an earlier run progressed past `resume_step`, a second run rolled back + to `disrupted_segment_min_step`, and a third run resumed from a snapshot at + `resume_step`, steps in `[disrupted_segment_min_step, resume_step)` were + already completed in the earlier run and are salvaged from + `self._historical_step_times`. + + Args: + disrupted_segment_min_step: The starting step of the disrupted segment. + resume_step: The step at which the new run resumed (`curr_step`). + completed_productive_steps: Set of productive step numbers to update. + + Returns: + Total salvaged productive time (in seconds) from historical steps. + """ + salvaged_historical_productive_time = 0.0 + for step_idx in range(disrupted_segment_min_step, resume_step): + salvaged_historical_productive_time += self._historical_step_times.get( + step_idx, 0.0 + ) + completed_productive_steps.add(step_idx) + return salvaged_historical_productive_time + + @staticmethod + def _compute_remaining_checkpoint_save_badput( + save_events: list[tuple[int, bool, Optional[float], float]], + completed_productive_steps: set[int], + deducted_save_keys: set[tuple[int, bool]], + ) -> float: + """Sums checkpoint save durations for completed steps not deducted in-window. + + Serves as a fallback when checkpoint save logs omit start timestamps or use + synthetic timestamps outside step windows, while ensuring uncompleted tail + steps never have checkpoint save time subtracted from productive time. + Tracks `(save_step, is_local)` separately and averages across multiple + worker occurrences per key to match `CheckpointBadputCalculator`. + + Args: + save_events: Extracted `(save_step, is_local, save_start_time, + save_duration)` tuples. + completed_productive_steps: Set of steps that completed productively. + deducted_save_keys: Set of `(save_step, is_local)` keys whose checkpoint + save time was already deducted in-window. + + Returns: + Remaining checkpoint save blocking time (in seconds) to deduct from total + productive training time. + """ + remaining_by_key: dict[tuple[int, bool], list[float]] = {} + for save_step, is_local, _, save_duration in save_events: + save_key = (save_step, is_local) + if ( + save_step in completed_productive_steps + and save_key not in deducted_save_keys + ): + remaining_by_key.setdefault(save_key, []).append(save_duration) + return sum( + sum(durations) / len(durations) + for durations in remaining_by_key.values() + if durations + ) + def _get_current_productive_and_unproductive_time( self, interval_query: Optional[bool] = False ) -> tuple[ @@ -801,10 +1091,20 @@ def _get_current_productive_and_unproductive_time( (dict of BadputType and unproductive time), the last productive step and the last recorded step. """ + if interval_query: + entries_to_process = self._interval_entries + else: + with self._goodput_cache_lock: + entries_to_process = list(self._goodput_cache.get_cached_entries()) + + save_events = self._extract_checkpoint_save_events(entries_to_process) + deducted_save_keys: set[tuple[int, bool]] = set() + completed_productive_steps: set[int] = set() + def _compute_adjusted_segment_productive_and_unproductive_time( step_items: list[tuple[int, float]], curr_step: int, - min_step: int, + min_step: int, # pylint: disable=unused-argument custom_sync_intervals: list[tuple[float, float, str]], ) -> tuple[ float, @@ -817,25 +1117,23 @@ def _compute_adjusted_segment_productive_and_unproductive_time( """Computes adjusted productive and unproductive time for a segment of steps. This helper function calculates the total productive time, and the - breakdown of time lost due to custom badput events, as well as - wasted progress caused by disruptions. + breakdown of time lost due to custom badput events and synchronous + checkpoint saves, as well as wasted progress caused by disruptions. Args: step_items: A list of tuples, where each tuple contains a step number (int) and its start timestamp (float). curr_step: The current step number indicating the end of the segment. min_step: The minimum step number indicating the start of the segment. - custom_sync_intervals: A list of tuples, where each tuple consists of: - - start_time (float): Start timestamp of the sync event. - - end_time (float): End timestamp of the sync event. - - sync_type (str): The type of sync event. + custom_sync_intervals: A list of (start_time, end_time, sync_type) + tuples representing custom sync events. Returns: A tuple containing: - total_productive_time (float): Adjusted time excluding custom - sync durations. - - step_times (list[float]): List of adjusted times for all productive - steps in the segment. + sync and synchronous checkpoint save durations. + - step_times (list[float]): List of adjusted times for all + productive steps in the segment. - wasted_progress (float): Total unproductive time due to possible disruptions. - custom_sync_breakdown (dict[str, float]): @@ -870,9 +1168,25 @@ def _compute_adjusted_segment_productive_and_unproductive_time( ) adjusted_delta = max(0.0, raw_delta - custom_sync_in_interval) - self._historical_step_times[prev_step] = adjusted_delta + + # Deduct synchronous checkpoint save blocking duration that occurred + # during `prev_step` before computing startup overhead so save time is + # not misclassified as `PROGRAM_STARTUP` or subtracted twice. + save_by_type = self._get_checkpoint_save_by_type_in_window( + save_events, prev_step, prev_time, curr_time + ) + checkpoint_save_in_step = sum(save_by_type.values()) + if ( + checkpoint_save_in_step > 0.0 + and adjusted_delta >= checkpoint_save_in_step + ): + adjusted_delta -= checkpoint_save_in_step + if curr_step_num <= curr_step and curr_step_num - 1 == prev_step: + for is_local in save_by_type: + deducted_save_keys.add((prev_step, is_local)) if curr_step_num <= curr_step: + self._historical_step_times[prev_step] = adjusted_delta if curr_step_num - 1 != prev_step: continue @@ -880,10 +1194,21 @@ def _compute_adjusted_segment_productive_and_unproductive_time( step_times.append(adjusted_delta) steps_in_segment += 1 max_productive_step_count = prev_step + completed_productive_steps.add(prev_step) else: + # Preserve rolled-back step times for potential jump-forward salvage + # only if not already recorded by an earlier productive segment. + if prev_step not in self._historical_step_times: + self._historical_step_times[prev_step] = adjusted_delta + if ( + checkpoint_save_in_step > 0.0 + and raw_delta >= checkpoint_save_in_step + ): + for is_local in save_by_type: + deducted_save_keys.add((prev_step, is_local)) # These steps are after curr_step, they are lost due to disruption. - wasted_progress += raw_delta + wasted_progress += max(0.0, raw_delta - checkpoint_save_in_step) return ( total_productive_time, @@ -968,26 +1293,10 @@ def _get_segment_productive_and_unproductive_time( 0, ) - # Attribute the maximum excess step time in the startup window to program startup. - startup_window = min(_STARTUP_COMPILATION_WINDOW, steps_in_segment) - baseline_step_time = compute_baseline_step_time(step_times) - compilation_step_offset = max( - range(startup_window), key=lambda i: step_times[i] - ) - startup_extra_time = ( - max(0.0, step_times[compilation_step_offset] - baseline_step_time) - if steps_in_segment > 1 - else 0.0 + startup_extra_time = self._compute_segment_startup_overhead( + step_times, min_step ) - # Adjust the compilation step's historical time to exclude startup overhead. - compilation_step = min_step + compilation_step_offset - if ( - startup_extra_time > 0.0 - and compilation_step in self._historical_step_times - ): - self._historical_step_times[compilation_step] -= startup_extra_time - # Compute final segment metrics. ( final_adjusted_productive_time, @@ -1007,18 +1316,18 @@ def _get_segment_productive_and_unproductive_time( # Build a deserialized dictionary from cloud logging entries to store step # start times. The dictionary maps from step count to start time and will be - # used to each step's productive time by looking for its completion in the - # next step's start. + # used to compute each step's productive time by looking for its completion + # in the next step's start. # Note in the instance where progress is lost due to a disruption and the # last successful checkpoint did not include all the steps, the last set of # records of the step information will be kept and the previous set will be - # overwritten by design so as to correct for the the previously computed + # overwritten by design so as to correct for the previously computed # additional time that was counted as productive but lost due to a # disruption. productive_training_time = 0.0 total_unproductive_time = {} step_start_data = {} - job_start_time = None + pending_restart_time: Optional[float] = None job_end_time = None tpu_init_start_time = None training_prep_start_time = None @@ -1028,223 +1337,219 @@ def _get_segment_productive_and_unproductive_time( data_loading_badput = 0.0 sync_data_loading = True current_sync_data_loading = None - if interval_query: - entries_to_process = self._interval_entries - else: - with self._goodput_cache_lock: - entries_to_process = list(self._goodput_cache.get_cached_entries()) - self._number_of_interruptions = 0 - for payload in entries_to_process: - if _JOB_START_TIME in payload: - # Keep track of the latest start to compute badput due to disruption. - job_start_time = payload[_JOB_START_TIME] - - if _STEP_START_TIME in payload: - curr_step = int(payload[_STEP_COUNT]) - if curr_step not in step_start_data: - is_jump_forward = False - is_sequential_restart = False - if step_start_data: - prev_step = list(step_start_data.keys())[-1] - if curr_step > prev_step + 1: - is_jump_forward = True - self._number_of_interruptions += 1 - self._last_disrupted_step = prev_step - self._last_disruption_time = step_start_data[prev_step] - - # Compute segment productive and unproductive time. - ( - segment_productive_time, - segment_unproductive_time, - _, - ) = _get_segment_productive_and_unproductive_time( - step_start_data, curr_step, entries_to_process - ) - - # Calculate salvaged progress from historical step times. - salvaged_run_1 = 0.0 - segment_2_min_step = min(step_start_data.keys()) - for s in range(segment_2_min_step, curr_step): - salvaged_run_1 += self._historical_step_times.get(s, 0.0) - - # Adjust wasted progress: we lost the progress in current segment - # (segment_productive_time) but recovered salvaged_run_1. - net_wasted = segment_productive_time - salvaged_run_1 - segment_unproductive_time[ - BadputType.WASTED_PROGRESS_FROM_DISRUPTION - ] = net_wasted - - if ( - job_start_time is not None - and self._last_disruption_time is not None - and job_start_time > self._last_disruption_time - ): - infrastructure_disruption_badput = ( - job_start_time - self._last_disruption_time - ) - existing_infra = segment_unproductive_time.get( - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION, 0.0 - ) - infra_val = ( - existing_infra - if isinstance(existing_infra, (int, float)) - else 0.0 - ) - segment_unproductive_time[ - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION - ] = infra_val + infrastructure_disruption_badput - - # Accumulate metrics - productive_training_time += salvaged_run_1 - self._accumulate_unproductive_time( - segment_unproductive_time, total_unproductive_time - ) + def _add_badput_to_segment( + seg_unprod: UnproductiveTimeDict, badput_type: BadputType, amount: float + ) -> None: + """Adds `amount` to `badput_type` inside `seg_unprod` if positive.""" + if amount <= 0.0: + return + existing = seg_unprod.get(badput_type, 0.0) + val = existing if isinstance(existing, (int, float)) else 0.0 + seg_unprod[badput_type] = val + amount + + def _handle_segment_restart( + mode: str, + curr_step: int, + prev_step: int, + restart_ts: Optional[float], + ) -> None: + """Closes a disrupted step segment and accumulates its Goodput/Badput. + + Handles three restart resume modes: + - `jump_forward`: The workload resumed at `curr_step > prev_step + 1` + using a checkpoint/snapshot from an earlier run. Salvages historical + step times for `[disrupted_segment_min_step, curr_step)`. + - `sequential`: A restart occurred between `prev_step` and + `curr_step == prev_step + 1`. Credits `prev_step` with an estimated + productive duration so the last step before checkpoint/restart is not + discarded. + - `duplicate`: The workload rolled back to a previously executed step + (`curr_step <= prev_step`). All steps `< curr_step` are preserved as + productive, while steps `>= curr_step` become wasted progress. + + Unproductive time from the disruption is separated into two buckets: + 1. Wasted training progress (`WASTED_PROGRESS_FROM_DISRUPTION`) and + infrastructure downtime until the restart begins + (`INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION`). + 2. Re-initialization overhead (`TPU_INITIALIZATION`, `TRAINING_PREP`, + `DATA_LOADING_SYNC`, `PROGRAM_STARTUP`, checkpoint restore) after + the restart begins. - # Reset step_start_data for the new segment - step_start_data = {curr_step: payload[_STEP_START_TIME]} + Args: + mode: Restart resume mode (`'jump_forward'`, `'sequential'`, or + `'duplicate'`). + curr_step: The step number at which training resumed after the restart. + prev_step: The last recorded step number in the disrupted segment. + restart_ts: Timestamp of the earliest restart event following + `prev_step`, if available. + """ + nonlocal productive_training_time, sync_data_loading, current_sync_data_loading + self._number_of_interruptions += 1 + self._last_disrupted_step = prev_step + self._last_disruption_time = step_start_data[prev_step] - elif ( - job_start_time is not None - and job_start_time > step_start_data[prev_step] - ): - # Sequential restart with job_start_time between prev_step and curr_step - is_sequential_restart = True - self._number_of_interruptions += 1 - self._last_disrupted_step = prev_step - self._last_disruption_time = step_start_data[prev_step] - - ( - segment_productive_time, - segment_unproductive_time, - _, - ) = _get_segment_productive_and_unproductive_time( - step_start_data, prev_step, entries_to_process - ) + boundary_step = prev_step if mode == 'sequential' else curr_step + ( + segment_productive_time, + segment_unproductive_time, + _, + ) = _get_segment_productive_and_unproductive_time( + step_start_data, boundary_step, entries_to_process + ) - num_completed_steps = len(step_start_data) - 1 - if num_completed_steps > 0: - avg_step_time = segment_productive_time / num_completed_steps - elif prev_step in self._historical_step_times: - avg_step_time = self._historical_step_times[prev_step] - else: - avg_step_time = 0.0 - - # Credit prev_step productive time - productive_training_time += ( - segment_productive_time + avg_step_time - ) + estimated_last_step_time = 0.0 + if mode == 'jump_forward': + disrupted_segment_min_step = min(step_start_data.keys()) + salvaged_historical_productive_time = ( + self._compute_salvaged_historical_time( + disrupted_segment_min_step, + curr_step, + completed_productive_steps, + ) + ) + # Adjust wasted progress: we lost the progress in the current segment + # (`segment_productive_time`) but recovered + # `salvaged_historical_productive_time` from the earlier run. + segment_unproductive_time[ + BadputType.WASTED_PROGRESS_FROM_DISRUPTION + ] = (segment_productive_time - salvaged_historical_productive_time) + productive_training_time += salvaged_historical_productive_time + elif mode == 'sequential': + num_completed_steps = len(step_start_data) - 1 + estimated_last_step_time = self._estimate_sequential_restart_step_time( + prev_step, segment_productive_time, num_completed_steps + ) + productive_training_time += ( + segment_productive_time + estimated_last_step_time + ) + if estimated_last_step_time > 0.0: + self._historical_step_times[prev_step] = estimated_last_step_time + else: + # Duplicate-step rollback: steps >= curr_step were rolled back. If + # `curr_step` rolled back past the start of the current segment + # (`disrupted_segment_min_step`), also move previously completed + # historical steps in `[curr_step, disrupted_segment_min_step)` from + # productive time to wasted progress. + disrupted_segment_min_step = min(step_start_data.keys()) + rolled_back_historical_time = 0.0 + for step_idx in range(curr_step, disrupted_segment_min_step): + if step_idx in completed_productive_steps: + rolled_back_historical_time += self._historical_step_times.get( + step_idx, 0.0 + ) + for step_idx in list(completed_productive_steps): + if step_idx >= curr_step: + completed_productive_steps.discard(step_idx) + if rolled_back_historical_time > 0.0: + productive_training_time = max( + 0.0, productive_training_time - rolled_back_historical_time + ) + _add_badput_to_segment( + segment_unproductive_time, + BadputType.WASTED_PROGRESS_FROM_DISRUPTION, + rolled_back_historical_time, + ) + productive_training_time += segment_productive_time + + # When the job restarts, the first data loading event is synchronous. + sync_data_loading = True + if current_sync_data_loading is not None: + _add_badput_to_segment( + segment_unproductive_time, + BadputType.DATA_LOADING_SYNC, + current_sync_data_loading, + ) + current_sync_data_loading = None - # Add downtime minus the credited prev_step duration - infrastructure_disruption_badput = max( - 0.0, - (job_start_time - self._last_disruption_time) - avg_step_time, - ) - existing_infra = segment_unproductive_time.get( - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION, 0.0 - ) - infra_val = ( - existing_infra - if isinstance(existing_infra, (int, float)) - else 0.0 - ) - segment_unproductive_time[ - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION - ] = infra_val + infrastructure_disruption_badput + # Compute infrastructure recovery downtime between the disrupted step + # start time and the restart timestamp, excluding any credited productive + # step duration and synchronous checkpoint save time on `prev_step`. + if ( + restart_ts is not None + and self._last_disruption_time is not None + and restart_ts > self._last_disruption_time + ): + disrupted_save_by_type = self._get_checkpoint_save_by_type_in_window( + [e for e in save_events if e[2] is not None], + prev_step, + self._last_disruption_time, + restart_ts, + ) + disrupted_step_checkpoint_save_time = sum( + disrupted_save_by_type.values() + ) + for is_local in disrupted_save_by_type: + deducted_save_keys.add((prev_step, is_local)) + infrastructure_disruption_badput = max( + 0.0, + (restart_ts - self._last_disruption_time) + - estimated_last_step_time + - disrupted_step_checkpoint_save_time, + ) + _add_badput_to_segment( + segment_unproductive_time, + BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION, + infrastructure_disruption_badput, + ) - self._accumulate_unproductive_time( - segment_unproductive_time, total_unproductive_time - ) + self._accumulate_unproductive_time( + segment_unproductive_time, total_unproductive_time + ) - step_start_data = {curr_step: payload[_STEP_START_TIME]} + self._number_of_interruptions = 0 + for payload in entries_to_process: + restart_event = self._get_restart_event(payload) + if restart_event is not None: + restart_ts, restart_type = restart_event + if step_start_data: + last_step_ts = step_start_data[list(step_start_data.keys())[-1]] + if restart_ts > last_step_ts: + # Keep the earliest restart timestamp after `last_step_ts` (e.g. + # `elastic_wait_start_time` before `elastic_reinit_start_time`), + # unless overridden by a full `INFRA_RESTART` (`job_start_time`). + if ( + pending_restart_time is None + or pending_restart_time <= last_step_ts + or restart_type == RestartType.INFRA_RESTART + ): + pending_restart_time = restart_ts + sync_data_loading = True + current_sync_data_loading = None + else: + pending_restart_time = restart_ts - if not is_jump_forward and not is_sequential_restart: - step_start_data[curr_step] = payload[_STEP_START_TIME] + if _STEP_START_TIME in payload: + curr_step = int(payload[_STEP_COUNT]) + if not step_start_data: + step_start_data[curr_step] = payload[_STEP_START_TIME] + pending_restart_time = None else: - # In this case, the job restarted from Step (curr_step). It means that - # all progress till Step (curr_step - 1) has been preserved. So we - # can get the productive time since the previous start/restart and - # then clear the step_start_data dict. - self._number_of_interruptions += 1 - self._last_disrupted_step = list(step_start_data.keys())[-1] - self._last_disruption_time = step_start_data[ - self._last_disrupted_step - ] - - # Compute segment productive and unproductive time. - ( - segment_productive_time, - segment_unproductive_time, - _, - ) = _get_segment_productive_and_unproductive_time( - step_start_data, curr_step, entries_to_process - ) - # Accumulate the segment productive time. - productive_training_time += segment_productive_time - - # When the job restarts, data loading is synchronous. - sync_data_loading = True - if current_sync_data_loading is not None: - existing_sync = segment_unproductive_time.get( - BadputType.DATA_LOADING_SYNC, 0.0 + prev_step = list(step_start_data.keys())[-1] + if curr_step > prev_step + 1: + _handle_segment_restart( + 'jump_forward', curr_step, prev_step, pending_restart_time ) - sync_val = ( - existing_sync - if isinstance(existing_sync, (int, float)) - else 0.0 + step_start_data = {curr_step: payload[_STEP_START_TIME]} + pending_restart_time = None + elif curr_step <= prev_step: + _handle_segment_restart( + 'duplicate', curr_step, prev_step, pending_restart_time ) - segment_unproductive_time[BadputType.DATA_LOADING_SYNC] = ( - sync_val + current_sync_data_loading - ) - current_sync_data_loading = None - - # Since the current step has been recorded again, the progress - # between the previously recorded curr_step and recently recorded - # curr_step has been lost to a disruption and partially recovered - # due to a checkpoint of curr_step - 1. Accumulate the lost time in - # this segment as unproductive time. - # Note this unproductive time is divided into two buckets: - # 1. Wasted training progress after the last successfully - # checkpointed step and the disruption time until the job - # restarts. - # 2. TPU re-init, training prep, data loading, program startup, - # checkpoint loading etc. after the job restarts and before - # training progress resumes. - - # The first bucket can be calculated as the time between the start - # time of curr_step and the job restart time immediately prior. - if ( - job_start_time is not None - and self._last_disruption_time is not None - and job_start_time > self._last_disruption_time + step_start_data = {curr_step: payload[_STEP_START_TIME]} + pending_restart_time = None + elif ( + pending_restart_time is not None + and pending_restart_time > step_start_data[prev_step] ): - # Add the additional time it took for the job to restart after last - # interruption. These conditions are only met when the job is - # restarted after a disruption. - infrastructure_disruption_badput = ( - job_start_time - self._last_disruption_time - ) - existing_infra = segment_unproductive_time.get( - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION, 0.0 - ) - infra_val = ( - existing_infra - if isinstance(existing_infra, (int, float)) - else 0.0 + _handle_segment_restart( + 'sequential', curr_step, prev_step, pending_restart_time ) - segment_unproductive_time[ - BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION - ] = infra_val + infrastructure_disruption_badput - - # The second bucket is individually computed either from recorded - # logs (TPU initialization, training preparation, data loading) or - # computed from the first step time after start or restart - # (segment unproductive time). All unproductive time is accumulated - # as we go. - self._accumulate_unproductive_time( - segment_unproductive_time, total_unproductive_time - ) - step_start_data = {curr_step: payload[_STEP_START_TIME]} + step_start_data = {curr_step: payload[_STEP_START_TIME]} + pending_restart_time = None + else: + step_start_data[curr_step] = payload[_STEP_START_TIME] + pending_restart_time = None if _JOB_END_TIME in payload: # Locate the last instance of job's end time if the job has completed. @@ -1284,12 +1589,13 @@ def _get_segment_productive_and_unproductive_time( ) data_loading_badput += current_sync_data_loading if sync_data_loading: - # When the job starts, data loading is synchronous. + # When the job starts or restarts, data loading is synchronous. total_unproductive_time[BadputType.DATA_LOADING_SYNC] = ( total_unproductive_time.get(BadputType.DATA_LOADING_SYNC, 0) + current_sync_data_loading ) sync_data_loading = False + current_sync_data_loading = None data_loading_start_time = None # Compute unproductive time from checkpoint manager save and restore. @@ -1352,11 +1658,33 @@ def _get_segment_productive_and_unproductive_time( # Only consider the last step productive if the job has completed. if job_end_time is not None: - productive_training_time += job_end_time - step_start_data[last_step] + last_step_delta = job_end_time - step_start_data[last_step] + last_step_save_by_type = self._get_checkpoint_save_by_type_in_window( + save_events, last_step, step_start_data[last_step], job_end_time + ) + last_step_checkpoint_save = sum(last_step_save_by_type.values()) + if ( + last_step_checkpoint_save > 0.0 + and last_step_delta >= last_step_checkpoint_save + ): + last_step_delta -= last_step_checkpoint_save + for is_local in last_step_save_by_type: + deducted_save_keys.add((last_step, is_local)) + productive_training_time += last_step_delta max_productive_step_count = last_step + completed_productive_steps.add(last_step) - # Remove blocking checkpoint manager save time from productive time. - productive_training_time -= checkpoint_manager_save_badput + # Deduct any remaining checkpoint save time that occurred on completed + # productive steps but whose timestamp was not already deducted in-window. + remaining_completed_save_badput = ( + self._compute_remaining_checkpoint_save_badput( + save_events, completed_productive_steps, deducted_save_keys + ) + ) + if remaining_completed_save_badput > 0.0: + productive_training_time = max( + 0.0, productive_training_time - remaining_completed_save_badput + ) # Return a tuple of the total productive training time, the total # unproductive time (dict of BadputType and unproductive time) and the last @@ -1804,12 +2132,55 @@ def _compute_job_goodput_interval_metrics( self._number_of_interruptions, ) - def _get_step_times(self, entries: list[Any]): - """Helper function to compute step times from the log entries.""" + def _get_step_times(self, entries: list[Any]) -> dict[int, float]: + """Computes per-step execution durations from log entries. + + Excludes synchronous checkpoint save blocking durations and breaks step + pairings across job or elastic restarts so recovery downtime does not + pollute `ideal_step_time` or `step_time_deviation`. + + Args: + entries: List of raw or `(timestamp, payload)` cached log entries. + + Returns: + Dictionary mapping step count (`int`) to execution duration (`float`). + """ + normalized_entries = [] + for item in entries: + if ( + isinstance(item, tuple) + and len(item) == 2 + and isinstance(item[1], dict) + ): + normalized_entries.append(item[1]) + else: + normalized_entries.append(item) + + save_events = self._extract_checkpoint_save_events(normalized_entries) step_times = {} previous_step_start_time = None previous_step_count = None - for payload in entries: + + for payload in normalized_entries: + if not isinstance(payload, dict): + continue + + # When a restart occurs after `previous_step_start_time`, do not pair + # `previous_step_count` directly with the post-restart step timestamp. + restart_event = self._get_restart_event(payload) + if restart_event is not None and previous_step_start_time is not None: + restart_ts, _ = restart_event + if restart_ts > previous_step_start_time: + if ( + previous_step_count is not None + and previous_step_count in self._historical_step_times + ): + step_times[previous_step_count] = self._historical_step_times[ + previous_step_count + ] + previous_step_start_time = None + previous_step_count = None + if _STEP_START_TIME in payload: step_start_time = payload[_STEP_START_TIME] step_count = int(payload[_STEP_COUNT]) @@ -1818,9 +2189,25 @@ def _get_step_times(self, entries: list[Any]): and previous_step_count is not None and step_count == previous_step_count + 1 ): - step_times[previous_step_count] = ( - step_start_time - previous_step_start_time + step_delta = step_start_time - previous_step_start_time + checkpoint_save_in_step = ( + self._get_checkpoint_save_duration_in_window( + save_events, + previous_step_count, + previous_step_start_time, + step_start_time, + ) ) + if ( + checkpoint_save_in_step > 0.0 + and step_delta >= checkpoint_save_in_step + ): + step_delta -= checkpoint_save_in_step + if previous_step_count in self._historical_step_times: + step_delta = min( + step_delta, self._historical_step_times[previous_step_count] + ) + step_times[previous_step_count] = step_delta previous_step_count = step_count previous_step_start_time = step_start_time return step_times @@ -1845,7 +2232,7 @@ def _compute_step_info_and_update_cache( return step_info.step_deviations with self._goodput_cache_lock: - process_entries = self._goodput_cache.get_step_entries() + process_entries = self._goodput_cache.get_cached_entries() step_times = self._get_step_times(process_entries) diff --git a/ml_goodput_measurement/src/goodput_elastic.py b/ml_goodput_measurement/src/goodput_elastic.py index 989112c..d20c388 100644 --- a/ml_goodput_measurement/src/goodput_elastic.py +++ b/ml_goodput_measurement/src/goodput_elastic.py @@ -234,6 +234,84 @@ def _compute_time_weighted_efficiency( weighted_available / total_duration, ) + def _get_restart_event( + self, payload: dict[str, Any] + ) -> Optional[tuple[float, goodput_utils.RestartType]]: + """Identifies job or in-process elastic restart events from a log entry. + + In Pathways elastic training, slice-down and scale-up recoveries occur + in-process without logging a new `job_start_time`. Detecting + `elastic_wait_start_time` (or `elastic_reinit_start_time` as a fallback) + ensures pre-restart step segments are closed and post-restart recovery + is properly classified. + + Args: + payload: Dictionary payload of a Cloud Logging entry. + + Returns: + A tuple of `(restart_timestamp, RestartType)` if the entry marks a + restart, or None otherwise. + """ + restart_event = super()._get_restart_event(payload) + if restart_event is not None: + return restart_event + + # Elastic wait marks the beginning of an in-process scale-up or slice-down. + if _ELASTIC_WAIT_START_TIME in payload: + wait_event_type = str(payload.get(_ELASTIC_WAIT_EVENT_TYPE, '')).lower() + restart_type = ( + goodput_utils.RestartType.ELASTIC_SCALE_UP + if 'scale_up' in wait_event_type + else goodput_utils.RestartType.ELASTIC_SLICE_DOWN + ) + return float(payload[_ELASTIC_WAIT_START_TIME]), restart_type + + # Fallback if elastic_reinit_start_time is logged without a preceding wait. + if _ELASTIC_REINIT_START_TIME in payload: + return ( + float(payload[_ELASTIC_REINIT_START_TIME]), + goodput_utils.RestartType.ELASTIC_SLICE_DOWN, + ) + return None + + @staticmethod + def _extract_restore_intervals( + entries: list[dict[str, Any]], + ) -> list[tuple[float, float]]: + """Extracts `[start_ts, end_ts]` intervals for checkpoint restore events. + + Used to discount `UNPRODUCTIVE_CHECKPOINT_RESTORE_TIME` when a checkpoint + restore occurs inside an `ELASTIC_REINITIALIZATION` window so the restore + duration is not double-counted. Groups by `(step, is_local)` and averages + across multiple worker occurrences to match `CheckpointBadputCalculator`. + + Args: + entries: List of Cloud Logging entry dictionaries. + + Returns: + A list of `(restore_start_ts, restore_end_ts)` tuples. + """ + grouped_restores: dict[tuple[Any, bool], list[tuple[float, float]]] = {} + for entry in entries: + if entry.get('event_type') in ('restore', 'emergency_restore'): + restore_start = entry.get('checkpoint_manager_start_time') + restore_duration = entry.get('checkpoint_manager_duration_secs') + if restore_start is not None and restore_duration is not None: + restore_start_ts = float(restore_start) + restore_duration_secs = float(restore_duration) + if restore_duration_secs > 0.0: + is_local = not str(entry.get('directory', '')).startswith('gs://') + key = (entry.get('step'), is_local) + grouped_restores.setdefault(key, []).append( + (restore_start_ts, restore_duration_secs) + ) + restore_intervals = [] + for items in grouped_restores.values(): + avg_start = sum(s for s, _ in items) / len(items) + avg_duration = sum(d for _, d in items) / len(items) + restore_intervals.append((avg_start, avg_start + avg_duration)) + return restore_intervals + def _get_current_productive_and_unproductive_time( self, interval_query: Optional[bool] = False ) -> tuple[float, goodput.UnproductiveTimeDict, int, int]: @@ -268,13 +346,18 @@ def _get_current_productive_and_unproductive_time( + (end - start) ) - # On elastic retries TPU_INIT and TRAINING_PREP are already counted inside - # ELASTIC_REINITIALIZATION, so discount the overlap. + # On elastic retries TPU_INIT, TRAINING_PREP, and CHECKPOINT_RESTORE are + # already counted inside ELASTIC_REINITIALIZATION, so discount the overlap. if reinit_intervals: tpu_intervals, prep_intervals = self._extract_init_intervals(entries) + restore_intervals = self._extract_restore_intervals(entries) for bt, intervals in ( (goodput_utils.BadputType.TPU_INITIALIZATION, tpu_intervals), (goodput_utils.BadputType.TRAINING_PREP, prep_intervals), + ( + goodput_utils.BadputType.UNPRODUCTIVE_CHECKPOINT_RESTORE_TIME, + restore_intervals, + ), ): overlap = self._overlap_with_reinit(intervals, reinit_intervals) if overlap > 0 and bt in unproductive_time: diff --git a/ml_goodput_measurement/src/goodput_utils.py b/ml_goodput_measurement/src/goodput_utils.py index 96c748d..eaca367 100644 --- a/ml_goodput_measurement/src/goodput_utils.py +++ b/ml_goodput_measurement/src/goodput_utils.py @@ -13,7 +13,7 @@ Retry = retry.Retry -_TIME_ENTRY = 'time' +_TIMESTAMP_SUFFIXES = ('_time', '_timestamp', '_at') _METADATA_SERVER_URL = 'http://metadata.google.internal/computeMetadata/v1/' _METADATA_HEADERS = {'Metadata-Flavor': 'Google'} @@ -97,6 +97,14 @@ class BadputType(enum.Enum): ELASTIC_REINITIALIZATION = 14 +class RestartType(enum.Enum): + """The type of job or elastic restart.""" + + INFRA_RESTART = 'infra_restart' + ELASTIC_SLICE_DOWN = 'elastic_slice_down' + ELASTIC_SCALE_UP = 'elastic_scale_up' + + class WorkloadMetricDetails(TypedDict): goodput_time: dict[GoodputType, float] badput_time: dict[BadputType, float | dict[str, float]] @@ -318,16 +326,37 @@ def get_extra_time_from_anomalous_steps(step_times: list[Any]) -> float: def get_timestamp_from_log_entry( entry: dict[str, Any], ) -> Optional[datetime.datetime]: - """Helper function to get the timestamp from a log entry.""" - timestamp_posix_time = [ - entry_value + """Extracts the earliest valid POSIX timestamp from a structured log entry. + + Matches keys ending with `_time`, `_timestamp`, or `_at` + (`_TIMESTAMP_SUFFIXES`) + to structurally exclude duration fields (such as `*_duration_secs` or + `time_between_consecutive_saves_sec`) and ignores `None` or boolean values. + + Args: + entry: Dictionary payload of a Cloud Logging entry. + + Returns: + The earliest UTC datetime found in the entry, or None if no valid timestamp + field exists. + """ + # Only consider numeric (non-boolean) values for keys with timestamp suffixes. + valid_timestamps = [ + float(entry_value) for entry_label, entry_value in entry.items() - if _TIME_ENTRY in entry_label + if entry_label.endswith(_TIMESTAMP_SUFFIXES) + and isinstance(entry_value, (int, float)) + and not isinstance(entry_value, bool) ] - if timestamp_posix_time: - return datetime.datetime.fromtimestamp( - timestamp_posix_time[0], datetime.timezone.utc - ) + if valid_timestamps: + try: + # When an entry contains multiple phase timestamps (e.g., checkpoint save + # sub-operations), return the earliest timestamp representing event start. + return datetime.datetime.fromtimestamp( + min(valid_timestamps), datetime.timezone.utc + ) + except (ValueError, TypeError, OverflowError): + pass return None diff --git a/ml_goodput_measurement/tests/goodput_elastic_test.py b/ml_goodput_measurement/tests/goodput_elastic_test.py index 386b4f5..3c2fb7e 100644 --- a/ml_goodput_measurement/tests/goodput_elastic_test.py +++ b/ml_goodput_measurement/tests/goodput_elastic_test.py @@ -178,6 +178,7 @@ def setUp(self): self.goodput_calculator = goodput_elastic.ElasticGoodputCalculator( self.job_name, self.logger_name, self.mock_cloud_logger ) + self.goodput_calculator._goodput_cache.clear_cache() def test_compute_time_weighted_efficiency(self): @@ -313,6 +314,149 @@ def test_get_job_goodput_details_elastic(self): details['available_slice_efficiency'], expected_avail_eff ) + def test_elastic_sequential_restart_without_job_start_time(self): + """Validates in-process elastic restart splits segment, credits step, and discounts restore overlap.""" + job_start = datetime.datetime( + 2026, 8, 27, 16, 0, 0, tzinfo=datetime.timezone.utc + ) + self.goodput_recorder.record_job_start_time(job_start) + + # Segment 1: Steps 0, 1, 2 (each 2s steady-state; Step 0 has +5s startup) + s0 = job_start + self.goodput_recorder.record_step_start_time(0, s0) + s1 = s0 + datetime.timedelta(seconds=7) # 2s + 5s compilation + self.goodput_recorder.record_step_start_time(1, s1) + s2 = s1 + datetime.timedelta(seconds=2) + self.goodput_recorder.record_step_start_time(2, s2) + + # Step 2 executes for 2s (finishes at s2 + 2s), then 10s uninstrumented gap + # before elastic_wait_start_time (no job_start_time logged!). + wait_start = s2 + datetime.timedelta(seconds=12) + self.goodput_recorder.record_elastic_wait_start_time( + 'slice_down', wait_start + ) + wait_end = wait_start + datetime.timedelta(seconds=3) + self.goodput_recorder.record_elastic_wait_end_time('slice_down', wait_end) + + # Elastic reinit: 20s total, containing 12s checkpoint restore + reinit_start = wait_end + reinit_end = reinit_start + datetime.timedelta(seconds=20) + self.goodput_recorder.record_elastic_reinit_start_time(reinit_start) + self.mock_cloud_logger.write_cloud_logging_entry({ + 'event_type': 'restore', + 'step': 2, + 'directory': 'gs://bucket/path', + 'checkpoint_manager_start_time': reinit_start.timestamp() + 2.0, + 'checkpoint_manager_duration_secs': 12.0, + 'checkpointer_duration_secs': 12.0, + }) + self.goodput_recorder.record_elastic_reinit_end_time(reinit_end) + + # Synchronous data loading after elastic reinit: 4s + dl_start = reinit_end + dl_end = dl_start + datetime.timedelta(seconds=4) + self.goodput_recorder.record_data_loading_start_time(dl_start) + self.goodput_recorder.record_data_loading_end_time(dl_end) + + # Segment 2 resumes sequentially at Step 3 (+6s recompilation on Step 3) + s3 = dl_end + self.goodput_recorder.record_step_start_time(3, s3) + s4 = s3 + datetime.timedelta(seconds=8) # 2s + 6s compilation + self.goodput_recorder.record_step_start_time(4, s4) + job_end = s4 + datetime.timedelta(seconds=2) + self.goodput_recorder.record_job_end_time(job_end) + + self.goodput_calculator._fetch_new_entries(job_end) + prod, unprod, max_step, _ = ( + self.goodput_calculator._get_current_productive_and_unproductive_time() + ) + + # 5 steps (0, 1, 2, 3, 4) * 2s = 10s productive time + self.assertAlmostEqual(prod, 10.0, delta=0.1) + self.assertEqual(max_step, 4) + # Startup compilation = 5s (Segment 1) + 6s (Segment 2) = 11s + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.PROGRAM_STARTUP], 11.0, delta=0.1 + ) + # Infra recovery gap before elastic_wait = (12s - 2s credited Step 2) = 10s + self.assertAlmostEqual( + unprod[ + goodput_utils.BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION + ], + 10.0, + delta=0.1, + ) + # Elastic slice down = 3s, Elastic reinit = 20s, Restore discounted to 0s + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.ELASTIC_SLICE_DOWN], 3.0, delta=0.1 + ) + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.ELASTIC_REINITIALIZATION], + 20.0, + delta=0.1, + ) + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.UNPRODUCTIVE_CHECKPOINT_RESTORE_TIME], + 0.0, + delta=0.1, + ) + # Post-reinit data loading classified as DATA_LOADING_SYNC (4s) + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.DATA_LOADING_SYNC], 4.0, delta=0.1 + ) + + def test_multi_worker_restore_inside_elastic_reinit(self): + """Validates multi-worker restore logs inside elastic_reinit are averaged.""" + t0 = datetime.datetime(2026, 8, 27, 16, 0, 0, tzinfo=datetime.timezone.utc) + self.goodput_recorder.record_job_start_time(t0) + + # Initial job restore at Step 0 outside elastic_reinit: 5s + self.mock_cloud_logger.write_cloud_logging_entry({ + 'event_type': 'restore', + 'step': 0, + 'directory': 'gs://bucket/path', + 'checkpoint_manager_start_time': t0.timestamp() + 1.0, + 'checkpoint_manager_duration_secs': 5.0, + 'checkpointer_duration_secs': 5.0, + }) + + s0 = t0 + datetime.timedelta(seconds=6) + self.goodput_recorder.record_step_start_time(0, s0) + s1 = s0 + datetime.timedelta(seconds=2) + self.goodput_recorder.record_step_start_time(1, s1) + + # Elastic reinit at Step 1: 15s total, with 2 workers logging a 6s restore. + reinit_start = s1 + datetime.timedelta(seconds=2) + reinit_end = reinit_start + datetime.timedelta(seconds=15) + self.goodput_recorder.record_elastic_reinit_start_time(reinit_start) + for _ in range(2): + self.mock_cloud_logger.write_cloud_logging_entry({ + 'event_type': 'restore', + 'step': 1, + 'directory': 'gs://bucket/path', + 'checkpoint_manager_start_time': reinit_start.timestamp() + 2.0, + 'checkpoint_manager_duration_secs': 6.0, + 'checkpointer_duration_secs': 6.0, + }) + self.goodput_recorder.record_elastic_reinit_end_time(reinit_end) + + s2 = reinit_end + self.goodput_recorder.record_step_start_time(2, s2) + job_end = s2 + datetime.timedelta(seconds=2) + self.goodput_recorder.record_job_end_time(job_end) + + self.goodput_calculator._fetch_new_entries(job_end) + _, unprod, _, _ = ( + self.goodput_calculator._get_current_productive_and_unproductive_time() + ) + # Initial 5s restore outside reinit is preserved; the 6s restore inside + # reinit (logged by 2 workers) is averaged to 6s and discounted once. + self.assertAlmostEqual( + unprod[goodput_utils.BadputType.UNPRODUCTIVE_CHECKPOINT_RESTORE_TIME], + 5.0, + delta=0.1, + ) + if __name__ == '__main__': googletest.main() diff --git a/ml_goodput_measurement/tests/goodput_test.py b/ml_goodput_measurement/tests/goodput_test.py index 6932f48..8bcead7 100644 --- a/ml_goodput_measurement/tests/goodput_test.py +++ b/ml_goodput_measurement/tests/goodput_test.py @@ -2767,5 +2767,239 @@ def test_exclusion_masking_training_prep(self): breakdown[BadputType.TRAINING_PREP], 25, delta=0.1 ) + def test_checkpoint_badput_skipped_for_uncompleted_step(self): + """Validates checkpoint save badput does not cause negative productive time if step did not complete.""" + job_start_time = datetime.datetime( + 2026, 8, 27, 15, 50, 0, tzinfo=datetime.timezone.utc + ) + self.goodput_recorder.record_job_start_time(job_start_time) + step_start = job_start_time + datetime.timedelta(seconds=10) + self.goodput_recorder.record_step_start_time(0, step_start) + + save_stats = MockSaveStepStatistics( + step=0, + event_type='save', + directory='gs://bucket/path', + checkpoint_manager_blocking_start_time=step_start.timestamp() + 5.0, + checkpoint_manager_blocking_duration_secs=3.0, + synchronous=True, + ) + self.mock_cloud_logger.write_cloud_logging_entry(asdict(save_stats)) + + goodput_pct, breakdown, max_step = self.goodput_calculator.get_job_goodput( + include_badput_breakdown=True + ) + self.assertEqual(goodput_pct, 0.0) + self.assertEqual(max_step, 0) + self.assertGreater( + breakdown.get(BadputType.UNPRODUCTIVE_CHECKPOINT_SAVE_TIME, 0.0), 0.0 + ) + + def test_checkpoint_save_on_compilation_step_not_double_subtracted(self): + """Validates checkpoint save on compilation step is deducted before startup_extra_time.""" + job_start_time = datetime.datetime( + 2026, 8, 27, 16, 0, 0, tzinfo=datetime.timezone.utc + ) + self.goodput_recorder.record_job_start_time(job_start_time) + + # Step 0: 3s execution + 10s compilation + 4s save = 17s total delta + s0 = job_start_time + self.goodput_recorder.record_step_start_time(0, s0) + save_stats = MockSaveStepStatistics( + step=0, + event_type='save', + directory='gs://bucket/path', + checkpoint_manager_blocking_start_time=s0.timestamp() + 13.0, + checkpoint_manager_blocking_duration_secs=4.0, + synchronous=True, + ) + self.mock_cloud_logger.write_cloud_logging_entry(asdict(save_stats)) + + # Step 1: 3s execution + s1 = s0 + datetime.timedelta(seconds=17) + self.goodput_recorder.record_step_start_time(1, s1) + + # Step 2: 3s execution + s2 = s1 + datetime.timedelta(seconds=3) + self.goodput_recorder.record_step_start_time(2, s2) + job_end = s2 + datetime.timedelta(seconds=3) + self.goodput_recorder.record_job_end_time(job_end) + + # Total time = 23s. Productive = 3 steps * 3s = 9s. + # Startup = 10s. Checkpoint save = 4s. + goodput_pct, breakdown, _ = self.goodput_calculator.get_job_goodput( + include_badput_breakdown=True + ) + self.assertAlmostEqual(goodput_pct, (9.0 / 23.0) * 100, delta=0.1) + self.assertAlmostEqual( + breakdown[BadputType.PROGRAM_STARTUP], (10.0 / 23.0) * 100, delta=0.1 + ) + self.assertAlmostEqual( + breakdown[BadputType.UNPRODUCTIVE_CHECKPOINT_SAVE_TIME], + (4.0 / 23.0) * 100, + delta=0.1, + ) + + def test_get_timestamp_from_log_entry_ignores_durations_and_nones(self): + """Validates get_timestamp_from_log_entry ignores duration fields and None values.""" + expected_ts = 1787865289.9484715 + entry = { + 'time_between_consecutive_saves_sec': 1008.38, + 'preemption_received_at': None, + 'checkpoint_manager_blocking_duration_secs': 2.86, + 'checkpoint_manager_blocking_start_time': expected_ts, + 'get_old_steps_start_time': expected_ts + 2.86, + 'synchronous': True, + } + ts = get_timestamp_from_log_entry(entry) + self.assertIsNotNone(ts) + self.assertAlmostEqual(ts.timestamp(), expected_ts, places=3) + + def test_checkpoint_save_multi_worker_averaging_and_local_plus_persistent( + self, + ): + """Validates multi-worker save averaging and local+persistent saves.""" + job_start_time = datetime.datetime( + 2026, 8, 27, 17, 0, 0, tzinfo=datetime.timezone.utc + ) + self.goodput_recorder.record_job_start_time(job_start_time) + + # Step 0: 4s compute + 4s avg persistent save (3s and 5s across 2 workers) + # + 2s local save = 10s total delta. + s0 = job_start_time + self.goodput_recorder.record_step_start_time(0, s0) + for worker_save_dur in (3.0, 5.0): + self.mock_cloud_logger.write_cloud_logging_entry( + asdict( + MockSaveStepStatistics( + step=0, + event_type='save', + directory='gs://bucket/persistent', + checkpoint_manager_blocking_start_time=s0.timestamp() + 4.0, + checkpoint_manager_blocking_duration_secs=worker_save_dur, + synchronous=True, + ) + ) + ) + self.mock_cloud_logger.write_cloud_logging_entry( + asdict( + MockSaveStepStatistics( + step=0, + event_type='save', + directory='/local/ramdisk/ckpt', + checkpoint_manager_blocking_start_time=s0.timestamp() + 8.0, + checkpoint_manager_blocking_duration_secs=2.0, + synchronous=True, + ) + ) + ) + + # Step 1: 4s compute + 1.5s local save + 2.5s persistent save (with + # checkpoint_manager_blocking_start_time=None to exercise + # _compute_remaining_checkpoint_save_badput). + s1 = s0 + datetime.timedelta(seconds=10) + self.goodput_recorder.record_step_start_time(1, s1) + self.mock_cloud_logger.write_cloud_logging_entry( + asdict( + MockSaveStepStatistics( + step=1, + event_type='save', + directory='/local/ramdisk/ckpt', + checkpointer_blocking_start_time=s1.timestamp() + 1.0, + checkpoint_manager_blocking_duration_secs=1.5, + synchronous=True, + ) + ) + ) + self.mock_cloud_logger.write_cloud_logging_entry( + asdict( + MockSaveStepStatistics( + step=1, + event_type='save', + directory='gs://bucket/persistent', + checkpointer_blocking_start_time=s1.timestamp() + 2.0, + checkpoint_manager_blocking_duration_secs=2.5, + synchronous=True, + ) + ) + ) + + # Step 2: 4s compute. + s2 = s1 + datetime.timedelta(seconds=8) + self.goodput_recorder.record_step_start_time(2, s2) + job_end = s2 + datetime.timedelta(seconds=4) + self.goodput_recorder.record_job_end_time(job_end) + + # Total job time = 22s. Productive time = 3 steps * 4s = 12s. + # Checkpoint save badput = (4s + 2s on Step 0) + (1.5s + 2.5s on Step 1). + goodput_pct, breakdown, _ = self.goodput_calculator.get_job_goodput( + include_badput_breakdown=True + ) + self.assertAlmostEqual(goodput_pct, (12.0 / 22.0) * 100, delta=0.1) + self.assertAlmostEqual( + breakdown[BadputType.UNPRODUCTIVE_CHECKPOINT_SAVE_TIME], + (10.0 / 22.0) * 100, + delta=0.1, + ) + + def test_multi_segment_rollback_to_earlier_step_than_previous_segment(self): + """Validates rollback to a step earlier than the disrupted segment start.""" + t0 = datetime.datetime(2026, 8, 27, 18, 0, 0, tzinfo=datetime.timezone.utc) + self.goodput_recorder.record_job_start_time(t0) + + # Run 1 executes steps 0, 1, 2 (5s each), disrupted at step 3 (t=15s). + self.goodput_recorder.record_step_start_time(0, t0) + self.goodput_recorder.record_step_start_time( + 1, t0 + datetime.timedelta(seconds=5) + ) + self.goodput_recorder.record_step_start_time( + 2, t0 + datetime.timedelta(seconds=10) + ) + self.goodput_recorder.record_step_start_time( + 3, t0 + datetime.timedelta(seconds=15) + ) + + # Run 2 restarts at t=20s from local checkpoint at step 2, runs 2->3 (5s). + t_r2 = t0 + datetime.timedelta(seconds=20) + self.goodput_recorder.record_job_start_time(t_r2) + self.goodput_recorder.record_step_start_time(2, t_r2) + self.goodput_recorder.record_step_start_time( + 3, t_r2 + datetime.timedelta(seconds=5) + ) + + # Run 3 restarts at t=30s from older persistent checkpoint at step 1 + # (curr_step=1 < Run 2 min_step=2) and runs steps 1, 2, 3 to completion. + t_r3 = t0 + datetime.timedelta(seconds=30) + self.goodput_recorder.record_job_start_time(t_r3) + self.goodput_recorder.record_step_start_time(1, t_r3) + self.goodput_recorder.record_step_start_time( + 2, t_r3 + datetime.timedelta(seconds=5) + ) + self.goodput_recorder.record_step_start_time( + 3, t_r3 + datetime.timedelta(seconds=10) + ) + job_end = t_r3 + datetime.timedelta(seconds=15) + self.goodput_recorder.record_job_end_time(job_end) + + # Total time = 45s. Productive = 4 steps * 5s = 20s. + # Wasted progress = Run 1 steps 1,2 (10s) + Run 2 step 2 (5s) = 15s. + # Infra recovery = (20-15) + (30-25) = 10s. + goodput_pct, breakdown, max_step = self.goodput_calculator.get_job_goodput( + include_badput_breakdown=True + ) + self.assertEqual(max_step, 3) + self.assertAlmostEqual(goodput_pct, (20.0 / 45.0) * 100, delta=0.1) + self.assertAlmostEqual( + breakdown[BadputType.WASTED_PROGRESS_FROM_DISRUPTION], + (15.0 / 45.0) * 100, + delta=0.1, + ) + self.assertAlmostEqual( + breakdown[BadputType.INFRASTRUCTURE_RECOVERY_FROM_DISRUPTION], + (10.0 / 45.0) * 100, + delta=0.1, + ) + + if __name__ == '__main__': googletest.main()