diff --git a/bert_squeeze/assistants/distil_assistant.py b/bert_squeeze/assistants/distil_assistant.py index 28c59f3..b0b0a0e 100644 --- a/bert_squeeze/assistants/distil_assistant.py +++ b/bert_squeeze/assistants/distil_assistant.py @@ -1,6 +1,7 @@ import logging +from copy import deepcopy from importlib import resources -from typing import Dict, List, Optional, Union +from typing import Dict, List, Optional, Union, cast import lightning.pytorch as pl import torch.nn @@ -19,6 +20,16 @@ "distil-seq2seq": "distil_seq2seq.yaml", } +DATA_SECTION_KEYS = { + "_target_", + "teacher_module", + "student_module", + "soft_data_config", + "hard_labeler", + "train_batch_size", + "eval_batch_size", +} + class DistilAssistant(object): """ @@ -77,7 +88,7 @@ def __init__( general_kwargs: Optional[Dict[str, object]] = None, train_kwargs: Optional[Dict[str, object]] = None, student_kwargs: Optional[Dict[str, object]] = None, - teacher_kwargs: Dict[str, object] = {}, + teacher_kwargs: Optional[Dict[str, object]] = None, data_kwargs: Optional[Dict[str, object]] = None, logger_kwargs: Optional[Dict[str, object]] = None, callbacks: Optional[List[Callback]] = None, @@ -95,41 +106,63 @@ def __init__( ) with resources.as_file(config_path) as resolved_path: conf = OmegaConf.load(resolved_path) - self._teacher_checkpoint = teacher_kwargs.pop("checkpoint_path", None) + teacher_overrides = deepcopy(teacher_kwargs) if teacher_kwargs is not None else {} + data_overrides = deepcopy(data_kwargs) if data_kwargs is not None else None + self._teacher_checkpoint = teacher_overrides.pop("checkpoint_path", None) - for name in ["teacher_module", "student_module"]: - conf["data"][name]["dataset_config"] = deep_update( - conf["data"][name]["dataset_config"], data_kwargs + if data_overrides is not None: + shared_dataset_overrides = cast( + Dict[str, object], data_overrides.pop("dataset_config", {}) + ) + shared_dataset_overrides.update( + { + key: value + for key, value in data_overrides.items() + if key not in DATA_SECTION_KEYS + } ) - for name, kws in zip( + for module_name in ["teacher_module", "student_module"]: + conf["data"][module_name]["dataset_config"] = deep_update( + conf["data"][module_name]["dataset_config"], + shared_dataset_overrides, + ) + + for section_name, overrides in zip( ["general", "train", "data", "logger", "callbacks"], - [general_kwargs, train_kwargs, data_kwargs, logger_kwargs, callbacks], + [general_kwargs, train_kwargs, data_overrides, logger_kwargs, callbacks], ): - if kws is not None: - base = conf.get(name) + if overrides is not None: + base = conf.get(section_name) if base is None: - conf[name] = kws + conf[section_name] = overrides continue - if "_target_" in kws and kws["_target_"] != conf[name]["_target_"]: - del conf[name] - conf[name] = kws - elif name == "data": - for module in ["teacher_module", "student_module"]: + if ( + isinstance(overrides, dict) + and "_target_" in overrides + and overrides["_target_"] != conf[section_name]["_target_"] + ): + del conf[section_name] + conf[section_name] = overrides + elif section_name == "data": + for module_name in ["teacher_module", "student_module"]: if ( - module in kws - and "_target_" in kws[module] - and conf[name][module]["_target_"] != kws[module]["_target_"] + module_name in overrides + and "_target_" in overrides[module_name] + and conf[section_name][module_name]["_target_"] + != overrides[module_name]["_target_"] ): - del conf[name][module] - conf[name][module] = kws[module] + del conf[section_name][module_name] + conf[section_name][module_name] = overrides[module_name] - conf[name] = deep_update(conf[name], kws) + conf[section_name] = deep_update(conf[section_name], overrides) - for name, kws in zip(["teacher", "student"], [teacher_kwargs, student_kwargs]): - if kws is not None: - conf["model"][name] = deep_update(conf["model"][name], kws) + for role, overrides in zip( + ["teacher", "student"], [teacher_overrides, student_kwargs] + ): + if overrides is not None: + conf["model"][role] = deep_update(conf["model"][role], overrides) self.name = name self.general = conf["general"] diff --git a/bert_squeeze/assistants/train_assistant.py b/bert_squeeze/assistants/train_assistant.py index be4d951..d1a4f74 100644 --- a/bert_squeeze/assistants/train_assistant.py +++ b/bert_squeeze/assistants/train_assistant.py @@ -1,4 +1,4 @@ -import logging +from copy import deepcopy from importlib import resources from typing import Dict, List, Optional @@ -76,34 +76,30 @@ def __init__( ) with resources.as_file(config_path) as resolved_path: conf = OmegaConf.load(resolved_path) - if ( - data_kwargs is not None - and data_kwargs.get("dataset_config", {}).get("path") is not None - ): - logging.warning( - "Found value for `dataset_config.path` which conflicts with parameter" - " `dataset_path`, usingvalue from the later." - ) - - conf["data"]["dataset_config"] = deep_update( - conf["data"]["dataset_config"], data_kwargs["dataset_config"] - ) - del data_kwargs["dataset_config"] - - for name, kws in zip( + data_overrides = deepcopy(data_kwargs) if data_kwargs is not None else None + if data_overrides is not None: + dataset_overrides = data_overrides.pop("dataset_config", None) + if dataset_overrides is not None: + conf["data"]["dataset_config"] = deep_update( + conf["data"]["dataset_config"], dataset_overrides + ) + + for section_name, overrides in zip( ["general", "train", "model", "data", "logger", "callbacks"], [ general_kwargs, train_kwargs, model_kwargs, - data_kwargs, + data_overrides, logger_kwargs, callbacks, ], ): - if kws is not None: - base = conf.get(name) - conf[name] = kws if base is None else deep_update(base, kws) + if overrides is not None: + base = conf.get(section_name) + conf[section_name] = ( + overrides if base is None else deep_update(base, overrides) + ) self.name = name self.general = conf["general"] diff --git a/tests/assistants/test_assistant_defaults.py b/tests/assistants/test_assistant_defaults.py new file mode 100644 index 0000000..7c4ad49 --- /dev/null +++ b/tests/assistants/test_assistant_defaults.py @@ -0,0 +1,41 @@ +from bert_squeeze.assistants import DistilAssistant, TrainAssistant + + +def test_train_assistant_uses_default_data_config_and_keeps_name(): + assistant = TrainAssistant("lr") + + assert assistant.name == "lr" + assert str(assistant) == "TrainAssistant_lr" + + +def test_train_assistant_does_not_mutate_data_overrides(): + data_kwargs = {"dataset_config": {"path": "custom-dataset"}} + + assistant = TrainAssistant("lr", data_kwargs=data_kwargs) + + assert data_kwargs == {"dataset_config": {"path": "custom-dataset"}} + assert assistant._data_conf.dataset_config.path == "custom-dataset" + + +def test_distil_assistant_uses_default_data_config_and_keeps_name(): + assistant = DistilAssistant("distil") + + assert assistant.name == "distil" + assert str(assistant) == "DistilAssistant_distil" + + +def test_distil_assistant_does_not_mutate_overrides(): + teacher_kwargs = {"checkpoint_path": "teacher.ckpt"} + data_kwargs = {"path": "custom-dataset"} + + assistant = DistilAssistant( + "distil", + teacher_kwargs=teacher_kwargs, + data_kwargs=data_kwargs, + ) + + assert teacher_kwargs == {"checkpoint_path": "teacher.ckpt"} + assert data_kwargs == {"path": "custom-dataset"} + assert assistant._teacher_checkpoint == "teacher.ckpt" + assert assistant._data_conf.teacher_module.dataset_config.path == "custom-dataset" + assert assistant._data_conf.student_module.dataset_config.path == "custom-dataset"