CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
test_trainer_callback.py246 linesDownload Raw Back to trainer
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