From 2cc7f2cbacf94012f1d287b6a77e0045cb94a667 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Thu, 26 Jan 2023 10:12:51 +0000 Subject: [PATCH] add config tests --- src/peft/utils/config.py | 5 +++ tests/test_config.py | 79 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 84 insertions(+) create mode 100644 tests/test_config.py diff --git a/src/peft/utils/config.py b/src/peft/utils/config.py index 41bef20..ef985e3 100644 --- a/src/peft/utils/config.py +++ b/src/peft/utils/config.py @@ -38,10 +38,15 @@ class TaskType(str, enum.Enum): @dataclass class PeftConfigMixin(object): + peft_type: Optional[PeftType] = field(default=None, metadata={"help": "The type of PEFT model."}) + @property def __dict__(self): return asdict(self) + def to_dict(self): + return self.__dict__ + def save_pretrained(self, save_directory): if os.path.isfile(save_directory): raise AssertionError(f"Provided path ({save_directory}) should be a directory, not a file") diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..7d3df71 --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,79 @@ +import unittest +import tempfile +import os + +from peft import LoraConfig, PromptEncoderConfig, PrefixTuningConfig, PromptTuningConfig + +class PeftConfigMixin: + all_config_classes = ( + LoraConfig, + PromptEncoderConfig, + PrefixTuningConfig, + PromptTuningConfig, + ) + + +class PeftConfigTester(unittest.TestCase, PeftConfigMixin): + def test_methods(self): + r""" + Test if all configs have the expected methods. Here we test + - to_dict + - save_pretrained + - from_pretrained + - from_json_file + """ + # test if all configs have the expected methods + for config_class in self.all_config_classes: + config = config_class() + self.assertTrue(hasattr(config, "to_dict")) + self.assertTrue(hasattr(config, "save_pretrained")) + self.assertTrue(hasattr(config, "from_pretrained")) + self.assertTrue(hasattr(config, "from_json_file")) + + + def test_save_pretrained(self): + r""" + Test if the config is correctly saved and loaded using + - save_pretrained + """ + for config_class in self.all_config_classes: + config = config_class() + with tempfile.TemporaryDirectory() as tmp_dirname: + config.save_pretrained(tmp_dirname) + + config_from_pretrained = config_class.from_pretrained(tmp_dirname) + self.assertEqual(config.to_dict(), config_from_pretrained.to_dict()) + + def test_from_json_file(self): + for config_class in self.all_config_classes: + config = config_class() + with tempfile.TemporaryDirectory() as tmp_dirname: + config.save_pretrained(tmp_dirname) + + config_from_json = config_class.from_json_file(os.path.join(tmp_dirname, "adapter_config.json")) + self.assertEqual(config.to_dict(), config_from_json) + + + def test_to_dict(self): + r""" + Test if the config can be correctly converted to a dict using: + - to_dict + - __dict__ + """ + for config_class in self.all_config_classes: + config = config_class() + self.assertEqual(config.to_dict(), config.__dict__) + self.assertTrue(isinstance(config.to_dict(), dict)) + + + def test_set_attributes(self): + # manually set attributes and check if they are correctly written + for config_class in self.all_config_classes: + config = config_class(peft_type="test") + + # save pretrained + with tempfile.TemporaryDirectory() as tmp_dirname: + config.save_pretrained(tmp_dirname) + + config_from_pretrained = config_class.from_pretrained(tmp_dirname) + self.assertEqual(config.to_dict(), config_from_pretrained.to_dict()) \ No newline at end of file