Skip to content

Commit 1bd0f3e

Browse files
galrotemmeta-codesync[bot]
authored andcommitted
Surface data-wait-inclusive iteration time in IterationTimeLogger (#1069)
Summary: Pull Request resolved: #1069 Reviewed By: cniii Differential Revision: D111447685 fbshipit-source-id: fa4bf20d49a29cae8ee50f1cb89edf6134671fe0
1 parent a3541c1 commit 1bd0f3e

2 files changed

Lines changed: 108 additions & 37 deletions

File tree

tests/framework/callbacks/test_iteration_time_logger.py

Lines changed: 79 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -82,6 +82,53 @@ def test_iteration_time_logger_test_on_train_step_end(self) -> None:
8282
]
8383
)
8484

85+
def test_logs_and_data_iteration_time(self) -> None:
86+
"""
87+
Test the data-wait-inclusive "*_and_data_iteration_time" metrics are
88+
logged alongside the step-only metrics for train/eval/predict.
89+
"""
90+
logger = MagicMock(spec=MetricLogger)
91+
state = MagicMock(spec=State)
92+
recorded_durations = {
93+
"train_iteration_time": [1, 3, 5, 7, 9],
94+
"train_and_data_iteration_time": [2, 4, 6, 8, 10],
95+
"eval_iteration_time": [1, 3, 5, 7, 9],
96+
"eval_and_data_iteration_time": [2, 4, 6, 8, 10],
97+
"predict_iteration_time": [1, 3, 5, 7, 9],
98+
"predict_and_data_iteration_time": [2, 4, 6, 8, 10],
99+
}
100+
state.train_state.iteration_timer.recorded_durations = recorded_durations.copy()
101+
state.eval_state.iteration_timer.recorded_durations = recorded_durations.copy()
102+
state.predict_state.iteration_timer.recorded_durations = (
103+
recorded_durations.copy()
104+
)
105+
106+
callback = IterationTimeLogger(logger=logger, moving_avg_window=4)
107+
108+
train_unit = DummyTrainUnit(input_dim=2)
109+
train_unit.train_progress.increment_step()
110+
train_unit.train_progress.increment_step()
111+
eval_unit = DummyEvalUnit(input_dim=2)
112+
eval_unit.eval_progress.increment_step()
113+
eval_unit.eval_progress.increment_step()
114+
predict_unit = DummyPredictUnit(input_dim=2)
115+
predict_unit.predict_progress.increment_step()
116+
predict_unit.predict_progress.increment_step()
117+
118+
callback.on_train_step_end(state, train_unit)
119+
callback.on_eval_step_end(state, eval_unit)
120+
callback.on_predict_step_end(state, predict_unit)
121+
122+
# avg of last 4 of [2,4,6,8,10] is 7.0; reported for step-1 == 1
123+
logger.log.assert_has_calls(
124+
[
125+
call("Train Iteration and Data Wait Time (seconds)", 7.0, 1),
126+
call("Eval Iteration and Data Wait Time (seconds)", 7.0, 1),
127+
call("Prediction Iteration and Data Wait Time (seconds)", 7.0, 1),
128+
],
129+
any_order=True,
130+
)
131+
85132
def test_with_train_epoch(self) -> None:
86133
"""
87134
Test IterationTimeLogger callback with train entry point
@@ -94,8 +141,9 @@ def test_with_train_epoch(self) -> None:
94141
num_samples=12, input_dim=2, batch_size=2
95142
)
96143
train(my_unit, dataloader, max_epochs=2, callbacks=[callback])
97-
# 2 epochs, 6 iterations each, logging every third step
98-
self.assertEqual(logger.log.call_count, 4)
144+
# 2 epochs, 6 iterations each, logging every third step; each logged step
145+
# emits both "train_iteration_time" and "train_and_data_iteration_time"
146+
self.assertEqual(logger.log.call_count, 8)
99147

100148
def test_comparing_step_logging_time(self) -> None:
101149
"""
@@ -131,21 +179,46 @@ def test_comparing_step_logging_time(self) -> None:
131179
train_iteration_timer = none_throws(
132180
state.train_state
133181
).iteration_timer.recorded_durations["train_iteration_time"]
182+
train_and_data_iteration_timer = none_throws(
183+
state.train_state
184+
).iteration_timer.recorded_durations["train_and_data_iteration_time"]
134185
eval_iteration_timer = none_throws(
135186
state.eval_state
136187
).iteration_timer.recorded_durations["eval_iteration_time"]
188+
eval_and_data_iteration_timer = none_throws(
189+
state.eval_state
190+
).iteration_timer.recorded_durations["eval_and_data_iteration_time"]
137191

