Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 58 additions & 25 deletions bert_squeeze/assistants/distil_assistant.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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):
"""
Expand Down Expand Up @@ -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,
Expand All @@ -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"]
Expand Down
36 changes: 16 additions & 20 deletions bert_squeeze/assistants/train_assistant.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
import logging
from copy import deepcopy
from importlib import resources
from typing import Dict, List, Optional

Expand Down Expand Up @@ -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"]
Expand Down
41 changes: 41 additions & 0 deletions tests/assistants/test_assistant_defaults.py
Original file line number Diff line number Diff line change
@@ -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"
Loading