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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions src/megatron/bridge/training/utils/train_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -862,6 +862,16 @@ def report_throughput(
elapsed_samples = int(history_samples[-1]) - int(history_samples[0])
elapsed_tokens = int(history_tokens[-1]) - int(history_tokens[0])
elapsed_wct = history_wct[-1] - history_wct[0]

# Skip throughput calculation if elapsed_wct is zero or negative
# This can happen during checkpoint resumption when history_wct is reinitialized
# and the first few iterations are very fast or have identical timestamps
if elapsed_wct <= 0:
print_rank_0(
f"Warning: elapsed_wct is {elapsed_wct}, skipping throughput calculation at iteration {iteration}"
)
return {}

batches_per_sec = elapsed_batches / elapsed_wct
samples_per_sec = elapsed_samples / elapsed_wct
dev_batches_per_sec = batches_per_sec / world_size
Expand Down
68 changes: 68 additions & 0 deletions tests/unit_tests/training/utils/test_train_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1384,6 +1384,74 @@ def test_report_throughput(self):
assert throughput_report["throughput/micro_batch_size"] == 4
assert throughput_report["throughput/device/samples_per_sec"] == 51.2

@mock.patch("megatron.bridge.training.utils.train_utils.print_rank_0")
def test_report_throughput_zero_elapsed_wct(self, mock_print_rank_0):
"""Test throughput metrics when elapsed_wct is zero (identical timestamps).

This can happen during checkpoint resumption when history_wct is reinitialized
and the first few iterations have identical timestamps.
"""
global_batch_size = 64
micro_batch_size = 4
iteration = 100
seq_length = 4096
# All timestamps are identical - elapsed_wct will be 0
history_wct = [1.5, 1.5, 1.5, 1.5, 1.5]
window_size = len(history_wct)
train_config = MockTrainConfig(global_batch_size=global_batch_size, micro_batch_size=micro_batch_size)

throughput_report = report_throughput(
train_config=train_config,
iteration=iteration,
seq_length=seq_length,
history_wct=history_wct,
window_size=window_size,
)

# Should return empty dict when elapsed_wct is 0
assert throughput_report == {}

# Verify warning was printed
mock_print_rank_0.assert_called_once()
warning_message = mock_print_rank_0.call_args[0][0]
assert "Warning: elapsed_wct is 0" in warning_message
assert "skipping throughput calculation" in warning_message
assert f"iteration {iteration}" in warning_message

@mock.patch("megatron.bridge.training.utils.train_utils.print_rank_0")
def test_report_throughput_negative_elapsed_wct(self, mock_print_rank_0):
"""Test throughput metrics when elapsed_wct is negative.

This shouldn't happen in normal operation, but the code guards against it
to prevent division by zero or negative throughput values.
"""
global_batch_size = 64
micro_batch_size = 4
iteration = 100
seq_length = 4096
# Timestamps go backwards - elapsed_wct will be negative
history_wct = [5.9, 4.2, 2.9, 1.7, 0.9]
window_size = len(history_wct)
train_config = MockTrainConfig(global_batch_size=global_batch_size, micro_batch_size=micro_batch_size)

throughput_report = report_throughput(
train_config=train_config,
iteration=iteration,
seq_length=seq_length,
history_wct=history_wct,
window_size=window_size,
)

# Should return empty dict when elapsed_wct is negative
assert throughput_report == {}

# Verify warning was printed with negative value
mock_print_rank_0.assert_called_once()
warning_message = mock_print_rank_0.call_args[0][0]
assert "Warning: elapsed_wct is -5.0" in warning_message
assert "skipping throughput calculation" in warning_message
assert f"iteration {iteration}" in warning_message

def test_l2_norm_grad(self):
"""Test l2 norm grad metrics."""
num_chunks = 10
Expand Down