138192
expected_training_iteration_time_calls = [
139193
call("Train Iteration Time (seconds)", train_iteration_timer[i], i + 1)
140194
for i in range(4)
141195
]
196+
expected_train_and_data_iteration_time_calls = [
197+
call(
198+
"Train Iteration and Data Wait Time (seconds)",
199+
train_and_data_iteration_timer[i],
200+
i + 1,
201+
)
202+
for i in range(4)
203+
]
142204
expected_eval_iteration_time_calls = [
143205
call("Eval Iteration Time (seconds)", eval_iteration_timer[i], i + 1)
144206
for i in range(4)
145207
]
208+
expected_eval_and_data_iteration_time_calls = [
209+
call(
210+
"Eval Iteration and Data Wait Time (seconds)",
211+
eval_and_data_iteration_timer[i],
212+
i + 1,
213+
)
214+
for i in range(4)
215+
]
146216

147217
logger.log.assert_has_calls(
148-
expected_training_iteration_time_calls + expected_eval_iteration_time_calls,
218+
expected_training_iteration_time_calls
219+
+ expected_train_and_data_iteration_time_calls
220+
+ expected_eval_iteration_time_calls
221+
+ expected_eval_and_data_iteration_time_calls,
149222
any_order=True,
150223
)
151224

@@ -161,8 +234,9 @@ def test_with_summary_writer(self) -> None:
161234
num_samples=12, input_dim=2, batch_size=2
162235
)
163236
train(my_unit, dataloader, max_epochs=2, callbacks=[callback])
164-
# 2 epochs, 6 iterations each, logging every third step
165-
self.assertEqual(logger.add_scalar.call_count, 4)
237+
# 2 epochs, 6 iterations each, logging every third step; each logged step
238+
# emits both "train_iteration_time" and "train_and_data_iteration_time"
239+
self.assertEqual(logger.add_scalar.call_count, 8)
166240

167241
def test_warmup_steps(self) -> None:
168242
logger = MagicMock(spec=MetricLogger)

torchtnt/framework/callbacks/iteration_time_logger.py

