@@ -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 )
0 commit comments