From b497d424bedc0efc5416bec3ecd0402104a40bf6 Mon Sep 17 00:00:00 2001 From: Jules Belveze Date: Sat, 1 Aug 2026 11:24:06 +0200 Subject: [PATCH 1/2] [optimizers] - refactor: make discriminative learning rates model-aware --- bert_squeeze/distillation/base_distiller.py | 127 ++----- bert_squeeze/models/base_lt_module.py | 133 ++----- bert_squeeze/models/lt_berxit.py | 183 +++------- bert_squeeze/models/lt_deebert.py | 148 ++------ bert_squeeze/utils/optimizers/__init__.py | 5 + .../utils/optimizers/parameter_groups.py | 280 +++++++++++++++ bert_squeeze/utils/schedulers/__init__.py | 1 + .../utils/schedulers/reduce_on_plateau.py | 34 ++ tests/test_optimizer_parameter_groups.py | 325 ++++++++++++++++++ 9 files changed, 753 insertions(+), 483 deletions(-) create mode 100644 bert_squeeze/utils/optimizers/parameter_groups.py create mode 100644 bert_squeeze/utils/schedulers/reduce_on_plateau.py create mode 100644 tests/test_optimizer_parameter_groups.py diff --git a/bert_squeeze/distillation/base_distiller.py b/bert_squeeze/distillation/base_distiller.py index b3af2c2..6bc94c4 100644 --- a/bert_squeeze/distillation/base_distiller.py +++ b/bert_squeeze/distillation/base_distiller.py @@ -3,11 +3,16 @@ import lightning.pytorch as pl import numpy as np import torch -from omegaconf import DictConfig, ListConfig -from torch.optim.lr_scheduler import ReduceLROnPlateau +from omegaconf import DictConfig from ..utils.experiment_logging import ExperimentLogger -from ..utils.optimizers import BertAdam +from ..utils.optimizers import ( + BertAdam, + OptimizerParameterGroup, + build_optimizer_parameter_groups, + register_legacy_optimizer_state_migration, +) +from ..utils.schedulers import GroupCompatibleReduceLROnPlateau from ..utils.types import DistillationLoss @@ -52,107 +57,14 @@ def _set_scorers(self) -> None: """""" raise NotImplementedError() - def _get_student_parameters(self) -> List[Dict]: - """ - Method that defines the student's parameters to optimize. - - Returns: - List[Dict]: group of parameters to optimize - """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight', 'layer_norm.weight'] - - if self.params.discriminative_learning: - if ( - isinstance(self.params.learning_rates, ListConfig) - and len(self.params.learning_rates) > 1 - ): - groups = [ - (f'layer.{i}.', self.params.learning_rates[i]) for i in range(12) - ] - else: - lr = ( - self.params.learning_rates[0] - if isinstance(self.params.learning_rates, ListConfig) - else self.params.learning_rates - ) - groups = [ - (f'layer.{i}.', lr * pow(self.params.layer_lr_decay, 11 - i)) - for i in range(12) - ] - - group_all = [f'layer.{i}.' for i in range(12)] - no_decay_optimizer_parameters, decay_optimizer_parameters = [], [] - for g, l in groups: - no_decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.student.named_parameters() - if not any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': self.params.weight_decay, - 'lr': l, - } - ) - decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.student.named_parameters() - if any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': 0.0, - 'lr': l, - } - ) - - group_all_parameters = [ - { - 'params': [ - p - for n, p in self.student.named_parameters() - if not any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': self.params.weight_decay, - }, - { - 'params': [ - p - for n, p in self.student.named_parameters() - if any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': 0.0, - }, - ] - optimizer_grouped_parameters = ( - no_decay_optimizer_parameters - + decay_optimizer_parameters - + group_all_parameters - ) - else: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.student.named_parameters() - if not any(nd in n for nd in no_decay) - ], - 'weight_decay': self.params.weight_decay, - }, - { - 'params': [ - p - for n, p in self.student.named_parameters() - if any(nd in n for nd in no_decay) - ], - 'weight_decay': 0.0, - }, - ] - return optimizer_grouped_parameters + def _get_student_parameters(self) -> List[OptimizerParameterGroup]: + return build_optimizer_parameter_groups( + self.student.named_parameters(), + discriminative_learning=self.params.discriminative_learning, + learning_rates=self.params.learning_rates, + layer_lr_decay=self.params.get("layer_lr_decay", 1.0), + weight_decay=self.params.weight_decay, + ) def configure_optimizers(self) -> Tuple[List, List]: """ @@ -184,8 +96,13 @@ def configure_optimizers(self) -> Tuple[List, List]: else: raise ValueError(f"Optimizer '{self.params.optimizer}' not supported.") + if self.params.discriminative_learning: + register_legacy_optimizer_state_migration( + optimizer, self.student.named_parameters() + ) + if self.params.lr_scheduler: - scheduler = ReduceLROnPlateau(optimizer) + scheduler = GroupCompatibleReduceLROnPlateau(optimizer) lr_scheduler = { 'scheduler': scheduler, 'name': 'NeptuneLogger', diff --git a/bert_squeeze/models/base_lt_module.py b/bert_squeeze/models/base_lt_module.py index ec1988f..011c6d2 100644 --- a/bert_squeeze/models/base_lt_module.py +++ b/bert_squeeze/models/base_lt_module.py @@ -13,7 +13,6 @@ import torch.nn.functional as F from omegaconf import DictConfig, ListConfig from torch.nn import CrossEntropyLoss -from torch.optim.lr_scheduler import ReduceLROnPlateau from transformers import ( AutoConfig, AutoModelForSeq2SeqLM, @@ -22,7 +21,13 @@ from ..utils.experiment_logging import ExperimentLogger from ..utils.losses import LabelSmoothingLoss -from ..utils.optimizers import BertAdam +from ..utils.optimizers import ( + BertAdam, + OptimizerParameterGroup, + build_optimizer_parameter_groups, + register_legacy_optimizer_state_migration, +) +from ..utils.schedulers import GroupCompatibleReduceLROnPlateau from ..utils.scorers import BaseSequenceClassificationScorer, LMScorer, Scorer from ..utils.types import ( FastBertLoss, @@ -31,7 +36,7 @@ ) -class _IdentityParamList(list): +class _IdentityParamList(list[nn.Parameter]): def __contains__(self, item: object) -> bool: return any(param is item for param in self) @@ -130,12 +135,6 @@ def configure_optimizers(self) -> Tuple[List, List]: lr=learning_rate, eps=self.config.adam_eps, ) - - if self.config.lr_scheduler: - scheduler = ReduceLROnPlateau(optimizer) - lr_scheduler = {'scheduler': scheduler, 'name': 'NeptuneLogger'} - return [optimizer], [lr_scheduler] - elif optimizer_name == "bertadam": optimizer = BertAdam( optimizer_parameters, @@ -156,6 +155,14 @@ def configure_optimizers(self) -> Tuple[List, List]: else: raise ValueError(f"Optimizer '{self.config.optimizer}' not supported.") + if self.config.discriminative_learning: + register_legacy_optimizer_state_migration(optimizer, self.named_parameters()) + + if optimizer_name == "adamw" and self.config.lr_scheduler: + scheduler = GroupCompatibleReduceLROnPlateau(optimizer) + lr_scheduler = {'scheduler': scheduler, 'name': 'NeptuneLogger'} + return [optimizer], [lr_scheduler] + return [optimizer], [] def _set_objective(self) -> None: @@ -176,106 +183,14 @@ def _sanity_checks(training_config: DictConfig) -> None: training_config.logging_steps > training_config.accumulation_steps ), "'logging_steps' should be greater than 'accumulation_steps'" - def _get_optimizer_parameters(self) -> List[Dict]: - """ - Method that defines the parameters to optimize. - - Returns: - List[Dict]: group of parameters to optimize - """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight', 'layer_norm.weight'] - - if self.config.discriminative_learning: - if ( - isinstance(self.config.learning_rates, ListConfig) - and len(self.config.learning_rates) > 1 - ): - groups = [ - (f'layer.{i}.', self.config.learning_rates[i]) for i in range(12) - ] - else: - lr = ( - self.config.learning_rates[0] - if isinstance(self.config.learning_rates, ListConfig) - else self.config.learning_rates - ) - groups = [ - (f'layer.{i}.', lr * pow(self.config.layer_lr_decay, 11 - i)) - for i in range(12) - ] - - group_all = [f'layer.{i}.' for i in range(12)] - no_decay_optimizer_parameters, decay_optimizer_parameters = [], [] - for g, l in groups: - no_decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if not any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': self.config.weight_decay, - 'lr': l, - } - ) - decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': 0.0, - 'lr': l, - } - ) - - group_all_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if not any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': 0.0, - }, - ] - optimizer_grouped_parameters = ( - no_decay_optimizer_parameters - + decay_optimizer_parameters - + group_all_parameters - ) - else: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if not any(nd in n for nd in no_decay) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if any(nd in n for nd in no_decay) - ], - 'weight_decay': 0.0, - }, - ] + def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: + optimizer_grouped_parameters = build_optimizer_parameter_groups( + self.named_parameters(), + discriminative_learning=self.config.discriminative_learning, + learning_rates=self.config.learning_rates, + layer_lr_decay=self.config.get("layer_lr_decay", 1.0), + weight_decay=self.config.weight_decay, + ) for group in optimizer_grouped_parameters: group["params"] = _IdentityParamList(list(group["params"])) return optimizer_grouped_parameters diff --git a/bert_squeeze/models/lt_berxit.py b/bert_squeeze/models/lt_berxit.py index 38e7263..9d1978b 100644 --- a/bert_squeeze/models/lt_berxit.py +++ b/bert_squeeze/models/lt_berxit.py @@ -6,11 +6,15 @@ import lightning.pytorch as pl import torch import torch.nn as nn -from omegaconf import DictConfig, ListConfig +from omegaconf import DictConfig from overrides import overrides from torch.nn import CrossEntropyLoss from transformers import AutoConfig +from bert_squeeze.utils.optimizers import ( + OptimizerParameterGroup, + build_optimizer_parameter_groups, +) from bert_squeeze.utils.scorers import Scorer from bert_squeeze.utils.types import RampOutput, SequenceClassificationOutput @@ -173,151 +177,44 @@ def _maybe_switch_stage(self) -> None: self._has_switched_stage = True @overrides - def _get_optimizer_parameters(self) -> List[Dict]: - # Mirror LtDeeBert grouping for backbone stage and provide a gate-only - # variant for the "gates" stage. - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight', 'layer_norm.weight'] - - # Gate-only training stage: optimize only gate parameters + def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: if getattr(self, "train_stage", "backbone") == "gates": - gate_params = [ - (n, p) - for n, p in self.named_parameters() - if "gates" in n and p.requires_grad - ] - optimizer_grouped_parameters = [ - { - 'params': [ - p for n, p in gate_params if not any(nd in n for nd in no_decay) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p for n, p in gate_params if any(nd in n for nd in no_decay) - ], - 'weight_decay': 0.0, - }, - ] - return optimizer_grouped_parameters - - if self.config.discriminative_learning: - if ( - isinstance(self.config.learning_rates, ListConfig) - and len(self.config.learning_rates) > 1 - ): - groups = [ - (f'layer.{i}.', self.config.learning_rates[i]) for i in range(12) - ] - else: - lr = ( - self.config.learning_rates[0] - if isinstance(self.config.learning_rates, ListConfig) - else self.config.learning_rates - ) - groups = [ - (f'layer.{i}.', lr * pow(self.config.layer_lr_decay, 11 - i)) - for i in range(12) - ] - - group_all = [f'layer.{i}.' for i in range(12)] - no_decay_optimizer_parameters, decay_optimizer_parameters = [], [] - for g, l in groups: - no_decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and not any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': self.config.weight_decay, - 'lr': l, - } - ) - decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': 0.0, - 'lr': l, - } - ) - - group_all_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and not any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': 0.0, - }, - ] - optimizer_grouped_parameters = ( - no_decay_optimizer_parameters - + decay_optimizer_parameters - + group_all_parameters + named_parameters = ( + (name, parameter) + for name, parameter in self.named_parameters() + if "gates" in name and parameter.requires_grad + ) + return build_optimizer_parameter_groups( + named_parameters, + discriminative_learning=False, + learning_rates=self.config.learning_rates, + layer_lr_decay=self.config.get("layer_lr_decay", 1.0), + weight_decay=self.config.weight_decay, ) + + discriminative_learning = self.config.discriminative_learning + if discriminative_learning: + named_parameters = self.named_parameters() else: - if self.config.train_highway: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" in n) and (not any(nd in n for nd in no_decay)) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" in n) and (any(nd in n for nd in no_decay)) - ], - 'weight_decay': 0.0, - }, - ] - else: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and (not any(nd in n for nd in no_decay)) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) and (any(nd in n for nd in no_decay)) - ], - 'weight_decay': 0.0, - }, - ] - return optimizer_grouped_parameters + named_parameters = ( + (name, parameter) + for name, parameter in self.named_parameters() + if self._is_active_training_parameter(name) + ) + return build_optimizer_parameter_groups( + named_parameters, + discriminative_learning=discriminative_learning, + learning_rates=self.config.learning_rates, + layer_lr_decay=self.config.get("layer_lr_decay", 1.0), + weight_decay=self.config.weight_decay, + ) + + def _is_active_training_parameter(self, parameter_name: str) -> bool: + if ".ramp." in parameter_name: + return self.config.train_highway + if ".gates." in parameter_name: + return self.train_gates + return not self.config.train_highway @overrides def loss( diff --git a/bert_squeeze/models/lt_deebert.py b/bert_squeeze/models/lt_deebert.py index c359bdb..1d9e64e 100644 --- a/bert_squeeze/models/lt_deebert.py +++ b/bert_squeeze/models/lt_deebert.py @@ -5,11 +5,15 @@ import lightning.pytorch as pl import torch import torch.nn as nn -from omegaconf import DictConfig, ListConfig +from omegaconf import DictConfig from overrides import overrides from torch.nn import CrossEntropyLoss from transformers import AutoConfig +from bert_squeeze.utils.optimizers import ( + OptimizerParameterGroup, + build_optimizer_parameter_groups, +) from bert_squeeze.utils.scorers import Scorer from bert_squeeze.utils.types import RampOutput, SequenceClassificationOutput @@ -147,132 +151,24 @@ def _before_prediction_step(self) -> None: self.bert.set_inference_mode(inference=True) @overrides - def _get_optimizer_parameters(self) -> List[Dict]: - """ - Method that defines the parameter to optimize. - - Returns: - List[Dict]: group of parameters to optimize - """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight', 'layer_norm.weight'] - - if self.config.discriminative_learning: - if ( - isinstance(self.config.learning_rates, ListConfig) - and len(self.config.learning_rates) > 1 - ): - groups = [ - (f'layer.{i}.', self.config.learning_rates[i]) for i in range(12) - ] - else: - lr = ( - self.config.learning_rates[0] - if isinstance(self.config.learning_rates, ListConfig) - else self.config.learning_rates - ) - groups = [ - (f'layer.{i}.', lr * pow(self.config.layer_lr_decay, 11 - i)) - for i in range(12) - ] - - group_all = [f'layer.{i}.' for i in range(12)] - no_decay_optimizer_parameters, decay_optimizer_parameters = [], [] - for g, l in groups: - no_decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and not any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': self.config.weight_decay, - 'lr': l, - } - ) - decay_optimizer_parameters.append( - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and any(nd in n for nd in no_decay) - and any(nd in n for nd in [g]) - ], - 'weight_decay': 0.0, - 'lr': l, - } - ) - - group_all_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and not any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and any(nd in n for nd in no_decay) - and not any(nd in n for nd in group_all) - ], - 'weight_decay': 0.0, - }, - ] - optimizer_grouped_parameters = ( - no_decay_optimizer_parameters - + decay_optimizer_parameters - + group_all_parameters - ) + def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: + discriminative_learning = self.config.discriminative_learning + if discriminative_learning: + named_parameters = self.named_parameters() else: - if self.config.train_highway: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" in n) and (not any(nd in n for nd in no_decay)) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" in n) and (any(nd in n for nd in no_decay)) - ], - 'weight_decay': 0.0, - }, - ] - else: - optimizer_grouped_parameters = [ - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) - and (not any(nd in n for nd in no_decay)) - ], - 'weight_decay': self.config.weight_decay, - }, - { - 'params': [ - p - for n, p in self.named_parameters() - if ("highway" not in n) and (any(nd in n for nd in no_decay)) - ], - 'weight_decay': 0.0, - }, - ] - return optimizer_grouped_parameters + ramp_only = self.config.train_highway + named_parameters = ( + (name, parameter) + for name, parameter in self.named_parameters() + if (".ramp." in name) == ramp_only + ) + return build_optimizer_parameter_groups( + named_parameters, + discriminative_learning=discriminative_learning, + learning_rates=self.config.learning_rates, + layer_lr_decay=self.config.get("layer_lr_decay", 1.0), + weight_decay=self.config.weight_decay, + ) @overrides def loss( diff --git a/bert_squeeze/utils/optimizers/__init__.py b/bert_squeeze/utils/optimizers/__init__.py index 71962b3..abdf7f3 100644 --- a/bert_squeeze/utils/optimizers/__init__.py +++ b/bert_squeeze/utils/optimizers/__init__.py @@ -1 +1,6 @@ from .bert_adam import BertAdam +from .parameter_groups import ( + OptimizerParameterGroup, + build_optimizer_parameter_groups, + register_legacy_optimizer_state_migration, +) diff --git a/bert_squeeze/utils/optimizers/parameter_groups.py b/bert_squeeze/utils/optimizers/parameter_groups.py new file mode 100644 index 0000000..9075090 --- /dev/null +++ b/bert_squeeze/utils/optimizers/parameter_groups.py @@ -0,0 +1,280 @@ +from __future__ import annotations + +import re +from collections.abc import Iterable, Sequence +from typing import Optional, TypedDict, Union, cast + +from torch import nn +from torch.optim import Optimizer +from torch.optim.optimizer import StateDict + +__all__ = [ + "OptimizerParameterGroup", + "build_optimizer_parameter_groups", + "register_legacy_optimizer_state_migration", +] + + +class _RequiredOptimizerParameterGroup(TypedDict): + params: list[nn.Parameter] + weight_decay: float + + +class OptimizerParameterGroup(_RequiredOptimizerParameterGroup, total=False): + lr: float + + +_LAYER_PATTERNS = ( + (re.compile(r"(?:^|\.)block\.(\d+)\."), False), + (re.compile(r"(?:^|\.)layers\.(\d+)\."), False), + (re.compile(r"(?:^|\.)h\.(\d+)\."), False), + (re.compile(r"(?:^|\.)layer\.(\d+)\."), True), +) +_NO_DECAY_NAMES = ("bias", "gamma", "beta", "LayerNorm.weight", "layer_norm.weight") + + +def build_optimizer_parameter_groups( + named_parameters: Iterable[tuple[str, nn.Parameter]], + *, + discriminative_learning: bool, + learning_rates: Union[float, Sequence[float]], + layer_lr_decay: float, + weight_decay: float, +) -> list[OptimizerParameterGroup]: + parameters = list(named_parameters) + if not discriminative_learning: + return _weight_decay_groups(parameters, weight_decay) + + parameters_by_layer, remaining_parameters, uses_legacy_layout = ( + _split_parameters_by_layer(parameters) + ) + layer_indices = sorted(parameters_by_layer) + if not layer_indices: + raise ValueError("No encoder layers found for discriminative learning.") + + layer_rates = _layer_rates(learning_rates, layer_lr_decay, len(layer_indices)) + layer_rate_by_index = dict(zip(layer_indices, layer_rates)) + preserve_legacy_slots = uses_legacy_layout and all( + 0 <= index < 12 for index in layer_indices + ) + group_indices = list(range(12)) if preserve_legacy_slots else layer_indices + groups = [] + for use_weight_decay in (True, False): + for layer_index in group_indices: + layer_parameters = [ + parameter + for name, parameter in parameters_by_layer.get(layer_index, []) + if _uses_weight_decay(name) == use_weight_decay + ] + if layer_parameters or preserve_legacy_slots: + groups.append( + _parameter_group( + layer_parameters, + weight_decay, + use_weight_decay, + layer_rate_by_index.get(layer_index, layer_rates[-1]), + ) + ) + + groups.extend(_weight_decay_groups(remaining_parameters, weight_decay)) + return groups + + +def register_legacy_optimizer_state_migration( + optimizer: Optimizer, + named_parameters: Iterable[tuple[str, nn.Parameter]], +) -> None: + legacy_parameter_groups = _legacy_parameter_groups(list(named_parameters)) + + def migrate_state_dict( + current_optimizer: Optimizer, state_dict: StateDict + ) -> Optional[StateDict]: + saved_groups = cast(list[dict[str, object]], state_dict["param_groups"]) + if len(saved_groups) != len(legacy_parameter_groups): + return None + + saved_metadata_by_parameter = _saved_parameter_metadata( + saved_groups, legacy_parameter_groups + ) + if saved_metadata_by_parameter is None: + return None + + current_groups = cast(list[dict[str, object]], current_optimizer.param_groups) + serialized_groups = cast( + list[dict[str, object]], current_optimizer.state_dict()["param_groups"] + ) + migrated_groups = [] + for current_group, serialized_group in zip(current_groups, serialized_groups): + current_parameters = cast(list[nn.Parameter], current_group["params"]) + if any( + id(parameter) not in saved_metadata_by_parameter + for parameter in current_parameters + ): + return None + source_group_indices = { + saved_metadata_by_parameter[id(parameter)][1] + for parameter in current_parameters + } + if len(source_group_indices) == 1: + source_group_index = next(iter(source_group_indices)) + source_group = saved_groups[source_group_index] + migrated_group = { + key: value for key, value in source_group.items() if key != "params" + } + else: + migrated_group = { + key: value + for key, value in serialized_group.items() + if key != "params" + } + migrated_group["params"] = [ + saved_metadata_by_parameter[id(parameter)][0] + for parameter in current_parameters + ] + migrated_groups.append(migrated_group) + + migrated_state_dict = dict(state_dict) + migrated_state_dict["param_groups"] = migrated_groups + return cast(StateDict, migrated_state_dict) + + optimizer.register_load_state_dict_pre_hook(migrate_state_dict) + + +def _split_parameters_by_layer( + parameters: Sequence[tuple[str, nn.Parameter]], +) -> tuple[ + dict[int, list[tuple[str, nn.Parameter]]], + list[tuple[str, nn.Parameter]], + bool, +]: + parameters_by_layer: dict[int, list[tuple[str, nn.Parameter]]] = {} + remaining_parameters = [] + uses_legacy_layout = True + for name, parameter in parameters: + layer_match = _layer_match(name) + if layer_match is None: + remaining_parameters.append((name, parameter)) + continue + layer_index, is_legacy_layer = layer_match + uses_legacy_layout = uses_legacy_layout and is_legacy_layer + parameters_by_layer.setdefault(layer_index, []).append((name, parameter)) + return parameters_by_layer, remaining_parameters, uses_legacy_layout + + +def _legacy_parameter_groups( + parameters: Sequence[tuple[str, nn.Parameter]], +) -> list[list[nn.Parameter]]: + layer_keys = [f"layer.{index}." for index in range(12)] + groups: list[list[nn.Parameter]] = [] + for use_weight_decay in (True, False): + groups.extend( + [ + parameter + for name, parameter in parameters + if layer_key in name and _uses_weight_decay(name) == use_weight_decay + ] + for layer_key in layer_keys + ) + for use_weight_decay in (True, False): + groups.append( + [ + parameter + for name, parameter in parameters + if not any(layer_key in name for layer_key in layer_keys) + and _uses_weight_decay(name) == use_weight_decay + ] + ) + return groups + + +def _saved_parameter_metadata( + saved_groups: Sequence[dict[str, object]], + legacy_parameter_groups: Sequence[list[nn.Parameter]], +) -> Optional[dict[int, tuple[int, int]]]: + saved_metadata_by_parameter: dict[int, tuple[int, int]] = {} + for group_index, (saved_group, legacy_parameters) in enumerate( + zip(saved_groups, legacy_parameter_groups) + ): + saved_ids = cast(list[int], saved_group["params"]) + if len(saved_ids) != len(legacy_parameters): + return None + saved_metadata_by_parameter.update( + (id(parameter), (saved_id, group_index)) + for parameter, saved_id in zip(legacy_parameters, saved_ids) + ) + return saved_metadata_by_parameter + + +def _layer_match(parameter_name: str) -> Optional[tuple[int, bool]]: + for pattern, is_legacy_layer in _LAYER_PATTERNS: + match = pattern.search(parameter_name) + if match is not None: + return int(match.group(1)), is_legacy_layer + return None + + +def _layer_rates( + learning_rates: Union[float, Sequence[float]], + layer_lr_decay: float, + layer_count: int, +) -> list[float]: + rates = ( + [float(rate) for rate in learning_rates] + if isinstance(learning_rates, Sequence) + else [float(learning_rates)] + ) + if not rates: + raise ValueError("At least one learning rate is required.") + if len(rates) == 1: + return [ + rates[0] * pow(layer_lr_decay, layer_count - index - 1) + for index in range(layer_count) + ] + if len(rates) != layer_count: + raise ValueError( + f"Expected {layer_count} layer learning rates, received {len(rates)}." + ) + return rates + + +def _weight_decay_groups( + named_parameters: Sequence[tuple[str, nn.Parameter]], + weight_decay: float, +) -> list[OptimizerParameterGroup]: + groups = [] + for use_weight_decay in (True, False): + parameters = [ + parameter + for name, parameter in named_parameters + if _uses_weight_decay(name) == use_weight_decay + ] + if not parameters: + continue + groups.append( + _parameter_group( + parameters, + weight_decay, + use_weight_decay, + None, + ) + ) + return groups + + +def _parameter_group( + parameters: list[nn.Parameter], + weight_decay: float, + use_weight_decay: bool, + learning_rate: Optional[float], +) -> OptimizerParameterGroup: + group = OptimizerParameterGroup( + params=parameters, + weight_decay=weight_decay if use_weight_decay else 0.0, + ) + if learning_rate is not None: + group["lr"] = learning_rate + return group + + +def _uses_weight_decay(parameter_name: str) -> bool: + return not any(no_decay in parameter_name for no_decay in _NO_DECAY_NAMES) diff --git a/bert_squeeze/utils/schedulers/__init__.py b/bert_squeeze/utils/schedulers/__init__.py index e69de29..96969c2 100644 --- a/bert_squeeze/utils/schedulers/__init__.py +++ b/bert_squeeze/utils/schedulers/__init__.py @@ -0,0 +1 @@ +from .reduce_on_plateau import GroupCompatibleReduceLROnPlateau diff --git a/bert_squeeze/utils/schedulers/reduce_on_plateau.py b/bert_squeeze/utils/schedulers/reduce_on_plateau.py new file mode 100644 index 0000000..55777fc --- /dev/null +++ b/bert_squeeze/utils/schedulers/reduce_on_plateau.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +from typing import Union, cast + +from overrides import overrides +from torch.optim.lr_scheduler import ReduceLROnPlateau + +__all__ = ["GroupCompatibleReduceLROnPlateau"] + + +class GroupCompatibleReduceLROnPlateau(ReduceLROnPlateau): + @overrides + def load_state_dict(self, state_dict: dict[str, object]) -> None: + migrated_state = dict(state_dict) + migrated_state["min_lrs"] = self._migrated_min_lrs(state_dict.get("min_lrs")) + migrated_state["_last_lr"] = [ + float(group["lr"]) for group in self.optimizer.param_groups + ] + super().load_state_dict(migrated_state) + + def _migrated_min_lrs(self, saved_min_lrs: object) -> list[float]: + current_min_lrs = [float(value) for value in self.min_lrs] + if not isinstance(saved_min_lrs, list) or not all( + isinstance(value, (int, float)) for value in saved_min_lrs + ): + return current_min_lrs + + min_lrs = cast(list[Union[int, float]], saved_min_lrs) + group_count = len(self.optimizer.param_groups) + if len(min_lrs) == group_count: + return [float(value) for value in min_lrs] + if min_lrs and all(value == min_lrs[0] for value in min_lrs): + return [float(min_lrs[0])] * group_count + return current_min_lrs diff --git a/tests/test_optimizer_parameter_groups.py b/tests/test_optimizer_parameter_groups.py new file mode 100644 index 0000000..b11ccd2 --- /dev/null +++ b/tests/test_optimizer_parameter_groups.py @@ -0,0 +1,325 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Optional, Union + +import pytest +import torch +import torch.nn as nn +from omegaconf import DictConfig, OmegaConf +from torch.optim import AdamW +from torch.optim.lr_scheduler import ReduceLROnPlateau +from transformers import BertConfig, T5Config, T5ForConditionalGeneration + +from bert_squeeze.models.custom_transformers.berxit import BerxitModel +from bert_squeeze.models.custom_transformers.deebert import DeeBertModel +from bert_squeeze.models.lt_berxit import LtBerxit +from bert_squeeze.models.lt_deebert import LtDeeBert +from bert_squeeze.utils.optimizers import ( + OptimizerParameterGroup, + build_optimizer_parameter_groups, + register_legacy_optimizer_state_migration, +) +from bert_squeeze.utils.schedulers import GroupCompatibleReduceLROnPlateau + +_NO_DECAY_NAMES = ("bias", "gamma", "beta", "LayerNorm.weight", "layer_norm.weight") + + +class _Encoder(nn.Module): + def __init__(self, layer_count: int) -> None: + super().__init__() + self.layer = nn.ModuleList([nn.Linear(2, 2) for _ in range(layer_count)]) + + +class _LayeredModel(nn.Module): + def __init__(self, layer_count: int) -> None: + super().__init__() + self.encoder = _Encoder(layer_count) + self.classifier = nn.Linear(2, 2) + + +def _learning_rate_for( + groups: list[OptimizerParameterGroup], parameter: nn.Parameter +) -> Optional[float]: + for group in groups: + if any(group_parameter is parameter for group_parameter in group["params"]): + return group.get("lr") + raise AssertionError("Parameter is missing from optimizer groups.") + + +@pytest.mark.parametrize( + ("learning_rates", "expected_rates"), + [ + ([0.1], [0.025, 0.05, 0.1]), + ([0.01, 0.02, 0.03], [0.01, 0.02, 0.03]), + ], +) +def test_optimizer_groups_follow_model_depth( + learning_rates: list[float], expected_rates: list[float] +) -> None: + model = _LayeredModel(layer_count=3) + + groups = build_optimizer_parameter_groups( + model.named_parameters(), + discriminative_learning=True, + learning_rates=learning_rates, + layer_lr_decay=0.5, + weight_decay=0.01, + ) + + actual_rates = [ + _learning_rate_for(groups, layer.weight) for layer in model.encoder.layer + ] + assert actual_rates == pytest.approx(expected_rates) + assert [groups[index]["lr"] for index in range(3)] == pytest.approx(expected_rates) + assert [groups[12 + index]["lr"] for index in range(3)] == pytest.approx( + expected_rates + ) + assert all("lr" not in group for group in groups[24:]) + assert _learning_rate_for(groups, model.classifier.weight) is None + assert sum(len(group["params"]) for group in groups) == len(list(model.parameters())) + + +def test_optimizer_groups_reject_mismatched_layer_rates() -> None: + model = _LayeredModel(layer_count=3) + + with pytest.raises(ValueError, match="Expected 3 layer learning rates"): + build_optimizer_parameter_groups( + model.named_parameters(), + discriminative_learning=True, + learning_rates=[0.1, 0.2], + layer_lr_decay=0.5, + weight_decay=0.01, + ) + + +def _legacy_optimizer_groups(model: nn.Module) -> list[OptimizerParameterGroup]: + named_parameters = list(model.named_parameters()) + layer_keys = [f"layer.{index}." for index in range(12)] + legacy_rates = [0.1 * pow(0.5, 11 - index) for index in range(12)] + legacy_groups: list[OptimizerParameterGroup] = [] + for use_weight_decay in (True, False): + legacy_groups.extend( + OptimizerParameterGroup( + params=[ + parameter + for name, parameter in named_parameters + if layer_key in name + and (not any(no_decay in name for no_decay in _NO_DECAY_NAMES)) + == use_weight_decay + ], + weight_decay=0.01 if use_weight_decay else 0.0, + lr=legacy_rates[index], + ) + for index, layer_key in enumerate(layer_keys) + ) + for use_weight_decay in (True, False): + legacy_groups.append( + OptimizerParameterGroup( + params=[ + parameter + for name, parameter in named_parameters + if not any(layer_key in name for layer_key in layer_keys) + and (not any(no_decay in name for no_decay in _NO_DECAY_NAMES)) + == use_weight_decay + ], + weight_decay=0.01 if use_weight_decay else 0.0, + ) + ) + return legacy_groups + + +def _current_optimizer(model: nn.Module) -> AdamW: + optimizer = AdamW( + build_optimizer_parameter_groups( + model.named_parameters(), + discriminative_learning=True, + learning_rates=[0.1], + layer_lr_decay=0.5, + weight_decay=0.01, + ), + lr=0.1, + ) + register_legacy_optimizer_state_migration(optimizer, model.named_parameters()) + return optimizer + + +@pytest.mark.parametrize("layer_count", [3, 24]) +def test_optimizer_groups_restore_legacy_optimizer_state(layer_count: int) -> None: + model = _LayeredModel(layer_count=layer_count) + legacy_optimizer = AdamW( + _legacy_optimizer_groups(model), + lr=0.1, + betas=(0.8, 0.88), + eps=1e-6, + ) + sum(parameter.square().sum() for parameter in model.parameters()).backward() + legacy_optimizer.step() + legacy_optimizer.zero_grad() + + current_optimizer = _current_optimizer(model) + current_optimizer.load_state_dict(legacy_optimizer.state_dict()) + sum(parameter.square().sum() for parameter in model.parameters()).backward() + current_optimizer.step() + + assert all(torch.isfinite(parameter).all() for parameter in model.parameters()) + assert current_optimizer.param_groups[0]["betas"] == (0.8, 0.88) + assert current_optimizer.param_groups[0]["eps"] == 1e-6 + + +def test_scheduler_restores_after_optimizer_group_migration() -> None: + model = _LayeredModel(layer_count=24) + legacy_optimizer = AdamW(_legacy_optimizer_groups(model), lr=0.1) + legacy_scheduler = ReduceLROnPlateau(legacy_optimizer, factor=0.5, patience=0) + current_optimizer = _current_optimizer(model) + current_optimizer.load_state_dict(legacy_optimizer.state_dict()) + current_scheduler = GroupCompatibleReduceLROnPlateau( + current_optimizer, factor=0.5, patience=0 + ) + + current_scheduler.load_state_dict(legacy_scheduler.state_dict()) + current_scheduler.step(1.0) + current_scheduler.step(2.0) + + assert len(current_scheduler.min_lrs) == len(current_optimizer.param_groups) + + +def _t5_model(block_count: int) -> T5ForConditionalGeneration: + return T5ForConditionalGeneration( + T5Config( + vocab_size=32, + d_model=16, + d_ff=32, + num_layers=block_count, + num_decoder_layers=block_count, + num_heads=2, + ) + ) + + +def test_t5_optimizer_groups_follow_block_depth() -> None: + model = _t5_model(block_count=2) + + groups = build_optimizer_parameter_groups( + model.named_parameters(), + discriminative_learning=True, + learning_rates=[0.1], + layer_lr_decay=0.5, + weight_decay=0.01, + ) + + encoder_rates = [ + _learning_rate_for(groups, block.layer[0].SelfAttention.q.weight) + for block in model.encoder.block + ] + decoder_rates = [ + _learning_rate_for(groups, block.layer[0].SelfAttention.q.weight) + for block in model.decoder.block + ] + grouped_parameters = [parameter for group in groups for parameter in group["params"]] + + assert encoder_rates == pytest.approx([0.05, 0.1]) + assert decoder_rates == pytest.approx([0.05, 0.1]) + assert len(grouped_parameters) == len( + {id(parameter) for parameter in grouped_parameters} + ) + assert len(grouped_parameters) == len(list(model.parameters())) + + +def test_t5_legacy_state_migrates_when_group_counts_match() -> None: + model = _t5_model(block_count=12) + legacy_optimizer = AdamW(_legacy_optimizer_groups(model), lr=0.1) + current_optimizer = _current_optimizer(model) + + assert len(legacy_optimizer.param_groups) == len(current_optimizer.param_groups) + current_optimizer.load_state_dict(legacy_optimizer.state_dict()) + + sum(parameter.square().sum() for parameter in model.parameters()).backward() + current_optimizer.step() + assert all(torch.isfinite(parameter).all() for parameter in model.parameters()) + + +def _model_config(tmp_path: Path) -> BertConfig: + model_config = BertConfig( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + num_attention_heads=2, + num_hidden_layers=2, + num_labels=2, + ) + model_config.save_pretrained(tmp_path) + return model_config + + +def _ramp_training_config(**overrides: object) -> DictConfig: + config = { + "logging_steps": 2, + "accumulation_steps": 1, + "objective": "ce", + "lr_scheduler": False, + "discriminative_learning": True, + "learning_rates": [0.1], + "layer_lr_decay": 0.5, + "weight_decay": 0.01, + "train_highway": True, + "train_gates": True, + "early_exit_entropy": -1.0, + } + config.update(overrides) + return OmegaConf.create(config) + + +def _assert_parameters_are_grouped( + module: Union[LtDeeBert, LtBerxit], parameter_marker: str +) -> None: + groups = module._get_optimizer_parameters() + grouped_parameter_ids = { + id(parameter) for group in groups for parameter in group["params"] + } + expected_parameter_ids = { + id(parameter) + for name, parameter in module.named_parameters() + if parameter_marker in name + } + + assert expected_parameter_ids + assert expected_parameter_ids <= grouped_parameter_ids + + +def test_deebert_shipped_config_includes_ramp_parameters(tmp_path: Path) -> None: + model_config = _model_config(tmp_path) + module = LtDeeBert( + training_config=_ramp_training_config(), + pretrained_model=str(tmp_path), + num_labels=2, + model=DeeBertModel(model_config), + ) + + _assert_parameters_are_grouped(module, ".ramp.") + + +def test_berxit_shipped_config_includes_ramp_parameters(tmp_path: Path) -> None: + model_config = _model_config(tmp_path) + module = LtBerxit( + training_config=_ramp_training_config(), + pretrained_model=str(tmp_path), + num_labels=2, + model=BerxitModel(model_config), + ) + + _assert_parameters_are_grouped(module, ".ramp.") + + +def test_berxit_non_discriminative_training_includes_gate_parameters( + tmp_path: Path, +) -> None: + model_config = _model_config(tmp_path) + module = LtBerxit( + training_config=_ramp_training_config(discriminative_learning=False), + pretrained_model=str(tmp_path), + num_labels=2, + model=BerxitModel(model_config), + ) + + _assert_parameters_are_grouped(module, ".gates.") From f19f4ad99d6bf5c528854460eca621eac35a64a5 Mon Sep 17 00:00:00 2001 From: Jules Belveze Date: Mon, 3 Aug 2026 15:07:56 +0200 Subject: [PATCH 2/2] [bert_squeeze] - refactor: streamline training configurations and logging - Simplified training stage options and clarified comments for better understanding of the configuration. - Improved logging of training loss to enhance monitoring during model training. --- .../assistants/configs/train_berxit.yaml | 14 +- .../assistants/configs/train_deebert.yaml | 3 +- bert_squeeze/distillation/base_distiller.py | 26 +- .../distillation/seq2seq_distiller.py | 7 +- .../sequence_classification_distiller.py | 12 +- bert_squeeze/models/base_lt_module.py | 20 +- .../models/custom_transformers/berxit.py | 222 ++++------ bert_squeeze/models/lt_berxit.py | 153 +++---- bert_squeeze/models/lt_deebert.py | 153 ++----- bert_squeeze/models/lt_t5.py | 6 + bert_squeeze/utils/optimizers/__init__.py | 6 +- .../utils/optimizers/parameter_groups.py | 314 +++++-------- bert_squeeze/utils/schedulers/__init__.py | 1 - .../utils/schedulers/reduce_on_plateau.py | 34 -- bert_squeeze/utils/types.py | 22 +- tests/test_optimizer_parameter_groups.py | 412 +++++++++++------- tests/test_seq2seq_distillation_training.py | 3 +- 17 files changed, 607 insertions(+), 801 deletions(-) delete mode 100644 bert_squeeze/utils/schedulers/reduce_on_plateau.py diff --git a/bert_squeeze/assistants/configs/train_berxit.yaml b/bert_squeeze/assistants/configs/train_berxit.yaml index 0a55ae4..9802dc2 100644 --- a/bert_squeeze/assistants/configs/train_berxit.yaml +++ b/bert_squeeze/assistants/configs/train_berxit.yaml @@ -15,13 +15,9 @@ train: accumulation_steps: 1 auto_lr: false discriminative_learning: true - # Training stage for BERxiT: - # - "backbone": train encoder + ramps + final classifier (no gate loss unless train_gates=true) - # - "gates": freeze backbone/ramps/classifier and train only gates - # You can also keep "backbone" here and use `switch_step` to switch to gate training mid-run. + # "backbone" trains the model; "gates" calibrates only the exit gate train_stage: "backbone" - # Optional global-step at which to switch from backbone to gate training within a single run. - # If null, no automatic switch is performed. + # Set a step to switch from backbone training to gate calibration. switch_step: dropout: 0.2 layer_lr_decay: 0.95 @@ -37,12 +33,12 @@ train: warmup_steps: true weight_decay: 0.01 + # Alternate between final-exit and all-exit objectives. train_highway: true early_exit_entropy: -1 - # BERxiT-specific options + # Train the shared learning-to-exit gate. train_gates: true - gate_hidden_dim: 32 - # Either a single float applied to all layers or a list of floats per layer + # Use one threshold for every layer or provide one value per layer. gate_thresholds: 0.5 model: diff --git a/bert_squeeze/assistants/configs/train_deebert.yaml b/bert_squeeze/assistants/configs/train_deebert.yaml index 651c112..befd798 100644 --- a/bert_squeeze/assistants/configs/train_deebert.yaml +++ b/bert_squeeze/assistants/configs/train_deebert.yaml @@ -29,6 +29,7 @@ train: warmup_steps: true weight_decay: 0.01 + # Set false for the backbone and final exit, or true for frozen-backbone exits. train_highway: true early_exit_entropy: -1 @@ -51,4 +52,4 @@ data: label_col: label truncate_mode: head tokenizer_name: ${model.pretrained_model} - max_length: 256 \ No newline at end of file + max_length: 256 diff --git a/bert_squeeze/distillation/base_distiller.py b/bert_squeeze/distillation/base_distiller.py index 6bc94c4..72bfa08 100644 --- a/bert_squeeze/distillation/base_distiller.py +++ b/bert_squeeze/distillation/base_distiller.py @@ -1,18 +1,17 @@ -from typing import Dict, List, Tuple, Union +from typing import Dict, List, Optional, Tuple, Union import lightning.pytorch as pl import numpy as np import torch from omegaconf import DictConfig +from torch.optim.lr_scheduler import ReduceLROnPlateau from ..utils.experiment_logging import ExperimentLogger from ..utils.optimizers import ( BertAdam, OptimizerParameterGroup, build_optimizer_parameter_groups, - register_legacy_optimizer_state_migration, ) -from ..utils.schedulers import GroupCompatibleReduceLROnPlateau from ..utils.types import DistillationLoss @@ -36,7 +35,7 @@ def __init__( teacher: Union["pl.LightningModule", "torch.nn.Module"], student: Union[pl.LightningModule, torch.nn.Module], training_config: DictConfig, - teacher_checkpoint: str = None, + teacher_checkpoint: Optional[str] = None, **kwargs, ): super().__init__() @@ -96,22 +95,25 @@ def configure_optimizers(self) -> Tuple[List, List]: else: raise ValueError(f"Optimizer '{self.params.optimizer}' not supported.") - if self.params.discriminative_learning: - register_legacy_optimizer_state_migration( - optimizer, self.student.named_parameters() - ) - if self.params.lr_scheduler: - scheduler = GroupCompatibleReduceLROnPlateau(optimizer) + scheduler = ReduceLROnPlateau(optimizer) lr_scheduler = { 'scheduler': scheduler, - 'name': 'NeptuneLogger', - 'monitor': 'loss', + 'name': 'learning_rate', + 'monitor': self.params.get("lr_scheduler_monitor", "train/epoch_loss"), } return [optimizer], [lr_scheduler] return [optimizer], [] + def _log_training_loss(self, loss: torch.Tensor) -> None: + self.log( + "train/epoch_loss", + loss, + on_step=False, + on_epoch=True, + ) + def training_step(self, batch, _) -> torch.Tensor: raise NotImplementedError() diff --git a/bert_squeeze/distillation/seq2seq_distiller.py b/bert_squeeze/distillation/seq2seq_distiller.py index 41f4d06..72b07b7 100644 --- a/bert_squeeze/distillation/seq2seq_distiller.py +++ b/bert_squeeze/distillation/seq2seq_distiller.py @@ -1,9 +1,7 @@ -from typing import Any, Dict, TypeVar, Union +from typing import Dict, Optional, Union import lightning.pytorch as pl -import numpy as np import torch -import torch.nn.functional as F from omegaconf import DictConfig from overrides import overrides from torch.nn import CrossEntropyLoss @@ -35,7 +33,7 @@ def __init__( teacher: Union["pl.LightningModule", "torch.nn.Module"], student: Union["pl.LightningModule", "torch.nn.Module"], training_config: DictConfig, - teacher_checkpoint: str = None, + teacher_checkpoint: Optional[str] = None, **kwargs, ): super().__init__(teacher, student, training_config, teacher_checkpoint, **kwargs) @@ -127,6 +125,7 @@ def training_step(self, batch, _) -> torch.Tensor: for key, val in self.s_scorer.losses.items() } self.log_dict(logging_loss) + self._log_training_loss(loss.full_loss) return loss.full_loss @overrides diff --git a/bert_squeeze/distillation/sequence_classification_distiller.py b/bert_squeeze/distillation/sequence_classification_distiller.py index f13f5d3..ff9718d 100644 --- a/bert_squeeze/distillation/sequence_classification_distiller.py +++ b/bert_squeeze/distillation/sequence_classification_distiller.py @@ -1,5 +1,5 @@ import logging -from typing import Dict, List, Tuple, Union +from typing import Dict, List, Optional, Tuple, Union import lightning.pytorch as pl import matplotlib.pyplot as plt @@ -43,7 +43,7 @@ def __init__( student: Union["pl.LightningModule", "torch.nn.Module"], training_config: DictConfig, labels: Union[List[str], List[int]], - teacher_checkpoint: str = None, + teacher_checkpoint: Optional[str] = None, **kwargs, ): super().__init__(teacher, student, training_config, teacher_checkpoint, **kwargs) @@ -158,7 +158,7 @@ def __init__( student: Union["pl.LightningModule", "torch.nn.Module"], training_config: DictConfig, labels: Union[List[str], List[int]], - teacher_checkpoint: str = None, + teacher_checkpoint: Optional[str] = None, **kwargs, ): super().__init__( @@ -220,7 +220,7 @@ def loss( self, teacher_logits: torch.Tensor, student_logits: torch.Tensor, - labels: torch.Tensor = None, + labels: Optional[torch.Tensor] = None, ignore_index: int = -100, *args, **kwargs, @@ -268,6 +268,7 @@ def training_step(self, batch, _) -> torch.Tensor: self.log_dict(logging_loss) self.log("train/acc", self.scorer.acc) + self._log_training_loss(loss.full_loss) return loss.full_loss @overrides @@ -350,7 +351,7 @@ def __init__( student: Union["pl.LightningModule", "torch.nn.Module"], training_config: DictConfig, labels: Union[List[str], List[int]], - teacher_checkpoint: str = None, + teacher_checkpoint: Optional[str] = None, **kwargs, ): super().__init__( @@ -464,6 +465,7 @@ def training_step(self, batch, _) -> torch.Tensor: s_logits_original, s_logits_translated = self.get_student_logits(batch) loss = self.loss(t_logits, s_logits_original, s_logits_translated) + self._log_training_loss(loss.full_loss) return loss.full_loss @overrides diff --git a/bert_squeeze/models/base_lt_module.py b/bert_squeeze/models/base_lt_module.py index 011c6d2..98b1813 100644 --- a/bert_squeeze/models/base_lt_module.py +++ b/bert_squeeze/models/base_lt_module.py @@ -13,6 +13,7 @@ import torch.nn.functional as F from omegaconf import DictConfig, ListConfig from torch.nn import CrossEntropyLoss +from torch.optim.lr_scheduler import ReduceLROnPlateau from transformers import ( AutoConfig, AutoModelForSeq2SeqLM, @@ -25,9 +26,7 @@ BertAdam, OptimizerParameterGroup, build_optimizer_parameter_groups, - register_legacy_optimizer_state_migration, ) -from ..utils.schedulers import GroupCompatibleReduceLROnPlateau from ..utils.scorers import BaseSequenceClassificationScorer, LMScorer, Scorer from ..utils.types import ( FastBertLoss, @@ -155,12 +154,13 @@ def configure_optimizers(self) -> Tuple[List, List]: else: raise ValueError(f"Optimizer '{self.config.optimizer}' not supported.") - if self.config.discriminative_learning: - register_legacy_optimizer_state_migration(optimizer, self.named_parameters()) - if optimizer_name == "adamw" and self.config.lr_scheduler: - scheduler = GroupCompatibleReduceLROnPlateau(optimizer) - lr_scheduler = {'scheduler': scheduler, 'name': 'NeptuneLogger'} + scheduler = ReduceLROnPlateau(optimizer) + lr_scheduler = { + 'scheduler': scheduler, + 'name': 'learning_rate', + 'monitor': self.config.get("lr_scheduler_monitor", "train/epoch_loss"), + } return [optimizer], [lr_scheduler] return [optimizer], [] @@ -298,6 +298,12 @@ def training_step( ) -> torch.Tensor: self._before_training_step() step_output = self._classification_step(batch, training=True) + self.log( + "train/epoch_loss", + step_output.optimization_loss, + on_step=False, + on_epoch=True, + ) self._update_scorer(self.scorer, step_output) self._log_training_metrics() return step_output.optimization_loss diff --git a/bert_squeeze/models/custom_transformers/berxit.py b/bert_squeeze/models/custom_transformers/berxit.py index 4c8c43b..ff2c58c 100644 --- a/bert_squeeze/models/custom_transformers/berxit.py +++ b/bert_squeeze/models/custom_transformers/berxit.py @@ -1,19 +1,8 @@ -""" -This implementation mirrors the DeeBERT integration but under the BERxiT -name, following the repository's integration patterns so users can train -and use a BERxiT-style early-exiting model. - -Note: This module reuses the same architectural approach as DeeBERT in this -codebase (off-ramps between layers with entropy-based early exit), -so it integrates seamlessly with existing training loops and configs. -""" - from abc import ABC -from typing import List, Tuple, Union +from typing import List, Optional, Tuple, Union import torch import torch.nn as nn -import torch.nn.functional as F from transformers import PretrainedConfig from transformers.models.bert.modeling_bert import ( BertEmbeddings, @@ -27,51 +16,20 @@ from .deebert import OffRamp -class ExitGate(nn.Module): - """A small MLP gate that predicts whether to exit at a given layer. - - Inputs are hand-crafted features from the ramp logits/probs. - """ - - def __init__(self, in_dim: int = 3, hidden: int = 32): - super().__init__() - self.net = nn.Sequential( - nn.Linear(in_dim, hidden), - nn.ReLU(), - nn.Linear(hidden, 1), - ) - - def forward(self, feats: torch.Tensor) -> torch.Tensor: - # Returns logits for BCEWithLogitsLoss - return self.net(feats) - - class BerxitEncoder(nn.Module): - """ - Encoder that inserts off-ramps between each Transformer block and - supports early exiting using an entropy threshold. - """ + """BERT encoder with classifier ramps and a shared exit gate.""" def __init__(self, config: PretrainedConfig, inference: bool): super(BerxitEncoder, self).__init__() self.config = config - # Ensure distinct modules per layer self.layer = nn.ModuleList( [BertLayer(config) for _ in range(config.num_hidden_layers)] ) self.ramp = nn.ModuleList( [OffRamp(config) for _ in range(config.num_hidden_layers)] ) - # BERxiT gates - gate_hidden = getattr(config, "gate_hidden_dim", 32) - self.gates = nn.ModuleList( - [ - ExitGate(in_dim=3, hidden=gate_hidden) - for _ in range(config.num_hidden_layers) - ] - ) + self.gates = nn.Linear(config.hidden_size, 1) - # Thresholds for DeeBERT entropy (fallback) and BERxiT gate self.early_exit_entropy = [-1.0] * config.num_hidden_layers self.gate_thresholds = [-1.0] * config.num_hidden_layers self.inference = inference @@ -98,25 +56,14 @@ def set_exit_gate_thresholds(self, x: Union[List[float], float]) -> None: if isinstance(x, float) or isinstance(x, int): self.gate_thresholds = [float(x)] * self.config.num_hidden_layers elif isinstance(x, list): - assert ( - len(x) == self.config.num_hidden_layers - ), "gate thresholds size mismatch" + if len(x) != self.config.num_hidden_layers: + raise ValueError("Gate threshold count must match the encoder depth.") self.gate_thresholds = [float(v) for v in x] else: raise TypeError( f"Expected 'x' to be of type 'float' or 'list' but got :'{type(x)}'" ) - @staticmethod - def _gate_features(logits: torch.Tensor) -> torch.Tensor: - # logits: [B, C] - probs = F.softmax(logits, dim=-1) - pmax, _ = probs.max(dim=-1, keepdim=True) # [B,1] - top2 = torch.topk(probs, k=2, dim=-1).values # [B,2] - margin = (top2[:, 0] - top2[:, 1]).unsqueeze(-1) # [B,1] - ent = entropy(probs).unsqueeze(-1) # [B,1] - return torch.cat([pmax, margin, ent], dim=-1) # [B,3] - def forward( self, hidden_states: torch.Tensor, @@ -127,15 +74,15 @@ def forward( output_attentions: bool = False, output_hidden_states: bool = False, ) -> DeeBertEncoderOutput: - all_hidden_states = tuple() if output_hidden_states else None - all_attentions = tuple() if output_attentions else None + all_hidden_states: List[torch.Tensor] = [] + all_attentions: List[torch.Tensor] = [] if not self.inference: - all_ramps: Tuple[RampOutput, ...] = tuple() - all_gates: Tuple[torch.Tensor, ...] = tuple() + all_ramps: List[RampOutput] = [] + all_gates: List[torch.Tensor] = [] for i, layer_module in enumerate(self.layer): if output_hidden_states: - all_hidden_states += (hidden_states,) + all_hidden_states.append(hidden_states) layer_outputs = layer_module( hidden_states=hidden_states, @@ -149,98 +96,93 @@ def forward( if output_attentions: attention = layer_outputs[1] - all_attentions += (attention,) + all_attentions.append(attention) ramp_exit = self.ramp[i](hidden_states) - all_ramps += (ramp_exit,) - # Gate logits from features of current ramp - feats = self._gate_features(ramp_exit.logits) - gate_logit = self.gates[i](feats) # [B,1] - all_gates += (gate_logit,) + all_ramps.append(ramp_exit) + gate_logit = self.gates(hidden_states[:, 0]) + all_gates.append(gate_logit) if output_hidden_states: - all_hidden_states = all_hidden_states + (hidden_states,) + all_hidden_states.append(hidden_states) return DeeBertEncoderOutput( last_hidden_state=hidden_states, - hidden_states=all_hidden_states, - attentions=all_attentions, - ramps_exit=all_ramps, - gates_logits=all_gates, + hidden_states=( + tuple(all_hidden_states) if output_hidden_states else None + ), + attentions=tuple(all_attentions) if output_attentions else None, + ramps_exit=tuple(all_ramps), + gates_logits=tuple(all_gates), exit_layer=i, ) - else: - batch_size = hidden_states.shape[0] - all_ramps = [0] * batch_size - positions = torch.arange( - start=0, end=hidden_states.shape[0], device=hidden_states.device - ).long() - # Collect per-layer gate logits for diagnostics; fill with NaNs by default - gates_per_layer: Tuple[torch.Tensor, ...] = tuple( - torch.full( - (batch_size, 1), - float('nan'), - device=hidden_states.device, - dtype=hidden_states.dtype, - ) - for _ in range(len(self.layer)) + + batch_size = hidden_states.shape[0] + inference_ramps: List[Optional[RampOutput]] = [None] * batch_size + positions = torch.arange( + start=0, end=hidden_states.shape[0], device=hidden_states.device + ).long() + gates_per_layer: Tuple[torch.Tensor, ...] = tuple( + torch.full( + (batch_size, 1), + float("nan"), + device=hidden_states.device, + dtype=hidden_states.dtype, ) + for _ in range(len(self.layer)) + ) - for i, layer_module in enumerate(self.layer): - layer_outputs = layer_module( - hidden_states=hidden_states, - attention_mask=attention_mask, - head_mask=head_mask[i], - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - output_attentions=output_attentions, - ) - hidden_states = layer_outputs[0] - ramp_exit = self.ramp[i](hidden_states) - # Compute gate decision - feats = self._gate_features(ramp_exit.logits) - gate_logit = self.gates[i](feats) - # Scatter current gate logits back to original batch positions - gates_per_layer[i][positions] = gate_logit - gate_prob = torch.sigmoid(gate_logit).squeeze(-1) # [b_cur] - ramp_exit.entropy = entropy(ramp_exit.logits) - - if i == len(self.layer) - 1: - for idx, pos in enumerate(positions): - all_ramps[pos] = ramp_exit[idx] - else: - # Prefer gate thresholds; fallback to entropy if thresholds are negative - if self.gate_thresholds[i] >= 0: - enough_info = gate_prob >= self.gate_thresholds[i] - else: - enough_info = ramp_exit.entropy < self.early_exit_entropy[i] - right_pos = positions[enough_info] - - for idx, pos in enumerate(right_pos): - all_ramps[pos] = ramp_exit[idx] - - hidden_states = hidden_states[~enough_info] - attention_mask = attention_mask[~enough_info] - positions = positions[~enough_info] - - if positions.nelement() == 0: - return DeeBertEncoderOutput( - ramps_exit=all_ramps, - gates_logits=gates_per_layer, - exit_layer=i, - ) - return DeeBertEncoderOutput( - ramps_exit=all_ramps, - gates_logits=gates_per_layer, - exit_layer=i, + for i, layer_module in enumerate(self.layer): + layer_outputs = layer_module( + hidden_states=hidden_states, + attention_mask=attention_mask, + head_mask=head_mask[i], + encoder_hidden_states=encoder_hidden_states, + encoder_attention_mask=encoder_attention_mask, + output_attentions=output_attentions, ) + hidden_states = layer_outputs[0] + ramp_exit = self.ramp[i](hidden_states) + gate_logit = self.gates(hidden_states[:, 0]) + gates_per_layer[i][positions] = gate_logit + gate_prob = torch.sigmoid(gate_logit).squeeze(-1) + ramp_entropy = entropy(ramp_exit.logits) + ramp_exit.entropy = ramp_entropy + + is_final_layer = i == len(self.layer) - 1 + if is_final_layer: + enough_info = torch.ones_like(gate_prob, dtype=torch.bool) + elif self.gate_thresholds[i] >= 0: + enough_info = gate_prob >= self.gate_thresholds[i] + else: + enough_info = ramp_entropy < self.early_exit_entropy[i] + right_pos = positions[enough_info] + + for idx, pos in enumerate(right_pos): + inference_ramps[pos] = ramp_exit[idx] + + if is_final_layer: + continue + + hidden_states = hidden_states[~enough_info] + attention_mask = attention_mask[~enough_info] + positions = positions[~enough_info] + + if positions.nelement() == 0: + break + + completed_ramps = tuple(ramp for ramp in inference_ramps if ramp is not None) + if len(completed_ramps) != batch_size: + raise RuntimeError("BERxiT inference did not produce every sample.") + return DeeBertEncoderOutput( + ramps_exit=completed_ramps, + gates_logits=gates_per_layer, + exit_layer=i, + ) class BerxitModel(BertPreTrainedModel, ABC): - """ - BERxiT-like BERT model with off-ramps and early exit, matching the - interfaces used by the DeeBERT integration in this repository. - """ + """BERT model with BERxiT early exits.""" def __init__(self, config: PretrainedConfig, inference: bool = False): super(BerxitModel, self).__init__(config) diff --git a/bert_squeeze/models/lt_berxit.py b/bert_squeeze/models/lt_berxit.py index 9d1978b..ca3750b 100644 --- a/bert_squeeze/models/lt_berxit.py +++ b/bert_squeeze/models/lt_berxit.py @@ -23,12 +23,7 @@ class LtBerxit(BaseSequenceClassificationTransformerModule): - """ - Lightning module to fine-tune a BERxiT-style model for sequence classification. - - This mirrors LtDeeBert's integration, exposing the same training/inference - behavior and configuration hooks (e.g., train_highway, early_exit_entropy). - """ + """Fine-tune BERxiT models for sequence classification.""" def __init__( self, @@ -50,18 +45,15 @@ def __init__( super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) - # Training stage: "backbone" (default) or "gates" self.train_stage = getattr(training_config, "train_stage", "backbone") - # Optional global step at which to switch from backbone to gate training self.switch_step: Optional[int] = getattr(training_config, "switch_step", None) self.train_highway = training_config.train_highway - # In backbone stage, gates are optional; in gates stage, we always train them - if self.train_stage == "gates": - self.train_gates = True - else: - self.train_gates = getattr(training_config, "train_gates", False) + self.train_gates = self.train_stage == "gates" or getattr( + training_config, "train_gates", False + ) self._build_model() - # Guard to ensure stage switching happens at most once + if self.train_stage == "gates": + self._freeze_backbone_for_gates() self._has_switched_stage = False @overrides @@ -89,10 +81,8 @@ def forward( if not self.bert.encoder.inference: exit_layer = self.num_layers - pooled_output = outputs.pooled_output - pooled_output = self.dropout(pooled_output) - logits = self.classifier(pooled_output) ramps_exits = outputs.ramps_exits + logits = ramps_exits[-1].logits gates_logits = outputs.gates_logits else: ramps_exits = outputs.ramps_exits @@ -125,7 +115,7 @@ def _classification_loss( return self.loss( logits=output.logits, labels=labels, - train_ramps=self.train_highway, + train_ramps=self._train_all_exits(), ramps_exits=output.ramps_exits, train_gates=self.train_gates, gates_logits=output.gates_logits, @@ -139,23 +129,14 @@ def _before_training_step(self) -> None: def _before_prediction_step(self) -> None: self.bert.set_inference_mode(inference=True) + def _train_all_exits(self) -> bool: + return self.train_highway and self.global_step % 2 == 1 + def _freeze_backbone_for_gates(self) -> None: - """ - Freeze all parameters except the BERXiT gates so that gate training - happens on top of a fixed teacher model. - """ - for name, param in self.named_parameters(): - if "gates" in name: - param.requires_grad = True - else: - param.requires_grad = False + for name, parameter in self.named_parameters(): + parameter.requires_grad = "gates" in name def _maybe_switch_stage(self) -> None: - """ - If `switch_step` is set and we are still in backbone stage, switch to - gate training once `global_step` reaches the threshold. This triggers - a reconfiguration of optimizers in Lightning. - """ if ( self.switch_step is None or self.train_stage == "gates" @@ -172,7 +153,6 @@ def _maybe_switch_stage(self) -> None: self.train_gates = True self._freeze_backbone_for_gates() if self.trainer is not None: - # Rebuild optimizers so only gate parameters are optimized self.trainer.strategy.setup_optimizers(self.trainer) self._has_switched_stage = True @@ -193,14 +173,11 @@ def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: ) discriminative_learning = self.config.discriminative_learning - if discriminative_learning: - named_parameters = self.named_parameters() - else: - named_parameters = ( - (name, parameter) - for name, parameter in self.named_parameters() - if self._is_active_training_parameter(name) - ) + named_parameters = ( + (name, parameter) + for name, parameter in self.named_parameters() + if parameter.requires_grad + ) return build_optimizer_parameter_groups( named_parameters, discriminative_learning=discriminative_learning, @@ -209,13 +186,6 @@ def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: weight_decay=self.config.weight_decay, ) - def _is_active_training_parameter(self, parameter_name: str) -> bool: - if ".ramp." in parameter_name: - return self.config.train_highway - if ".gates." in parameter_name: - return self.train_gates - return not self.config.train_highway - @overrides def loss( self, @@ -228,64 +198,51 @@ def loss( *args, **kwargs, ) -> torch.Tensor: - # Same ramp loss mechanics as LtDeeBert for consistency - if train_ramps: - if ramps_exits is None or len(ramps_exits) < 2: - raise ValueError("Ramp training requires at least two ramp outputs.") - ramps_losses: List[torch.Tensor] = [] - for ramps_exit in ramps_exits[:-1]: - ramps_logits = ramps_exit.logits - loss_fct = CrossEntropyLoss() - ramps_loss = loss_fct( - ramps_logits.view(-1, self.model_config.num_labels), labels.view(-1) - ) - ramps_losses.append(ramps_loss) - loss = torch.stack(ramps_losses).sum() - else: + if ramps_exits is None: if logits is None: - raise ValueError( - "Classifier logits are required when ramps are disabled." - ) - loss_fct = CrossEntropyLoss() - loss = loss_fct( + raise ValueError("BERxiT training requires classifier outputs.") + return CrossEntropyLoss()( logits.view(-1, self.model_config.num_labels), labels.view(-1) ) - # Optional: add gate loss using pseudo-labels from final ramp - if train_gates: - if ramps_exits is None or gates_logits is None: - raise ValueError("Gate training requires ramp and gate outputs.") - with torch.no_grad(): - final_logits = ramps_exits[-1].logits # [B, C] - final_pred = final_logits.argmax(dim=-1) # [B] - bce = torch.nn.BCEWithLogitsLoss() - gate_losses: List[torch.Tensor] = [] - for i, gate_logit in enumerate(gates_logits[:-1]): - layer_pred = ramps_exits[i].logits.argmax(dim=-1) # [B] - target = (layer_pred == final_pred).float().unsqueeze(-1) # [B,1] - gate_losses.append(bce(gate_logit, target)) - if gate_losses: - loss = loss + torch.stack(gate_losses).sum() - return loss - def _build_model(self): - # Pass BERxiT-specific hyperparams via HF config attributes - if not hasattr(self.model_config, "gate_hidden_dim"): - self.model_config.gate_hidden_dim = getattr( - self.config, "gate_hidden_dim", 32 + exit_indices = ( + tuple(range(len(ramps_exits))) if train_ramps else (len(ramps_exits) - 1,) + ) + loss_fct = CrossEntropyLoss() + classification_losses = [ + loss_fct( + ramps_exits[index].logits.view(-1, self.model_config.num_labels), + labels.view(-1), ) + for index in exit_indices + ] + loss = torch.stack(classification_losses).sum() + + if not train_gates: + return loss + if gates_logits is None or len(gates_logits) != len(ramps_exits): + raise ValueError("Gate training requires one gate output per ramp.") + return loss + self._gate_loss(labels, ramps_exits, gates_logits, exit_indices) + + @staticmethod + def _gate_loss( + labels: torch.Tensor, + ramps_exits: Sequence[RampOutput], + gates_logits: Sequence[torch.Tensor], + exit_indices: Sequence[int], + ) -> torch.Tensor: + gate_losses = [] + for index in exit_indices: + prediction = ramps_exits[index].logits.argmax(dim=-1) + target = (prediction == labels).float() + certainty = torch.sigmoid(gates_logits[index]).squeeze(-1) + gate_losses.append(torch.nn.functional.mse_loss(certainty, target)) + return torch.stack(gate_losses).sum() + + def _build_model(self) -> None: self.bert = self.model self.num_layers = len(self.bert.encoder.layer) - self.dropout = nn.Dropout(self.model_config.hidden_dropout_prob) - self.classifier = torch.nn.Sequential( - torch.nn.Dropout(self.model_config.hidden_dropout_prob), - torch.nn.Linear(self.model_config.hidden_size, self.model_config.hidden_size), - torch.nn.ReLU(), - torch.nn.LayerNorm(self.model_config.hidden_size), - torch.nn.Linear(self.model_config.hidden_size, self.model_config.num_labels), - ) - self.bert.encoder.set_early_exit_entropy(self.config.early_exit_entropy) - # Optional: set gate thresholds for early exit if hasattr(self.config, "gate_thresholds"): self.bert.set_exit_gate_thresholds(self.config.gate_thresholds) self.bert.init_highway_pooler() diff --git a/bert_squeeze/models/lt_deebert.py b/bert_squeeze/models/lt_deebert.py index 1d9e64e..a178343 100644 --- a/bert_squeeze/models/lt_deebert.py +++ b/bert_squeeze/models/lt_deebert.py @@ -22,22 +22,7 @@ class LtDeeBert(BaseSequenceClassificationTransformerModule): - """ - Lightning module to fine-tune a DeeBert based model on a sequence classification - task (see `models.custom_transformers.deebert.py`) for detailed explanation. - - Args: - training_config (DictConfig): - training configuration - num_labels (int): - number of labels - pretrained_model (str): - name of the pretrained Transformer model to use - model (Optional[Union[pl.LightningModule, nn.Module]]): - optional instantiated model - scorer (Scorer): - helper object to compute performance metrics during training - """ + """Fine-tune DeeBERT models for sequence classification.""" def __init__( self, @@ -72,35 +57,6 @@ def forward( head_mask: torch.Tensor = None, **kwargs, ) -> Tuple[torch.Tensor, Sequence[RampOutput], int]: - """ - During training, we pass the hidden states through all layers and store all the off-ramps - outputs as well as the final classification layer. - During inference, we try to pass the hidden states through the whole BertLayer and OffRamps stack - which is exited as soon as the entropy of one layer is lower than a given threshold. - - Args: - input_ids (torch.Tensor): - sentence or sentences represented as tokens - attention_mask (torch.Tensor): - tells the model which tokens in the input_ids are words and which are padding. - 1 indicates a token and 0 indicates padding. - token_type_ids (torch.Tensor): - used when there are two sentences that need to be part of the input. It indicates which - tokens are part of sentence1 and which are part of sentence2. - position_ids (torch.Tensor): - indices of positions of each input sequence tokens in the position embeddings. Selected - in the range ``[0, config.max_position_embeddings - 1] - head_mask (torch.Tensor): - mask to nullify selected heads of the self-attention modules - Returns: - torch.Tensor: - output of the classification layer which uses the last ramp output during training and - the output of the exited ramp during inference. - Tuple[torch.Tensor]: - iterable containing all ramp exits - int: - index of the exited ramp - """ outputs = self.bert( input_ids, attention_mask=attention_mask, @@ -111,10 +67,8 @@ def forward( if not self.bert.encoder.inference: exit_layer = self.num_layers - pooled_output = outputs.pooled_output - pooled_output = self.dropout(pooled_output) - logits = self.classifier(pooled_output) ramps_exits = outputs.ramps_exits + logits = ramps_exits[-1].logits else: ramps_exits = outputs.ramps_exits exit_layer = outputs.exit_layer @@ -152,16 +106,14 @@ def _before_prediction_step(self) -> None: @overrides def _get_optimizer_parameters(self) -> List[OptimizerParameterGroup]: - discriminative_learning = self.config.discriminative_learning - if discriminative_learning: - named_parameters = self.named_parameters() - else: - ramp_only = self.config.train_highway - named_parameters = ( - (name, parameter) - for name, parameter in self.named_parameters() - if (".ramp." in name) == ramp_only - ) + named_parameters = ( + (name, parameter) + for name, parameter in self.named_parameters() + if parameter.requires_grad + ) + discriminative_learning = ( + self.config.discriminative_learning and not self.train_highway + ) return build_optimizer_parameter_groups( named_parameters, discriminative_learning=discriminative_learning, @@ -180,64 +132,41 @@ def loss( *args, **kwargs, ) -> torch.Tensor: - """ - Handles the loss computation part. - - If `train_ramps=False` we only use the logits of the final classification layer to compute - the cross entropy. If `train_ramps=True` we add up all the cross entropies of the off-ramps. - - Args: - labels (torch.Tensor): - ground truth labels - ramps_exits (Tuple[torch.Tensor]): - list containing the predicted logits from all the off-ramps - logits (torch.Tensor): - predicted logits by the final classification layer - train_ramps (bool): - whether to train the off-ramps or the final classification layer. - Returns: - - """ - # We want to fine-tune each individual ramp + loss_fct = CrossEntropyLoss() if train_ramps: if ramps_exits is None or len(ramps_exits) < 2: raise ValueError("Ramp training requires at least two ramp outputs.") - ramps_losses: List[torch.Tensor] = [] - # We train all but the last off-ramp (corresponds to stage 2 in paper) - for ramps_exit in ramps_exits[:-1]: - ramps_logits = ramps_exit.logits - - loss_fct = CrossEntropyLoss() - ramps_loss = loss_fct( - ramps_logits.view(-1, self.model_config.num_labels), labels.view(-1) - ) - ramps_losses.append(ramps_loss) - - loss = torch.stack(ramps_losses).sum() - else: - if logits is None: - raise ValueError( - "Classifier logits are required when ramps are disabled." - ) - # We only train the last off-ramp - loss_fct = CrossEntropyLoss() - loss = loss_fct( - logits.view(-1, self.model_config.num_labels), labels.view(-1) - ) - return loss - - def _build_model(self): - """""" + return torch.stack( + [ + loss_fct( + ramp.logits.view(-1, self.model_config.num_labels), + labels.view(-1), + ) + for ramp in ramps_exits[:-1] + ] + ).sum() + + if logits is None: + raise ValueError("Classifier logits are required when ramps are disabled.") + return loss_fct(logits.view(-1, self.model_config.num_labels), labels.view(-1)) + + def _build_model(self) -> None: self.bert = self.model self.num_layers = len(self.bert.encoder.layer) - self.dropout = nn.Dropout(self.model_config.hidden_dropout_prob) - self.classifier = torch.nn.Sequential( - torch.nn.Dropout(self.model_config.hidden_dropout_prob), - torch.nn.Linear(self.model_config.hidden_size, self.model_config.hidden_size), - torch.nn.ReLU(), - torch.nn.LayerNorm(self.model_config.hidden_size), - torch.nn.Linear(self.model_config.hidden_size, self.model_config.num_labels), - ) - self.bert.encoder.set_early_exit_entropy(self.config.early_exit_entropy) self.bert.init_highway_pooler() + self._set_trainable_parameters() + + def _set_trainable_parameters(self) -> None: + final_ramp = f".ramp.{self.num_layers - 1}." + for name, parameter in self.named_parameters(): + is_ramp = ".ramp." in name + is_final_ramp = final_ramp in name + if self.train_highway: + trainable = is_ramp and not is_final_ramp + else: + trainable = not is_ramp or is_final_ramp + + if ".pooler." in name and not is_ramp: + trainable = False + parameter.requires_grad = trainable diff --git a/bert_squeeze/models/lt_t5.py b/bert_squeeze/models/lt_t5.py index 4a69776..01fc7f1 100644 --- a/bert_squeeze/models/lt_t5.py +++ b/bert_squeeze/models/lt_t5.py @@ -84,6 +84,12 @@ def training_step(self, batch, batch_idx, *args, **kwargs): ) self.scorer.reset() + self.log( + "train/epoch_loss", + outputs.loss, + on_step=False, + on_epoch=True, + ) return outputs.loss def validation_step(self, batch, batch_idx, *args, **kwargs) -> dict: diff --git a/bert_squeeze/utils/optimizers/__init__.py b/bert_squeeze/utils/optimizers/__init__.py index abdf7f3..65740fa 100644 --- a/bert_squeeze/utils/optimizers/__init__.py +++ b/bert_squeeze/utils/optimizers/__init__.py @@ -1,6 +1,2 @@ from .bert_adam import BertAdam -from .parameter_groups import ( - OptimizerParameterGroup, - build_optimizer_parameter_groups, - register_legacy_optimizer_state_migration, -) +from .parameter_groups import OptimizerParameterGroup, build_optimizer_parameter_groups diff --git a/bert_squeeze/utils/optimizers/parameter_groups.py b/bert_squeeze/utils/optimizers/parameter_groups.py index 9075090..405a249 100644 --- a/bert_squeeze/utils/optimizers/parameter_groups.py +++ b/bert_squeeze/utils/optimizers/parameter_groups.py @@ -2,17 +2,11 @@ import re from collections.abc import Iterable, Sequence -from typing import Optional, TypedDict, Union, cast +from typing import Optional, TypedDict, Union from torch import nn -from torch.optim import Optimizer -from torch.optim.optimizer import StateDict -__all__ = [ - "OptimizerParameterGroup", - "build_optimizer_parameter_groups", - "register_legacy_optimizer_state_migration", -] +__all__ = ["OptimizerParameterGroup", "build_optimizer_parameter_groups"] class _RequiredOptimizerParameterGroup(TypedDict): @@ -24,13 +18,21 @@ class OptimizerParameterGroup(_RequiredOptimizerParameterGroup, total=False): lr: float +_LayerKey = tuple[str, int] + _LAYER_PATTERNS = ( - (re.compile(r"(?:^|\.)block\.(\d+)\."), False), - (re.compile(r"(?:^|\.)layers\.(\d+)\."), False), - (re.compile(r"(?:^|\.)h\.(\d+)\."), False), - (re.compile(r"(?:^|\.)layer\.(\d+)\."), True), + re.compile(r"(?:^|\.)(block)\.(\d+)\."), + re.compile(r"(?:^|\.)(layers)\.(\d+)\."), + re.compile(r"(?:^|\.)(h)\.(\d+)\."), + re.compile(r"(?:^|\.)(layer)\.(\d+)\."), +) +_EMBEDDING_MODULES = ( + "embeddings", + "embed_tokens", + "embed_positions", + "wte", + "wpe", ) -_NO_DECAY_NAMES = ("bias", "gamma", "beta", "LayerNorm.weight", "layer_norm.weight") def build_optimizer_parameter_groups( @@ -45,196 +47,120 @@ def build_optimizer_parameter_groups( if not discriminative_learning: return _weight_decay_groups(parameters, weight_decay) - parameters_by_layer, remaining_parameters, uses_legacy_layout = ( - _split_parameters_by_layer(parameters) - ) - layer_indices = sorted(parameters_by_layer) - if not layer_indices: + rates = _learning_rate_values(learning_rates) + layer_keys = [_layer_key(name) for name, _ in parameters] + stack_indices = _stack_indices(layer_keys) + if not stack_indices: raise ValueError("No encoder layers found for discriminative learning.") - layer_rates = _layer_rates(learning_rates, layer_lr_decay, len(layer_indices)) - layer_rate_by_index = dict(zip(layer_indices, layer_rates)) - preserve_legacy_slots = uses_legacy_layout and all( - 0 <= index < 12 for index in layer_indices - ) - group_indices = list(range(12)) if preserve_legacy_slots else layer_indices - groups = [] - for use_weight_decay in (True, False): - for layer_index in group_indices: - layer_parameters = [ - parameter - for name, parameter in parameters_by_layer.get(layer_index, []) - if _uses_weight_decay(name) == use_weight_decay - ] - if layer_parameters or preserve_legacy_slots: - groups.append( - _parameter_group( - layer_parameters, - weight_decay, - use_weight_decay, - layer_rate_by_index.get(layer_index, layer_rates[-1]), - ) - ) - - groups.extend(_weight_decay_groups(remaining_parameters, weight_decay)) - return groups + layer_rate_by_key = _layer_rates_by_key(rates, layer_lr_decay, stack_indices) + embedding_rate = min(layer_rate_by_key.values()) * layer_lr_decay + head_rate = rates[-1] + + grouped_parameters: dict[tuple[float, bool], list[nn.Parameter]] = {} + for (name, parameter), layer_key in zip(parameters, layer_keys): + learning_rate = _parameter_learning_rate( + name, + layer_key, + layer_rate_by_key, + embedding_rate, + head_rate, + ) + group_key = (learning_rate, _uses_weight_decay(name)) + grouped_parameters.setdefault(group_key, []).append(parameter) + return [ + _parameter_group(parameters, weight_decay, use_weight_decay, learning_rate) + for (learning_rate, use_weight_decay), parameters in grouped_parameters.items() + ] -def register_legacy_optimizer_state_migration( - optimizer: Optimizer, - named_parameters: Iterable[tuple[str, nn.Parameter]], -) -> None: - legacy_parameter_groups = _legacy_parameter_groups(list(named_parameters)) - - def migrate_state_dict( - current_optimizer: Optimizer, state_dict: StateDict - ) -> Optional[StateDict]: - saved_groups = cast(list[dict[str, object]], state_dict["param_groups"]) - if len(saved_groups) != len(legacy_parameter_groups): - return None - - saved_metadata_by_parameter = _saved_parameter_metadata( - saved_groups, legacy_parameter_groups - ) - if saved_metadata_by_parameter is None: - return None - current_groups = cast(list[dict[str, object]], current_optimizer.param_groups) - serialized_groups = cast( - list[dict[str, object]], current_optimizer.state_dict()["param_groups"] - ) - migrated_groups = [] - for current_group, serialized_group in zip(current_groups, serialized_groups): - current_parameters = cast(list[nn.Parameter], current_group["params"]) - if any( - id(parameter) not in saved_metadata_by_parameter - for parameter in current_parameters - ): - return None - source_group_indices = { - saved_metadata_by_parameter[id(parameter)][1] - for parameter in current_parameters - } - if len(source_group_indices) == 1: - source_group_index = next(iter(source_group_indices)) - source_group = saved_groups[source_group_index] - migrated_group = { - key: value for key, value in source_group.items() if key != "params" - } - else: - migrated_group = { - key: value - for key, value in serialized_group.items() - if key != "params" - } - migrated_group["params"] = [ - saved_metadata_by_parameter[id(parameter)][0] - for parameter in current_parameters - ] - migrated_groups.append(migrated_group) - - migrated_state_dict = dict(state_dict) - migrated_state_dict["param_groups"] = migrated_groups - return cast(StateDict, migrated_state_dict) - - optimizer.register_load_state_dict_pre_hook(migrate_state_dict) - - -def _split_parameters_by_layer( - parameters: Sequence[tuple[str, nn.Parameter]], -) -> tuple[ - dict[int, list[tuple[str, nn.Parameter]]], - list[tuple[str, nn.Parameter]], - bool, -]: - parameters_by_layer: dict[int, list[tuple[str, nn.Parameter]]] = {} - remaining_parameters = [] - uses_legacy_layout = True - for name, parameter in parameters: - layer_match = _layer_match(name) - if layer_match is None: - remaining_parameters.append((name, parameter)) +def _learning_rate_values( + learning_rates: Union[float, Sequence[float]], +) -> list[float]: + rates = ( + [float(rate) for rate in learning_rates] + if isinstance(learning_rates, Sequence) + else [float(learning_rates)] + ) + if not rates: + raise ValueError("At least one learning rate is required.") + return rates + + +def _stack_indices(layer_keys: Sequence[Optional[_LayerKey]]) -> dict[str, set[int]]: + indices: dict[str, set[int]] = {} + for layer_key in layer_keys: + if layer_key is None: continue - layer_index, is_legacy_layer = layer_match - uses_legacy_layout = uses_legacy_layout and is_legacy_layer - parameters_by_layer.setdefault(layer_index, []).append((name, parameter)) - return parameters_by_layer, remaining_parameters, uses_legacy_layout + stack, layer_index = layer_key + indices.setdefault(stack, set()).add(layer_index) + return indices -def _legacy_parameter_groups( - parameters: Sequence[tuple[str, nn.Parameter]], -) -> list[list[nn.Parameter]]: - layer_keys = [f"layer.{index}." for index in range(12)] - groups: list[list[nn.Parameter]] = [] - for use_weight_decay in (True, False): - groups.extend( - [ - parameter - for name, parameter in parameters - if layer_key in name and _uses_weight_decay(name) == use_weight_decay - ] - for layer_key in layer_keys - ) - for use_weight_decay in (True, False): - groups.append( - [ - parameter - for name, parameter in parameters - if not any(layer_key in name for layer_key in layer_keys) - and _uses_weight_decay(name) == use_weight_decay - ] +def _layer_rates_by_key( + rates: Sequence[float], + layer_lr_decay: float, + stack_indices: dict[str, set[int]], +) -> dict[_LayerKey, float]: + if layer_lr_decay <= 0 or layer_lr_decay > 1: + raise ValueError("layer_lr_decay must be in (0, 1].") + max_layer_count = max(len(indices) for indices in stack_indices.values()) + if len(rates) > 1 and len(rates) != max_layer_count: + raise ValueError( + f"Expected {max_layer_count} layer learning rates, received {len(rates)}." ) - return groups - -def _saved_parameter_metadata( - saved_groups: Sequence[dict[str, object]], - legacy_parameter_groups: Sequence[list[nn.Parameter]], -) -> Optional[dict[int, tuple[int, int]]]: - saved_metadata_by_parameter: dict[int, tuple[int, int]] = {} - for group_index, (saved_group, legacy_parameters) in enumerate( - zip(saved_groups, legacy_parameter_groups) - ): - saved_ids = cast(list[int], saved_group["params"]) - if len(saved_ids) != len(legacy_parameters): - return None - saved_metadata_by_parameter.update( - (id(parameter), (saved_id, group_index)) - for parameter, saved_id in zip(legacy_parameters, saved_ids) + rate_by_key: dict[_LayerKey, float] = {} + for stack, indices in stack_indices.items(): + sorted_indices = sorted(indices) + stack_rates = _stack_rates(rates, layer_lr_decay, len(sorted_indices)) + rate_by_key.update( + ((stack, layer_index), rate) + for layer_index, rate in zip(sorted_indices, stack_rates) ) - return saved_metadata_by_parameter + return rate_by_key -def _layer_match(parameter_name: str) -> Optional[tuple[int, bool]]: - for pattern, is_legacy_layer in _LAYER_PATTERNS: +def _stack_rates( + rates: Sequence[float], layer_lr_decay: float, layer_count: int +) -> list[float]: + if len(rates) > 1: + return list(rates[-layer_count:]) + return [ + rates[0] * pow(layer_lr_decay, layer_count - index - 1) + for index in range(layer_count) + ] + + +def _parameter_learning_rate( + name: str, + layer_key: Optional[_LayerKey], + layer_rate_by_key: dict[_LayerKey, float], + embedding_rate: float, + head_rate: float, +) -> float: + if layer_key is not None: + return layer_rate_by_key[layer_key] + if _is_embedding_parameter(name): + return embedding_rate + return head_rate + + +def _layer_key(parameter_name: str) -> Optional[_LayerKey]: + for pattern in _LAYER_PATTERNS: match = pattern.search(parameter_name) if match is not None: - return int(match.group(1)), is_legacy_layer + stack = parameter_name[: match.start(1)] + match.group(1) + return stack, int(match.group(2)) return None -def _layer_rates( - learning_rates: Union[float, Sequence[float]], - layer_lr_decay: float, - layer_count: int, -) -> list[float]: - rates = ( - [float(rate) for rate in learning_rates] - if isinstance(learning_rates, Sequence) - else [float(learning_rates)] +def _is_embedding_parameter(parameter_name: str) -> bool: + parts = parameter_name.lower().split(".") + return parts[-2:] == ["shared", "weight"] or any( + module in _EMBEDDING_MODULES for module in parts[:-1] ) - if not rates: - raise ValueError("At least one learning rate is required.") - if len(rates) == 1: - return [ - rates[0] * pow(layer_lr_decay, layer_count - index - 1) - for index in range(layer_count) - ] - if len(rates) != layer_count: - raise ValueError( - f"Expected {layer_count} layer learning rates, received {len(rates)}." - ) - return rates def _weight_decay_groups( @@ -248,16 +174,10 @@ def _weight_decay_groups( for name, parameter in named_parameters if _uses_weight_decay(name) == use_weight_decay ] - if not parameters: - continue - groups.append( - _parameter_group( - parameters, - weight_decay, - use_weight_decay, - None, + if parameters: + groups.append( + _parameter_group(parameters, weight_decay, use_weight_decay, None) ) - ) return groups @@ -277,4 +197,10 @@ def _parameter_group( def _uses_weight_decay(parameter_name: str) -> bool: - return not any(no_decay in parameter_name for no_decay in _NO_DECAY_NAMES) + parts = parameter_name.lower().split(".") + parameter = parts[-1] + if parameter in {"bias", "beta", "gamma"}: + return False + return parameter != "weight" or not any( + "norm" in module or module.startswith("ln_") for module in parts[:-1] + ) diff --git a/bert_squeeze/utils/schedulers/__init__.py b/bert_squeeze/utils/schedulers/__init__.py index 96969c2..e69de29 100644 --- a/bert_squeeze/utils/schedulers/__init__.py +++ b/bert_squeeze/utils/schedulers/__init__.py @@ -1 +0,0 @@ -from .reduce_on_plateau import GroupCompatibleReduceLROnPlateau diff --git a/bert_squeeze/utils/schedulers/reduce_on_plateau.py b/bert_squeeze/utils/schedulers/reduce_on_plateau.py deleted file mode 100644 index 55777fc..0000000 --- a/bert_squeeze/utils/schedulers/reduce_on_plateau.py +++ /dev/null @@ -1,34 +0,0 @@ -from __future__ import annotations - -from typing import Union, cast - -from overrides import overrides -from torch.optim.lr_scheduler import ReduceLROnPlateau - -__all__ = ["GroupCompatibleReduceLROnPlateau"] - - -class GroupCompatibleReduceLROnPlateau(ReduceLROnPlateau): - @overrides - def load_state_dict(self, state_dict: dict[str, object]) -> None: - migrated_state = dict(state_dict) - migrated_state["min_lrs"] = self._migrated_min_lrs(state_dict.get("min_lrs")) - migrated_state["_last_lr"] = [ - float(group["lr"]) for group in self.optimizer.param_groups - ] - super().load_state_dict(migrated_state) - - def _migrated_min_lrs(self, saved_min_lrs: object) -> list[float]: - current_min_lrs = [float(value) for value in self.min_lrs] - if not isinstance(saved_min_lrs, list) or not all( - isinstance(value, (int, float)) for value in saved_min_lrs - ): - return current_min_lrs - - min_lrs = cast(list[Union[int, float]], saved_min_lrs) - group_count = len(self.optimizer.param_groups) - if len(min_lrs) == group_count: - return [float(value) for value in min_lrs] - if min_lrs and all(value == min_lrs[0] for value in min_lrs): - return [float(min_lrs[0])] * group_count - return current_min_lrs diff --git a/bert_squeeze/utils/types.py b/bert_squeeze/utils/types.py index 5ecc653..9bd8112 100644 --- a/bert_squeeze/utils/types.py +++ b/bert_squeeze/utils/types.py @@ -26,24 +26,22 @@ def __getitem__(self, item): @dataclass class DeeBertEncoderOutput: exit_layer: int - last_hidden_state: Optional[torch.FloatTensor] = None - hidden_states: Optional[Tuple[torch.FloatTensor]] = None - attentions: Optional[Tuple[torch.FloatTensor]] = None - ramps_exit: Optional[Tuple[RampOutput]] = None - # Optional per-layer gate logits/probs for BERxiT - gates_logits: Optional[Tuple[torch.FloatTensor]] = None + last_hidden_state: Optional[torch.Tensor] = None + hidden_states: Optional[Tuple[torch.Tensor, ...]] = None + attentions: Optional[Tuple[torch.Tensor, ...]] = None + ramps_exit: Optional[Tuple[RampOutput, ...]] = None + gates_logits: Optional[Tuple[torch.Tensor, ...]] = None @dataclass class DeeBertModelOutput: exit_layer: int - sequence_output: Optional[torch.FloatTensor] = None - pooled_output: Optional[torch.FloatTensor] = None - hidden_states: Optional[torch.FloatTensor] = None - attentions: Optional[torch.FloatTensor] = None + sequence_output: Optional[torch.Tensor] = None + pooled_output: Optional[torch.Tensor] = None + hidden_states: Optional[Tuple[torch.Tensor, ...]] = None + attentions: Optional[Tuple[torch.Tensor, ...]] = None ramps_exits: Optional[Sequence[RampOutput]] = None - # Optional per-layer gate logits/probs for BERxiT - gates_logits: Optional[Tuple[torch.FloatTensor]] = None + gates_logits: Optional[Tuple[torch.Tensor, ...]] = None @property def logits(self) -> torch.Tensor: diff --git a/tests/test_optimizer_parameter_groups.py b/tests/test_optimizer_parameter_groups.py index b11ccd2..4d0dafb 100644 --- a/tests/test_optimizer_parameter_groups.py +++ b/tests/test_optimizer_parameter_groups.py @@ -1,15 +1,23 @@ from __future__ import annotations from pathlib import Path -from typing import Optional, Union +from typing import Optional import pytest import torch import torch.nn as nn +from lightning.pytorch import Trainer from omegaconf import DictConfig, OmegaConf -from torch.optim import AdamW -from torch.optim.lr_scheduler import ReduceLROnPlateau -from transformers import BertConfig, T5Config, T5ForConditionalGeneration +from torch.utils.data import DataLoader +from transformers import ( + BertConfig, + GPT2Config, + GPT2LMHeadModel, + LlamaConfig, + LlamaForCausalLM, + T5Config, + T5ForConditionalGeneration, +) from bert_squeeze.models.custom_transformers.berxit import BerxitModel from bert_squeeze.models.custom_transformers.deebert import DeeBertModel @@ -18,11 +26,8 @@ from bert_squeeze.utils.optimizers import ( OptimizerParameterGroup, build_optimizer_parameter_groups, - register_legacy_optimizer_state_migration, ) -from bert_squeeze.utils.schedulers import GroupCompatibleReduceLROnPlateau - -_NO_DECAY_NAMES = ("bias", "gamma", "beta", "LayerNorm.weight", "layer_norm.weight") +from bert_squeeze.utils.types import RampOutput class _Encoder(nn.Module): @@ -34,12 +39,13 @@ def __init__(self, layer_count: int) -> None: class _LayeredModel(nn.Module): def __init__(self, layer_count: int) -> None: super().__init__() + self.embeddings = nn.Embedding(4, 2) self.encoder = _Encoder(layer_count) self.classifier = nn.Linear(2, 2) def _learning_rate_for( - groups: list[OptimizerParameterGroup], parameter: nn.Parameter + groups: list[OptimizerParameterGroup], parameter: torch.Tensor ) -> Optional[float]: for group in groups: if any(group_parameter is parameter for group_parameter in group["params"]): @@ -47,15 +53,23 @@ def _learning_rate_for( raise AssertionError("Parameter is missing from optimizer groups.") +def _grouped_parameter_ids( + groups: list[OptimizerParameterGroup], +) -> list[int]: + return [id(parameter) for group in groups for parameter in group["params"]] + + @pytest.mark.parametrize( - ("learning_rates", "expected_rates"), + ("learning_rates", "expected_rates", "embedding_rate"), [ - ([0.1], [0.025, 0.05, 0.1]), - ([0.01, 0.02, 0.03], [0.01, 0.02, 0.03]), + ([0.1], [0.025, 0.05, 0.1], 0.0125), + ([0.01, 0.02, 0.03], [0.01, 0.02, 0.03], 0.005), ], ) def test_optimizer_groups_follow_model_depth( - learning_rates: list[float], expected_rates: list[float] + learning_rates: list[float], + expected_rates: list[float], + embedding_rate: float, ) -> None: model = _LayeredModel(layer_count=3) @@ -67,17 +81,19 @@ def test_optimizer_groups_follow_model_depth( weight_decay=0.01, ) - actual_rates = [ + assert [ _learning_rate_for(groups, layer.weight) for layer in model.encoder.layer - ] - assert actual_rates == pytest.approx(expected_rates) - assert [groups[index]["lr"] for index in range(3)] == pytest.approx(expected_rates) - assert [groups[12 + index]["lr"] for index in range(3)] == pytest.approx( - expected_rates + ] == pytest.approx(expected_rates) + assert _learning_rate_for(groups, model.embeddings.weight) == pytest.approx( + embedding_rate ) - assert all("lr" not in group for group in groups[24:]) - assert _learning_rate_for(groups, model.classifier.weight) is None - assert sum(len(group["params"]) for group in groups) == len(list(model.parameters())) + assert _learning_rate_for(groups, model.classifier.weight) == pytest.approx( + learning_rates[-1] + ) + grouped_parameter_ids = _grouped_parameter_ids(groups) + assert len(grouped_parameter_ids) == len(set(grouped_parameter_ids)) + assert len(grouped_parameter_ids) == len(list(model.parameters())) + assert all(group["params"] for group in groups) def test_optimizer_groups_reject_mismatched_layer_rates() -> None: @@ -93,117 +109,37 @@ def test_optimizer_groups_reject_mismatched_layer_rates() -> None: ) -def _legacy_optimizer_groups(model: nn.Module) -> list[OptimizerParameterGroup]: - named_parameters = list(model.named_parameters()) - layer_keys = [f"layer.{index}." for index in range(12)] - legacy_rates = [0.1 * pow(0.5, 11 - index) for index in range(12)] - legacy_groups: list[OptimizerParameterGroup] = [] - for use_weight_decay in (True, False): - legacy_groups.extend( - OptimizerParameterGroup( - params=[ - parameter - for name, parameter in named_parameters - if layer_key in name - and (not any(no_decay in name for no_decay in _NO_DECAY_NAMES)) - == use_weight_decay - ], - weight_decay=0.01 if use_weight_decay else 0.0, - lr=legacy_rates[index], - ) - for index, layer_key in enumerate(layer_keys) - ) - for use_weight_decay in (True, False): - legacy_groups.append( - OptimizerParameterGroup( - params=[ - parameter - for name, parameter in named_parameters - if not any(layer_key in name for layer_key in layer_keys) - and (not any(no_decay in name for no_decay in _NO_DECAY_NAMES)) - == use_weight_decay - ], - weight_decay=0.01 if use_weight_decay else 0.0, - ) - ) - return legacy_groups - - -def _current_optimizer(model: nn.Module) -> AdamW: - optimizer = AdamW( - build_optimizer_parameter_groups( - model.named_parameters(), - discriminative_learning=True, - learning_rates=[0.1], - layer_lr_decay=0.5, - weight_decay=0.01, - ), - lr=0.1, - ) - register_legacy_optimizer_state_migration(optimizer, model.named_parameters()) - return optimizer - - -@pytest.mark.parametrize("layer_count", [3, 24]) -def test_optimizer_groups_restore_legacy_optimizer_state(layer_count: int) -> None: - model = _LayeredModel(layer_count=layer_count) - legacy_optimizer = AdamW( - _legacy_optimizer_groups(model), - lr=0.1, - betas=(0.8, 0.88), - eps=1e-6, - ) - sum(parameter.square().sum() for parameter in model.parameters()).backward() - legacy_optimizer.step() - legacy_optimizer.zero_grad() - - current_optimizer = _current_optimizer(model) - current_optimizer.load_state_dict(legacy_optimizer.state_dict()) - sum(parameter.square().sum() for parameter in model.parameters()).backward() - current_optimizer.step() - - assert all(torch.isfinite(parameter).all() for parameter in model.parameters()) - assert current_optimizer.param_groups[0]["betas"] == (0.8, 0.88) - assert current_optimizer.param_groups[0]["eps"] == 1e-6 - - -def test_scheduler_restores_after_optimizer_group_migration() -> None: - model = _LayeredModel(layer_count=24) - legacy_optimizer = AdamW(_legacy_optimizer_groups(model), lr=0.1) - legacy_scheduler = ReduceLROnPlateau(legacy_optimizer, factor=0.5, patience=0) - current_optimizer = _current_optimizer(model) - current_optimizer.load_state_dict(legacy_optimizer.state_dict()) - current_scheduler = GroupCompatibleReduceLROnPlateau( - current_optimizer, factor=0.5, patience=0 - ) - - current_scheduler.load_state_dict(legacy_scheduler.state_dict()) - current_scheduler.step(1.0) - current_scheduler.step(2.0) - - assert len(current_scheduler.min_lrs) == len(current_optimizer.param_groups) - - -def _t5_model(block_count: int) -> T5ForConditionalGeneration: +def _t5_model(encoder_blocks: int, decoder_blocks: int) -> T5ForConditionalGeneration: return T5ForConditionalGeneration( T5Config( vocab_size=32, d_model=16, d_ff=32, - num_layers=block_count, - num_decoder_layers=block_count, + num_layers=encoder_blocks, + num_decoder_layers=decoder_blocks, num_heads=2, ) ) -def test_t5_optimizer_groups_follow_block_depth() -> None: - model = _t5_model(block_count=2) +@pytest.mark.parametrize( + ("learning_rates", "expected_encoder_rates", "expected_decoder_rates"), + [ + ([0.1], [0.05, 0.1], [0.0125, 0.025, 0.05, 0.1]), + ([0.01, 0.02, 0.03, 0.04], [0.03, 0.04], [0.01, 0.02, 0.03, 0.04]), + ], +) +def test_encoder_decoder_stacks_receive_independent_schedules( + learning_rates: list[float], + expected_encoder_rates: list[float], + expected_decoder_rates: list[float], +) -> None: + model = _t5_model(encoder_blocks=2, decoder_blocks=4) groups = build_optimizer_parameter_groups( model.named_parameters(), discriminative_learning=True, - learning_rates=[0.1], + learning_rates=learning_rates, layer_lr_decay=0.5, weight_decay=0.01, ) @@ -216,27 +152,63 @@ def test_t5_optimizer_groups_follow_block_depth() -> None: _learning_rate_for(groups, block.layer[0].SelfAttention.q.weight) for block in model.decoder.block ] - grouped_parameters = [parameter for group in groups for parameter in group["params"]] + assert encoder_rates == pytest.approx(expected_encoder_rates) + assert decoder_rates == pytest.approx(expected_decoder_rates) - assert encoder_rates == pytest.approx([0.05, 0.1]) - assert decoder_rates == pytest.approx([0.05, 0.1]) - assert len(grouped_parameters) == len( - {id(parameter) for parameter in grouped_parameters} - ) - assert len(grouped_parameters) == len(list(model.parameters())) +@pytest.mark.parametrize( + "architecture", + ["gpt2", "llama"], +) +def test_normalization_parameters_do_not_use_weight_decay( + architecture: str, +) -> None: + if architecture == "gpt2": + model = GPT2LMHeadModel( + GPT2Config( + n_layer=2, + n_head=2, + n_embd=16, + n_positions=16, + vocab_size=32, + ) + ) + else: + model = LlamaForCausalLM( + LlamaConfig( + hidden_size=16, + intermediate_size=32, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + vocab_size=32, + max_position_embeddings=16, + ) + ) -def test_t5_legacy_state_migrates_when_group_counts_match() -> None: - model = _t5_model(block_count=12) - legacy_optimizer = AdamW(_legacy_optimizer_groups(model), lr=0.1) - current_optimizer = _current_optimizer(model) - - assert len(legacy_optimizer.param_groups) == len(current_optimizer.param_groups) - current_optimizer.load_state_dict(legacy_optimizer.state_dict()) + groups = build_optimizer_parameter_groups( + model.named_parameters(), + discriminative_learning=True, + learning_rates=[0.1], + layer_lr_decay=0.5, + weight_decay=0.01, + ) + weight_decay_by_parameter = { + id(parameter): group["weight_decay"] + for group in groups + for parameter in group["params"] + } + normalization_parameters = [ + parameter + for name, parameter in model.named_parameters() + if "ln_" in name or "layernorm" in name or name.endswith("norm.weight") + ] - sum(parameter.square().sum() for parameter in model.parameters()).backward() - current_optimizer.step() - assert all(torch.isfinite(parameter).all() for parameter in model.parameters()) + assert normalization_parameters + assert all( + weight_decay_by_parameter[id(parameter)] == 0.0 + for parameter in normalization_parameters + ) def _model_config(tmp_path: Path) -> BertConfig: @@ -252,74 +224,182 @@ def _model_config(tmp_path: Path) -> BertConfig: return model_config -def _ramp_training_config(**overrides: object) -> DictConfig: +def _training_config(**overrides: object) -> DictConfig: config = { "logging_steps": 2, "accumulation_steps": 1, "objective": "ce", "lr_scheduler": False, + "optimizer": "adamw", + "adam_eps": 1e-8, "discriminative_learning": True, "learning_rates": [0.1], "layer_lr_decay": 0.5, "weight_decay": 0.01, "train_highway": True, "train_gates": True, + "train_stage": "backbone", "early_exit_entropy": -1.0, + "gate_thresholds": 0.5, } config.update(overrides) return OmegaConf.create(config) -def _assert_parameters_are_grouped( - module: Union[LtDeeBert, LtBerxit], parameter_marker: str -) -> None: - groups = module._get_optimizer_parameters() - grouped_parameter_ids = { - id(parameter) for group in groups for parameter in group["params"] - } - expected_parameter_ids = { - id(parameter) - for name, parameter in module.named_parameters() - if parameter_marker in name +def _batch() -> dict[str, torch.Tensor]: + return { + "input_ids": torch.tensor([[1, 2, 3], [3, 2, 1]]), + "attention_mask": torch.ones(2, 3, dtype=torch.long), + "labels": torch.tensor([0, 1]), } - assert expected_parameter_ids - assert expected_parameter_ids <= grouped_parameter_ids - -def test_deebert_shipped_config_includes_ramp_parameters(tmp_path: Path) -> None: +@pytest.mark.parametrize("train_highway", [False, True]) +def test_deebert_training_stages_update_the_inference_exits( + tmp_path: Path, train_highway: bool +) -> None: model_config = _model_config(tmp_path) module = LtDeeBert( - training_config=_ramp_training_config(), + training_config=_training_config(train_highway=train_highway), pretrained_model=str(tmp_path), num_labels=2, model=DeeBertModel(model_config), ) - _assert_parameters_are_grouped(module, ".ramp.") - + output = module._classification_output(_batch()) + module._classification_loss(output, _batch()["labels"]).backward() + grouped_ids = set(_grouped_parameter_ids(module._get_optimizer_parameters())) + trainable_ids = { + id(parameter) for parameter in module.parameters() if parameter.requires_grad + } -def test_berxit_shipped_config_includes_ramp_parameters(tmp_path: Path) -> None: + backbone_grad = module.bert.encoder.layer[0].attention.self.query.weight.grad + intermediate_grad = module.bert.encoder.ramp[0].classifier.weight.grad + final_grad = module.bert.encoder.ramp[-1].classifier.weight.grad + if train_highway: + assert backbone_grad is None + assert intermediate_grad is not None + assert final_grad is None + else: + assert backbone_grad is not None + assert intermediate_grad is None + assert final_grad is not None + assert grouped_ids == trainable_ids + + +def _berxit_module(tmp_path: Path, **overrides: object) -> LtBerxit: model_config = _model_config(tmp_path) - module = LtBerxit( - training_config=_ramp_training_config(), + return LtBerxit( + training_config=_training_config(**overrides), pretrained_model=str(tmp_path), num_labels=2, model=BerxitModel(model_config), ) - _assert_parameters_are_grouped(module, ".ramp.") +def test_berxit_gate_stage_only_trains_the_shared_gate(tmp_path: Path) -> None: + module = _berxit_module(tmp_path, train_stage="gates") -def test_berxit_non_discriminative_training_includes_gate_parameters( - tmp_path: Path, -) -> None: + trainable_parameters = { + name for name, parameter in module.named_parameters() if parameter.requires_grad + } + grouped_ids = set(_grouped_parameter_ids(module._get_optimizer_parameters())) + + assert trainable_parameters + assert all("gates" in name for name in trainable_parameters) + assert grouped_ids == { + id(parameter) for parameter in module.parameters() if parameter.requires_grad + } + + +def test_berxit_trains_final_and_intermediate_exits(tmp_path: Path) -> None: + module = _berxit_module(tmp_path) + batch = _batch() + output = module._classification_output(batch) + + module.loss( + labels=batch["labels"], + ramps_exits=output.ramps_exits, + train_ramps=False, + ).backward() + assert module.bert.encoder.ramp[0].classifier.weight.grad is None + assert module.bert.encoder.ramp[-1].classifier.weight.grad is not None + + module.zero_grad(set_to_none=True) + output = module._classification_output(batch) + module.loss( + labels=batch["labels"], + ramps_exits=output.ramps_exits, + train_ramps=True, + ).backward() + assert module.bert.encoder.ramp[0].classifier.weight.grad is not None + assert module.bert.encoder.ramp[-1].classifier.weight.grad is not None + + +def test_berxit_uses_label_based_certainty_targets(tmp_path: Path) -> None: + module = _berxit_module(tmp_path) + model_config = module.model_config + labels = torch.tensor([0]) + pooled_output = torch.zeros(1, model_config.hidden_size) + correct_early_exit = RampOutput( + logits=torch.tensor([[4.0, -4.0]]), pooled_output=pooled_output + ) + wrong_final_exit = RampOutput( + logits=torch.tensor([[-4.0, 4.0]]), pooled_output=pooled_output + ) + early_gate = torch.tensor([[2.0]], requires_grad=True) + final_gate = torch.tensor([[2.0]], requires_grad=True) + + loss = module.loss( + labels=labels, + ramps_exits=[correct_early_exit, wrong_final_exit], + gates_logits=(early_gate, final_gate), + train_ramps=True, + train_gates=True, + ) + loss.backward() + + assert early_gate.grad is not None and early_gate.grad.item() < 0 + assert final_gate.grad is not None and final_gate.grad.item() > 0 + assert isinstance(module.bert.encoder.gates, nn.Linear) + + +def test_berxit_inference_returns_every_sample_after_early_exit(tmp_path: Path) -> None: + module = _berxit_module(tmp_path, gate_thresholds=0.0) + module.bert.set_inference_mode(inference=True) + batch = _batch() + + logits, ramps_exits, exit_layer, _ = module.forward( + input_ids=batch["input_ids"], + attention_mask=batch["attention_mask"], + ) + + assert logits.shape == (2, 2) + assert len(ramps_exits) == 2 + assert exit_layer == 0 + + +def test_plateau_scheduler_uses_epoch_training_loss(tmp_path: Path) -> None: model_config = _model_config(tmp_path) - module = LtBerxit( - training_config=_ramp_training_config(discriminative_learning=False), + module = LtDeeBert( + training_config=_training_config( + train_highway=False, + lr_scheduler=True, + ), pretrained_model=str(tmp_path), num_labels=2, - model=BerxitModel(model_config), + model=DeeBertModel(model_config), ) + trainer = Trainer( + accelerator="cpu", + devices=1, + enable_checkpointing=False, + enable_model_summary=False, + enable_progress_bar=False, + logger=False, + max_epochs=1, + ) + + trainer.fit(module, train_dataloaders=DataLoader([_batch()], batch_size=None)) - _assert_parameters_are_grouped(module, ".gates.") + assert "train/epoch_loss" in trainer.callback_metrics diff --git a/tests/test_seq2seq_distillation_training.py b/tests/test_seq2seq_distillation_training.py index 103422e..50924dd 100644 --- a/tests/test_seq2seq_distillation_training.py +++ b/tests/test_seq2seq_distillation_training.py @@ -49,7 +49,7 @@ def _training_config() -> DictConfig: "discriminative_learning": False, "learning_rates": [0.1], "logging_steps": 10, - "lr_scheduler": False, + "lr_scheduler": True, "optimizer": "sgd", "weight_decay": 0.0, } @@ -87,5 +87,6 @@ def test_seq2seq_distillation_trains_student_with_synthetic_batches(): ) assert trainer.global_step == 2 + assert "train/epoch_loss" in trainer.callback_metrics assert not torch.equal(student.classifier.weight, initial_weights) assert all(parameter.grad is None for parameter in teacher.parameters())