chendl/compositional_test
1
1# Copyright 2020 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15import shutil16import tempfile17import unittest18from unittest.mock import patch19 20from transformers import (21 DefaultFlowCallback,22 IntervalStrategy,23 PrinterCallback,24 ProgressCallback,25 Trainer,26 TrainerCallback,27 TrainingArguments,28 is_torch_available,29)30from transformers.testing_utils import require_torch31 32 33if is_torch_available():34 from transformers.trainer import DEFAULT_CALLBACKS35 36 from .test_trainer import RegressionDataset, RegressionModelConfig, RegressionPreTrainedModel37 38 39class MyTestTrainerCallback(TrainerCallback):40 "A callback that registers the events that goes through."41 42 def __init__(self):43 self.events = []44 45 def on_init_end(self, args, state, control, **kwargs):46 self.events.append("on_init_end")47 48 def on_train_begin(self, args, state, control, **kwargs):49 self.events.append("on_train_begin")50 51 def on_train_end(self, args, state, control, **kwargs):52 self.events.append("on_train_end")53 54 def on_epoch_begin(self, args, state, control, **kwargs):55 self.events.append("on_epoch_begin")56 57 def on_epoch_end(self, args, state, control, **kwargs):58 self.events.append("on_epoch_end")59 60 def on_step_begin(self, args, state, control, **kwargs):61 self.events.append("on_step_begin")62 63 def on_step_end(self, args, state, control, **kwargs):64 self.events.append("on_step_end")65 66 def on_evaluate(self, args, state, control, **kwargs):67 self.events.append("on_evaluate")68 69 def on_predict(self, args, state, control, **kwargs):70 self.events.append("on_predict")71 72 def on_save(self, args, state, control, **kwargs):73 self.events.append("on_save")74 75 def on_log(self, args, state, control, **kwargs):76 self.events.append("on_log")77 78 def on_prediction_step(self, args, state, control, **kwargs):79 self.events.append("on_prediction_step")80 81 82@require_torch83class TrainerCallbackTest(unittest.TestCase):84 def setUp(self):85 self.output_dir = tempfile.mkdtemp()86 87 def tearDown(self):88 shutil.rmtree(self.output_dir)89 90 def get_trainer(self, a=0, b=0, train_len=64, eval_len=64, callbacks=None, disable_tqdm=False, **kwargs):91 # disable_tqdm in TrainingArguments has a flaky default since it depends on the level of logging. We make sure92 # its set to False since the tests later on depend on its value.93 train_dataset = RegressionDataset(length=train_len)94 eval_dataset = RegressionDataset(length=eval_len)95 config = RegressionModelConfig(a=a, b=b)96 model = RegressionPreTrainedModel(config)97 98 args = TrainingArguments(self.output_dir, disable_tqdm=disable_tqdm, report_to=[], **kwargs)99 return Trainer(100 model,101 args,102 train_dataset=train_dataset,103 eval_dataset=eval_dataset,104 callbacks=callbacks,105 )106 107 def check_callbacks_equality(self, cbs1, cbs2):108 self.assertEqual(len(cbs1), len(cbs2))109 110 # Order doesn't matter111 cbs1 = sorted(cbs1, key=lambda cb: cb.__name__ if isinstance(cb, type) else cb.__class__.__name__)112 cbs2 = sorted(cbs2, key=lambda cb: cb.__name__ if isinstance(cb, type) else cb.__class__.__name__)113 114 for cb1, cb2 in zip(cbs1, cbs2):115 if isinstance(cb1, type) and isinstance(cb2, type):116 self.assertEqual(cb1, cb2)117 elif isinstance(cb1, type) and not isinstance(cb2, type):118 self.assertEqual(cb1, cb2.__class__)119 elif not isinstance(cb1, type) and isinstance(cb2, type):120 self.assertEqual(cb1.__class__, cb2)121 else:122 self.assertEqual(cb1, cb2)123 124 def get_expected_events(self, trainer):125 expected_events = ["on_init_end", "on_train_begin"]126 step = 0127 train_dl_len = len(trainer.get_eval_dataloader())128 evaluation_events = ["on_prediction_step"] * len(trainer.get_eval_dataloader()) + ["on_log", "on_evaluate"]129 for _ in range(trainer.state.num_train_epochs):130 expected_events.append("on_epoch_begin")131 for _ in range(train_dl_len):132 step += 1133 expected_events += ["on_step_begin", "on_step_end"]134 if step % trainer.args.logging_steps == 0:135 expected_events.append("on_log")136 if trainer.args.evaluation_strategy == IntervalStrategy.STEPS and step % trainer.args.eval_steps == 0:137 expected_events += evaluation_events.copy()138 if step % trainer.args.save_steps == 0:139 expected_events.append("on_save")140 expected_events.append("on_epoch_end")141 if trainer.args.evaluation_strategy == IntervalStrategy.EPOCH:142 expected_events += evaluation_events.copy()143 expected_events += ["on_log", "on_train_end"]144 return expected_events145 146 def test_init_callback(self):147 trainer = self.get_trainer()148 expected_callbacks = DEFAULT_CALLBACKS.copy() + [ProgressCallback]149 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)150 151 # Callbacks passed at init are added to the default callbacks152 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback])153 expected_callbacks.append(MyTestTrainerCallback)154 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)155 156 # TrainingArguments.disable_tqdm controls if use ProgressCallback or PrinterCallback157 trainer = self.get_trainer(disable_tqdm=True)158 expected_callbacks = DEFAULT_CALLBACKS.copy() + [PrinterCallback]159 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)160 161 def test_add_remove_callback(self):162 expected_callbacks = DEFAULT_CALLBACKS.copy() + [ProgressCallback]163 trainer = self.get_trainer()164 165 # We can add, pop, or remove by class name166 trainer.remove_callback(DefaultFlowCallback)167 expected_callbacks.remove(DefaultFlowCallback)168 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)169 170 trainer = self.get_trainer()171 cb = trainer.pop_callback(DefaultFlowCallback)172 self.assertEqual(cb.__class__, DefaultFlowCallback)173 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)174 175 trainer.add_callback(DefaultFlowCallback)176 expected_callbacks.insert(0, DefaultFlowCallback)177 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)178 179 # We can also add, pop, or remove by instance180 trainer = self.get_trainer()181 cb = trainer.callback_handler.callbacks[0]182 trainer.remove_callback(cb)183 expected_callbacks.remove(DefaultFlowCallback)184 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)185 186 trainer = self.get_trainer()187 cb1 = trainer.callback_handler.callbacks[0]188 cb2 = trainer.pop_callback(cb1)189 self.assertEqual(cb1, cb2)190 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)191 192 trainer.add_callback(cb1)193 expected_callbacks.insert(0, DefaultFlowCallback)194 self.check_callbacks_equality(trainer.callback_handler.callbacks, expected_callbacks)195 196 def test_event_flow(self):197 import warnings198 199 # XXX: for now ignore scatter_gather warnings in this test since it's not relevant to what's being tested200 warnings.simplefilter(action="ignore", category=UserWarning)201 202 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback])203 trainer.train()204 events = trainer.callback_handler.callbacks[-2].events205 self.assertEqual(events, self.get_expected_events(trainer))206 207 # Independent log/save/eval208 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback], logging_steps=5)209 trainer.train()210 events = trainer.callback_handler.callbacks[-2].events211 self.assertEqual(events, self.get_expected_events(trainer))212 213 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback], save_steps=5)214 trainer.train()215 events = trainer.callback_handler.callbacks[-2].events216 self.assertEqual(events, self.get_expected_events(trainer))217 218 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback], eval_steps=5, evaluation_strategy="steps")219 trainer.train()220 events = trainer.callback_handler.callbacks[-2].events221 self.assertEqual(events, self.get_expected_events(trainer))222 223 trainer = self.get_trainer(callbacks=[MyTestTrainerCallback], evaluation_strategy="epoch")224 trainer.train()225 events = trainer.callback_handler.callbacks[-2].events226 self.assertEqual(events, self.get_expected_events(trainer))227 228 # A bit of everything229 trainer = self.get_trainer(230 callbacks=[MyTestTrainerCallback],231 logging_steps=3,232 save_steps=10,233 eval_steps=5,234 evaluation_strategy="steps",235 )236 trainer.train()237 events = trainer.callback_handler.callbacks[-2].events238 self.assertEqual(events, self.get_expected_events(trainer))239 240 # warning should be emitted for duplicated callbacks241 with patch("transformers.trainer_callback.logger.warning") as warn_mock:242 trainer = self.get_trainer(243 callbacks=[MyTestTrainerCallback, MyTestTrainerCallback],244 )245 assert str(MyTestTrainerCallback) in warn_mock.call_args[0][0]246 