Lines changed: 29 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -77,8 +77,11 @@ def _log_step_metrics(
7777

7878
human_metric_names = {
7979
"train_iteration_time": "Train Iteration Time (seconds)",
80+
"train_and_data_iteration_time": "Train Iteration and Data Wait Time (seconds)",
8081
"eval_iteration_time": "Eval Iteration Time (seconds)",
82+
"eval_and_data_iteration_time": "Eval Iteration and Data Wait Time (seconds)",
8183
"predict_iteration_time": "Prediction Iteration Time (seconds)",
84+
"predict_and_data_iteration_time": "Prediction Iteration and Data Wait Time (seconds)",
8285
}
8386

8487
time_list = iteration_timer.recorded_durations.get(metric_label, [])
@@ -101,52 +104,46 @@ def _log_step_metrics(
101104

102105
def on_train_step_end(self, state: State, unit: TTrainUnit) -> None:
103106
timer = none_throws(state.train_state).iteration_timer
104-
self._log_step_metrics(
105-
"train_iteration_time",
106-
timer,
107-
# on_train_step_end happens after the num steps is incremented, but before the timer list is populated,
108-
# so it logs for step-1
109-
unit.train_progress.num_steps_completed - 1,
110-
)
107+
# on_train_step_end happens after the num steps is incremented, but before the timer lists are populated,
108+
# so it logs for step-1
109+
step_logging_for = unit.train_progress.num_steps_completed - 1
110+
self._log_step_metrics("train_iteration_time", timer, step_logging_for)
111+
self._log_step_metrics("train_and_data_iteration_time", timer, step_logging_for)
111112

112113
def on_train_end(self, state: State, unit: TTrainUnit) -> None:
113-
self._log_step_metrics(
114-
"train_iteration_time",
115-
none_throws(state.train_state).iteration_timer,
116-
unit.train_progress.num_steps_completed,
117-
)
114+
timer = none_throws(state.train_state).iteration_timer
115+
step_logging_for = unit.train_progress.num_steps_completed
116+
self._log_step_metrics("train_iteration_time", timer, step_logging_for)
117+
self._log_step_metrics("train_and_data_iteration_time", timer, step_logging_for)
118118

119119
def on_eval_step_end(self, state: State, unit: TEvalUnit) -> None:
120120
timer = none_throws(state.eval_state).iteration_timer
121-
self._log_step_metrics(
122-
"eval_iteration_time",
123-
timer,
124-
# on_eval_step_end happens after the num steps is incremented, but before the timer list is populated,
125-
# so it logs for step-1
126-
unit.eval_progress.num_steps_completed - 1,
127-
)
121+
# on_eval_step_end happens after the num steps is incremented, but before the timer lists are populated,
122+
# so it logs for step-1
123+
step_logging_for = unit.eval_progress.num_steps_completed - 1
124+
self._log_step_metrics("eval_iteration_time", timer, step_logging_for)
125+
self._log_step_metrics("eval_and_data_iteration_time", timer, step_logging_for)
128126

129127
def on_eval_end(self, state: State, unit: TEvalUnit) -> None:
130-
self._log_step_metrics(
131-
"eval_iteration_time",
132-
none_throws(state.eval_state).iteration_timer,
133-
unit.eval_progress.num_steps_completed,
134-
)
128+
timer = none_throws(state.eval_state).iteration_timer
129+
step_logging_for = unit.eval_progress.num_steps_completed
130+
self._log_step_metrics("eval_iteration_time", timer, step_logging_for)
131+
self._log_step_metrics("eval_and_data_iteration_time", timer, step_logging_for)
135132

136133
def on_predict_step_end(self, state: State, unit: TPredictUnit) -> None:
137134
timer = none_throws(state.predict_state).iteration_timer
135+
# on_predict_step_end happens after the num steps is incremented, but before the timer lists are populated,
136+
# so it logs for step-1
137+
step_logging_for = unit.predict_progress.num_steps_completed - 1
138+
self._log_step_metrics("predict_iteration_time", timer, step_logging_for)
138139
self._log_step_metrics(
139-
"predict_iteration_time",
140-
timer,
141-
# on_predict_step_end happens after the num steps is incremented, but before the timer list is populated,
142-
# so it logs for step-1
143-
unit.predict_progress.num_steps_completed - 1,
140+
"predict_and_data_iteration_time", timer, step_logging_for
144141
)
145142

146143
def on_predict_end(self, state: State, unit: TPredictUnit) -> None:
147144
timer = none_throws(state.predict_state).iteration_timer
145+
step_logging_for = unit.predict_progress.num_steps_completed
146+
self._log_step_metrics("predict_iteration_time", timer, step_logging_for)
148147
self._log_step_metrics(
149-
"predict_iteration_time",
150-
timer,
151-
unit.predict_progress.num_steps_completed,
148+
"predict_and_data_iteration_time", timer, step_logging_for
152149
)

0 commit comments

Comments
 (0)