diff --git a/bert_squeeze/assistants/configs/train_adapter.yaml b/bert_squeeze/assistants/configs/train_adapter.yaml index 12fdd8b..dc05f42 100644 --- a/bert_squeeze/assistants/configs/train_adapter.yaml +++ b/bert_squeeze/assistants/configs/train_adapter.yaml @@ -5,6 +5,7 @@ general: get_mismatched: true evaluate_during_training: true labels: [ 0, 1 ] + num_labels: 2 output_dir: outputs save_steps: 500 validation_every_n_epoch: 1 @@ -34,7 +35,7 @@ model: training_config: ${train} task_name: adapter_config_name: "seq_bn" - labels: [ "0", "1" ] + labels: ${general.labels} scorer: _target_: bert_squeeze.utils.scorers.sequence_classification_scorer.BaseSequenceClassificationScorer labels: ${general.labels} diff --git a/bert_squeeze/assistants/configs/train_bert.yaml b/bert_squeeze/assistants/configs/train_bert.yaml index 2442061..105a66c 100644 --- a/bert_squeeze/assistants/configs/train_bert.yaml +++ b/bert_squeeze/assistants/configs/train_bert.yaml @@ -5,6 +5,7 @@ general: get_mismatched: true evaluate_during_training: true labels: [ 0, 1 ] + num_labels: 2 output_dir: outputs save_steps: 500 validation_every_n_epoch: 1 @@ -30,7 +31,7 @@ train: model: _target_: bert_squeeze.models.lt_bert.LtSequenceClassificationCustomBert - num_labels: 2 + num_labels: ${general.num_labels} pretrained_model: "bert-base-cased" training_config: ${train} scorer: diff --git a/bert_squeeze/assistants/configs/train_fastbert.yaml b/bert_squeeze/assistants/configs/train_fastbert.yaml index 30a1bd5..f307f87 100644 --- a/bert_squeeze/assistants/configs/train_fastbert.yaml +++ b/bert_squeeze/assistants/configs/train_fastbert.yaml @@ -5,6 +5,7 @@ general: get_mismatched: true evaluate_during_training: true labels: [ 0, 1 ] + num_labels: 2 output_dir: outputs save_steps: 500 validation_every_n_epoch: 1 @@ -35,7 +36,7 @@ model: _target_: bert_squeeze.models.lt_fastbert.LtFastBert training_config: ${train} pretrained_model: "bert-base-cased" - num_labels: 2 + num_labels: ${general.num_labels} scorer_type: "fast" scorer: _target_: bert_squeeze.utils.scorers.sequence_classification_scorer.FastBertSequenceClassificationScorer @@ -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/assistants/configs/train_theseus_bert.yaml b/bert_squeeze/assistants/configs/train_theseus_bert.yaml index 7578b3a..49326ed 100644 --- a/bert_squeeze/assistants/configs/train_theseus_bert.yaml +++ b/bert_squeeze/assistants/configs/train_theseus_bert.yaml @@ -5,6 +5,7 @@ general: get_mismatched: true evaluate_during_training: true labels: [ 0, 1 ] + num_labels: 2 output_dir: outputs save_steps: 500 validation_every_n_epoch: 1 @@ -32,7 +33,7 @@ model: _target_: bert_squeeze.models.lt_theseus_bert.LtTheseusBert training_config: ${train} pretrained_model: "bert-base-cased" - num_labels: 2 + num_labels: ${general.num_labels} replacement_scheduler: type: "linear" base_replacing_rate: 0.3 diff --git a/bert_squeeze/assistants/train_assistant.py b/bert_squeeze/assistants/train_assistant.py index 2bd9d07..ce877ab 100644 --- a/bert_squeeze/assistants/train_assistant.py +++ b/bert_squeeze/assistants/train_assistant.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from copy import deepcopy from importlib import resources from typing import Dict, List, Optional @@ -72,8 +74,11 @@ def __init__( f" following: {CONFIG_MAPPER.keys()}" ) - config_path = resources.files("bert_squeeze").joinpath( - "assistants/configs", config_name + config_path = ( + resources.files("bert_squeeze") + .joinpath("assistants") + .joinpath("configs") + .joinpath(config_name) ) with resources.as_file(config_path) as resolved_path: conf = OmegaConf.load(resolved_path) @@ -102,6 +107,14 @@ def __init__( overrides if base is None else deep_update(base, overrides) ) + labels = conf["general"].get("labels") + if labels is not None: + num_labels = len(labels) + configured_num_labels = conf["general"].get("num_labels") + if configured_num_labels is not None and configured_num_labels != num_labels: + raise ValueError("general.num_labels must match the number of labels.") + conf["general"]["num_labels"] = num_labels + self.name = name self.general = conf["general"] self.train = conf["train"] @@ -112,8 +125,8 @@ def __init__( self._model: Optional[pl.LightningModule] = None self._data: Optional[pl.LightningDataModule] = None - self._logger = None - self._callbacks = None + self._logger: Optional[Logger] = None + self._callbacks: Optional[List[Callback]] = None @property def model(self) -> pl.LightningModule: @@ -162,11 +175,11 @@ def callbacks(self) -> List[Callback]: """""" if self._callbacks is None: if self._callbacks_conf is not None: - self.callbacks = [ + self._callbacks = [ instantiate(callback) for callback in self._callbacks_conf ] else: - self.callbacks = [] + self._callbacks = [] return self._callbacks @callbacks.setter diff --git a/bert_squeeze/models/custom_transformers/deebert.py b/bert_squeeze/models/custom_transformers/deebert.py index e85ff58..7677418 100644 --- a/bert_squeeze/models/custom_transformers/deebert.py +++ b/bert_squeeze/models/custom_transformers/deebert.py @@ -1,7 +1,9 @@ # This is heavily inspired by the following repo: # https://github.com/castorini/DeeBERT +from __future__ import annotations + from abc import ABC -from typing import List, Union +from typing import List, Optional, Tuple, Union import torch import torch.nn as nn @@ -66,15 +68,17 @@ class DeeBertEncoder(nn.Module): def __init__(self, config: PretrainedConfig, inference: bool): super(DeeBertEncoder, self).__init__() self.config = config - self.layer = nn.ModuleList([BertLayer(config)] * config.num_hidden_layers) - self.ramp = nn.ModuleList([OffRamp(config)] * config.num_hidden_layers) + 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)] + ) - self.early_exit_entropy = [ - -1, - ] * config.num_hidden_layers + self.early_exit_entropy: List[float] = [-1.0] * config.num_hidden_layers self.inference = inference - def set_early_exit_entropy(self, x: Union[List[float], float]) -> None: + def set_early_exit_entropy(self, x: Union[List[float], float, int]) -> None: """ Assigning an entropy threshold to every layer. @@ -85,9 +89,9 @@ def set_early_exit_entropy(self, x: Union[List[float], float]) -> None: """ if isinstance(x, float) or isinstance(x, int): for i in range(self.config.num_hidden_layers): - self.early_exit_entropy[i] = x + self.early_exit_entropy[i] = float(x) elif isinstance(x, list): - self.early_exit_entropy = x + self.early_exit_entropy = list(x) else: raise TypeError( f"Expected 'x' to be of type 'float' or 'list' but got :'{type(x)}'" @@ -117,11 +121,11 @@ def forward( 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: Tuple[torch.Tensor, ...] = tuple() + all_attentions: Tuple[torch.Tensor, ...] = tuple() if not self.inference: - all_ramps = tuple() + all_ramps: Tuple[RampOutput, ...] = tuple() for i, layer_module in enumerate(self.layer): if output_hidden_states: @@ -149,16 +153,18 @@ def forward( return DeeBertEncoderOutput( last_hidden_state=hidden_states, - hidden_states=all_hidden_states, - attentions=all_attentions, + hidden_states=all_hidden_states if output_hidden_states else None, + attentions=all_attentions if output_attentions else None, ramps_exit=all_ramps, exit_layer=i, ) else: - all_ramps = [ - 0, - ] * hidden_states.shape[0] - positions = torch.arange(start=0, end=hidden_states.shape[0]).long() + batch_ramps: List[Optional[RampOutput]] = [None] * hidden_states.shape[0] + positions = torch.arange( + start=0, + end=hidden_states.shape[0], + device=hidden_states.device, + ).long() for i, layer_module in enumerate(self.layer): layer_outputs = layer_module( @@ -174,21 +180,35 @@ def forward( if i == len(self.layer) - 1: for idx, pos in enumerate(positions): - all_ramps[pos] = ramp_exit[idx] + batch_ramps[int(pos)] = ramp_exit[idx] 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] + batch_ramps[int(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, exit_layer=i) - return DeeBertEncoderOutput(ramps_exit=all_ramps, exit_layer=i) + return DeeBertEncoderOutput( + ramps_exit=self._completed_ramps(batch_ramps), + exit_layer=i, + ) + return DeeBertEncoderOutput( + ramps_exit=self._completed_ramps(batch_ramps), + exit_layer=i, + ) + + @staticmethod + def _completed_ramps( + ramps: List[Optional[RampOutput]], + ) -> Tuple[RampOutput, ...]: + if any(ramp is None for ramp in ramps): + raise RuntimeError("DeeBERT did not produce an output for every sample.") + return tuple(ramp for ramp in ramps if ramp is not None) class DeeBertModel(BertPreTrainedModel, ABC): diff --git a/bert_squeeze/models/custom_transformers/fastbert.py b/bert_squeeze/models/custom_transformers/fastbert.py index 18cfe47..8ebac75 100644 --- a/bert_squeeze/models/custom_transformers/fastbert.py +++ b/bert_squeeze/models/custom_transformers/fastbert.py @@ -2,6 +2,8 @@ # The main difference relies on the fact that I'm trying to use HuggingFace's # 'transformers' components as much as possible. +from __future__ import annotations + from typing import List, Tuple, Union import torch @@ -129,7 +131,9 @@ def forward( if inference: # positions will keep track of the original position of each element in the # batch when elements will be removed - final_probs = torch.zeros((hidden_states[0].shape[0], 2), device=device) + final_probs = hidden_states[0].new_zeros( + (hidden_states[0].shape[0], self.config.num_labels) + ) positions = torch.arange( start=0, end=hidden_states[0].shape[0], device=device ).long() diff --git a/bert_squeeze/models/lt_berxit.py b/bert_squeeze/models/lt_berxit.py index b66c4b0..0d0f969 100644 --- a/bert_squeeze/models/lt_berxit.py +++ b/bert_squeeze/models/lt_berxit.py @@ -1,5 +1,7 @@ +from __future__ import annotations + import logging -from typing import Dict, List, Optional, Tuple, Union +from typing import Dict, List, Optional, Sequence, Tuple, Union import lightning.pytorch as pl import torch @@ -11,6 +13,7 @@ from transformers import AutoConfig from bert_squeeze.utils.scorers import Scorer +from bert_squeeze.utils.types import RampOutput from .base_lt_module import BaseSequenceClassificationTransformerModule from .custom_transformers.berxit import BerxitModel @@ -67,7 +70,12 @@ def forward( position_ids: torch.Tensor = None, head_mask: torch.Tensor = None, **kwargs, - ) -> Tuple[torch.Tensor, Tuple[torch.Tensor], int, Optional[Tuple[torch.Tensor]]]: + ) -> Tuple[ + torch.Tensor, + Sequence[RampOutput], + int, + Optional[Tuple[torch.Tensor, ...]], + ]: outputs = self.bert( input_ids, attention_mask=attention_mask, @@ -76,7 +84,7 @@ def forward( head_mask=head_mask, ) - if self.training: + if not self.bert.encoder.inference: exit_layer = self.num_layers pooled_output = outputs.pooled_output pooled_output = self.dropout(pooled_output) @@ -153,6 +161,7 @@ def on_validation_epoch_end(self) -> None: self.log_eval_report(labels_probs) self.valid_scorer.reset() + self.validation_step_outputs.clear() @overrides def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: @@ -161,9 +170,14 @@ def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: "attention_mask": batch["attention_mask"], "token_type_ids": batch["token_type_ids"], } - logits, _, _, _ = self.forward(**inputs) + logits, ramps_exits, _, gates_logits = self.forward(**inputs) loss = self.loss( - logits=logits, labels=batch["labels"], train_ramps=self.train_highway + logits=logits, + labels=batch["labels"], + ramps_exits=ramps_exits, + train_ramps=self.train_highway, + train_gates=self.train_gates, + gates_logits=gates_logits, ) self.test_scorer.add(logits.cpu(), batch["labels"].cpu(), loss.cpu()) self.test_step_outputs.append( @@ -370,17 +384,19 @@ def _get_optimizer_parameters(self) -> List[Dict]: def loss( self, labels: torch.Tensor, - logits: torch.Tensor = None, - ramps_exits: Tuple[torch.Tensor] = None, + logits: Optional[torch.Tensor] = None, + ramps_exits: Optional[Sequence[RampOutput]] = None, train_ramps: bool = False, train_gates: bool = False, - gates_logits: Optional[Tuple[torch.Tensor]] = None, + gates_logits: Optional[Tuple[torch.Tensor, ...]] = None, *args, **kwargs, ) -> torch.Tensor: # Same ramp loss mechanics as LtDeeBert for consistency if train_ramps: - ramps_losses = [] + 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() @@ -388,25 +404,31 @@ def loss( ramps_logits.view(-1, self.model_config.num_labels), labels.view(-1) ) ramps_losses.append(ramps_loss) - loss = sum(ramps_losses) + loss = torch.stack(ramps_losses).sum() else: + if logits is None: + raise ValueError( + "Classifier logits are required when ramps are disabled." + ) loss_fct = CrossEntropyLoss() loss = loss_fct( logits.view(-1, self.model_config.num_labels), labels.view(-1) ) # Optional: add gate loss using pseudo-labels from final ramp - if train_gates and gates_logits is not None: + 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 = [] + 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 + sum(gate_losses) + loss = loss + torch.stack(gate_losses).sum() return loss def _build_model(self): diff --git a/bert_squeeze/models/lt_deebert.py b/bert_squeeze/models/lt_deebert.py index a48c3a6..259277f 100644 --- a/bert_squeeze/models/lt_deebert.py +++ b/bert_squeeze/models/lt_deebert.py @@ -1,5 +1,7 @@ +from __future__ import annotations + import logging -from typing import Dict, List, Optional, Tuple, Union +from typing import Dict, List, Optional, Sequence, Tuple, Union import lightning.pytorch as pl import torch @@ -11,6 +13,7 @@ from transformers import AutoConfig from bert_squeeze.utils.scorers import Scorer +from bert_squeeze.utils.types import RampOutput from .base_lt_module import BaseSequenceClassificationTransformerModule from .custom_transformers.deebert import DeeBertModel @@ -66,7 +69,7 @@ def forward( position_ids: torch.Tensor = None, head_mask: torch.Tensor = None, **kwargs, - ) -> Tuple[torch.Tensor, Tuple[torch.Tensor], int]: + ) -> 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. @@ -104,7 +107,7 @@ def forward( head_mask=head_mask, ) - if self.training: + if not self.bert.encoder.inference: exit_layer = self.num_layers pooled_output = outputs.pooled_output pooled_output = self.dropout(pooled_output) @@ -125,7 +128,7 @@ def training_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: "attention_mask": batch["attention_mask"], "token_type_ids": batch["token_type_ids"], } - logits, ramps_exits, exit_layer = self.forward(**inputs) + logits, ramps_exits, _ = self.forward(**inputs) loss = self.loss( logits=logits, labels=batch["labels"], @@ -155,7 +158,7 @@ def validation_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: "attention_mask": batch["attention_mask"], "token_type_ids": batch["token_type_ids"], } - logits, ramps_exits, exit_layer = self.forward(**inputs) + logits, ramps_exits, _ = self.forward(**inputs) loss = self.loss( logits=logits, labels=batch["labels"], @@ -176,6 +179,7 @@ def on_validation_epoch_end(self) -> None: self.log_eval_report(labels_probs) self.valid_scorer.reset() + self.validation_step_outputs.clear() @overrides def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: @@ -185,9 +189,12 @@ def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: "attention_mask": batch["attention_mask"], "token_type_ids": batch["token_type_ids"], } - logits, ramps_exits, exit_layer = self.forward(**inputs) + logits, ramps_exits, _ = self.forward(**inputs) loss = self.loss( - logits=logits, labels=batch["labels"], train_ramps=self.train_highway + logits=logits, + labels=batch["labels"], + ramps_exits=ramps_exits, + train_ramps=self.train_highway, ) self.test_scorer.add(logits.cpu(), batch["labels"].cpu(), loss.cpu()) self.test_step_outputs.append( @@ -199,6 +206,7 @@ def on_test_epoch_end(self) -> None: """""" logging.info(self.test_scorer.get_table()) self.test_scorer.reset() + self.test_step_outputs.clear() def predict_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: """""" @@ -345,8 +353,8 @@ def _get_optimizer_parameters(self) -> List[Dict]: def loss( self, labels: torch.Tensor, - logits: torch.Tensor = None, - ramps_exits: Tuple[torch.Tensor] = None, + logits: Optional[torch.Tensor] = None, + ramps_exits: Optional[Sequence[RampOutput]] = None, train_ramps: bool = False, *args, **kwargs, @@ -371,7 +379,9 @@ def loss( """ # We want to fine-tune each individual ramp if train_ramps: - ramps_losses = [] + 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 @@ -382,8 +392,12 @@ def loss( ) ramps_losses.append(ramps_loss) - loss = sum(ramps_losses) + 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( diff --git a/bert_squeeze/models/lt_distilbert.py b/bert_squeeze/models/lt_distilbert.py index dc0907e..8905c4f 100644 --- a/bert_squeeze/models/lt_distilbert.py +++ b/bert_squeeze/models/lt_distilbert.py @@ -1,13 +1,9 @@ -from typing import Optional, Tuple, Union +from __future__ import annotations + +from typing import Tuple, Union -import lightning.pytorch as pl import torch -import torch.nn as nn -from omegaconf import DictConfig from overrides import overrides -from transformers import AutoModel - -from bert_squeeze.utils.scorers import Scorer from .base_lt_module import BaseSequenceClassificationTransformerModule @@ -29,20 +25,6 @@ class LtCustomDistilBert(BaseSequenceClassificationTransformerModule): helper object to compute performance metrics during training """ - def __init__( - self, - training_config: DictConfig, - pretrained_model: str, - num_labels: int, - model: Optional[Union[pl.LightningModule, nn.Module]] = None, - scorer: Scorer = None, - **kwargs, - ): - super().__init__( - training_config, pretrained_model, num_labels, model, scorer, **kwargs - ) - self._build_model() - @overrides def forward( self, @@ -50,7 +32,7 @@ def forward( attention_mask: torch.Tensor = None, output_attentions: bool = False, **kwargs, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + ) -> Union[torch.Tensor, Tuple[torch.Tensor, Tuple[torch.Tensor, ...]]]: """ Args: input_ids (torch.Tensor): @@ -61,29 +43,42 @@ def forward( output_attentions (bool): whether to output attention scores. Returns: - Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: logits obtained from model pass - along with the attention scores if `output_attentions=True`. + Logits, optionally paired with per-layer attention tensors. """ - outputs = self.encoder( - input_ids, attention_mask=attention_mask, output_attentions=output_attentions + kwargs.pop("return_dict", None) + outputs = self.model( + input_ids=input_ids, + attention_mask=attention_mask, + output_attentions=output_attentions, + return_dict=True, + **kwargs, ) - hidden_state = outputs[0] - logits = self.classifier(hidden_state) + logits = getattr(outputs, "logits", None) + if not isinstance(logits, torch.Tensor): + raise TypeError("DistilBERT sequence classifiers must return tensor logits.") if output_attentions: - return logits, outputs.attentions + attentions = getattr(outputs, "attentions", None) + if not isinstance(attentions, tuple): + raise TypeError( + "DistilBERT did not return attentions when they were requested." + ) + return logits, attentions return logits + def _classification_logits( + self, input_ids: torch.Tensor, attention_mask: torch.Tensor + ) -> torch.Tensor: + outputs = self.forward(input_ids=input_ids, attention_mask=attention_mask) + if not isinstance(outputs, torch.Tensor): + raise TypeError("DistilBERT classification must return tensor logits.") + return outputs + @overrides def training_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: """""" - inputs = { - "input_ids": batch["input_ids"], - "attention_mask": batch["attention_mask"], - } - - logits = self.forward(**inputs) - loss = self.loss(logits, batch["labels"]) + logits = self._classification_logits(batch["input_ids"], batch["attention_mask"]) + loss = self.loss(labels=batch["labels"], logits=logits) self.scorer.add(logits.detach().cpu(), batch["labels"], loss.detach().cpu()) if self.global_step > 0 and self.global_step % self.config.logging_steps == 0: @@ -100,13 +95,8 @@ def training_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: @overrides def validation_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: """""" - inputs = { - "input_ids": batch["input_ids"], - "attention_mask": batch["attention_mask"], - } - - logits = self.forward(**inputs) - loss = self.loss(logits, batch["labels"]) + logits = self._classification_logits(batch["input_ids"], batch["attention_mask"]) + loss = self.loss(labels=batch["labels"], logits=logits) self.valid_scorer.add(logits.cpu(), batch["labels"].cpu(), loss.cpu()) self.validation_step_outputs.append( @@ -117,27 +107,11 @@ def validation_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: @overrides def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: """""" - inputs = { - "input_ids": batch["input_ids"], - "attention_mask": batch["attention_mask"], - } - - logits = self.forward(**inputs) - loss = self.loss(logits, batch["labels"]) + logits = self._classification_logits(batch["input_ids"], batch["attention_mask"]) + loss = self.loss(labels=batch["labels"], logits=logits) self.test_scorer.add(logits.cpu(), batch["labels"].cpu(), loss.cpu()) self.test_step_outputs.append( {"loss": loss, "logits": logits.cpu(), "labels": batch["labels"].cpu()} ) return loss - - def _build_model(self): - """""" - self.encoder = AutoModel.from_pretrained(self.pretrained_model) - self.classifier = torch.nn.Sequential( - torch.nn.Dropout(self.model_config.seq_classif_dropout), - 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), - ) diff --git a/bert_squeeze/models/lt_fastbert.py b/bert_squeeze/models/lt_fastbert.py index 7ff74a4..93c4623 100644 --- a/bert_squeeze/models/lt_fastbert.py +++ b/bert_squeeze/models/lt_fastbert.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import logging import os from collections import defaultdict @@ -45,14 +47,20 @@ def __init__( super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) - self.training_stage = getattr(kwargs, "training_stage", 0) + training_stage = kwargs.get("training_stage", 0) + if not isinstance(training_stage, int) or training_stage not in {0, 1}: + raise ValueError("training_stage must be 0 or 1.") + self.training_stage = training_stage self._build_model() if self.training_stage == 0: - self._load_pretrained_bert_model( - getattr(kwargs, "pretrained_model_path", None) - ) + pretrained_model_path = kwargs.get("pretrained_model_path") + if pretrained_model_path is not None and not isinstance( + pretrained_model_path, str + ): + raise TypeError("pretrained_model_path must be a string.") + self._load_pretrained_bert_model(pretrained_model_path) @overrides def forward( @@ -203,7 +211,9 @@ def _build_model(self): self.embeddings = BertEmbeddings(self.model_config) self.encoder = FastBertGraph(self.model_config) - def _load_pretrained_bert_model(self, pretrained_model_path: str = None) -> None: + def _load_pretrained_bert_model( + self, pretrained_model_path: Optional[str] = None + ) -> None: """ Loads the pretrained weights into the model. diff --git a/tests/assistants/test_assistant_defaults.py b/tests/assistants/test_assistant_defaults.py index 7c4ad49..8541067 100644 --- a/tests/assistants/test_assistant_defaults.py +++ b/tests/assistants/test_assistant_defaults.py @@ -1,3 +1,7 @@ +from __future__ import annotations + +import pytest + from bert_squeeze.assistants import DistilAssistant, TrainAssistant @@ -17,6 +21,39 @@ def test_train_assistant_does_not_mutate_data_overrides(): assert assistant._data_conf.dataset_config.path == "custom-dataset" +@pytest.mark.parametrize( + ("assistant_name", "model_field"), + [ + ("bert", "num_labels"), + ("fastbert", "num_labels"), + ("theseusbert", "num_labels"), + ("adapter", "labels"), + ], +) +def test_train_assistant_propagates_label_overrides( + assistant_name: str, model_field: str +) -> None: + labels = [0, 1, 2] + + assistant = TrainAssistant( + assistant_name, + general_kwargs={"labels": labels, "num_labels": len(labels)}, + ) + + model_value = assistant._model_conf[model_field] + expected_value = labels if model_field == "labels" else len(labels) + assert model_value == expected_value + assert assistant._model_conf.scorer.labels == labels + + +def test_train_assistant_rejects_mismatched_label_count() -> None: + with pytest.raises(ValueError, match="must match the number of labels"): + TrainAssistant( + "bert", + general_kwargs={"labels": [0, 1, 2], "num_labels": 2}, + ) + + def test_distil_assistant_uses_default_data_config_and_keeps_name(): assistant = DistilAssistant("distil") diff --git a/tests/test_custom_model_initialization.py b/tests/test_custom_model_initialization.py index 29576dc..be2ce3b 100644 --- a/tests/test_custom_model_initialization.py +++ b/tests/test_custom_model_initialization.py @@ -1,22 +1,32 @@ +from __future__ import annotations + import torch from omegaconf import OmegaConf -from transformers import BertConfig +from transformers import ( + BertConfig, + BertForSequenceClassification, + DistilBertConfig, + DistilBertForSequenceClassification, +) from bert_squeeze.models.custom_transformers.berxit import BerxitModel from bert_squeeze.models.custom_transformers.deebert import DeeBertModel +from bert_squeeze.models.custom_transformers.fastbert import FastBertGraph from bert_squeeze.models.custom_transformers.theseus_bert import TheseusBertModel from bert_squeeze.models.lt_berxit import LtBerxit from bert_squeeze.models.lt_deebert import LtDeeBert +from bert_squeeze.models.lt_distilbert import LtCustomDistilBert +from bert_squeeze.models.lt_fastbert import LtFastBert from bert_squeeze.models.lt_theseus_bert import LtTheseusBert -def _bert_config(num_hidden_layers: int = 1) -> BertConfig: +def _bert_config(num_hidden_layers: int = 1, num_labels: int = 2) -> BertConfig: return BertConfig( hidden_size=16, intermediate_size=32, num_attention_heads=2, num_hidden_layers=num_hidden_layers, - num_labels=2, + num_labels=num_labels, vocab_size=32, ) @@ -41,34 +51,44 @@ def _inputs(): def test_deebert_loads_and_uses_a_custom_pretrained_encoder(tmp_path): - source_encoder = DeeBertModel(_bert_config()) + source_encoder = DeeBertModel(_bert_config(num_hidden_layers=2)) source_encoder.save_pretrained(tmp_path) module = LtDeeBert( - training_config=_training_config(train_highway=False, early_exit_entropy=-1.0), + training_config=_training_config(train_highway=True, early_exit_entropy=-1.0), pretrained_model=str(tmp_path), num_labels=2, ) + module.eval() logits, _, _ = module(**_inputs()) + loss = module.test_step({**_inputs(), "labels": torch.tensor([0, 1])}, 0) + predictions = module.predict_step(_inputs(), 0) assert module.model is module.bert + assert module.bert.encoder.layer[0] is not module.bert.encoder.layer[1] + assert module.bert.encoder.ramp[0] is not module.bert.encoder.ramp[1] assert torch.equal( module.bert.embeddings.word_embeddings.weight, source_encoder.embeddings.word_embeddings.weight, ) assert logits.shape == (2, 2) + assert predictions.shape == (2, 2) + assert torch.isfinite(loss) def test_berxit_loads_and_uses_a_custom_pretrained_encoder(tmp_path): - source_encoder = BerxitModel(_bert_config()) + source_encoder = BerxitModel(_bert_config(num_hidden_layers=2)) source_encoder.save_pretrained(tmp_path) module = LtBerxit( - training_config=_training_config(train_highway=False, early_exit_entropy=-1.0), + training_config=_training_config(train_highway=True, early_exit_entropy=-1.0), pretrained_model=str(tmp_path), num_labels=2, ) + module.eval() logits, _, _, _ = module(**_inputs()) + loss = module.test_step({**_inputs(), "labels": torch.tensor([0, 1])}, 0) + predictions = module.predict_step(_inputs(), 0) assert module.model is module.bert assert torch.equal( @@ -76,6 +96,78 @@ def test_berxit_loads_and_uses_a_custom_pretrained_encoder(tmp_path): source_encoder.embeddings.word_embeddings.weight, ) assert logits.shape == (2, 2) + assert predictions.shape == (2, 2) + assert torch.isfinite(loss) + + +def test_distilbert_uses_injected_sequence_classifier(tmp_path): + config = DistilBertConfig( + vocab_size=32, + dim=16, + hidden_dim=32, + n_layers=2, + n_heads=2, + num_labels=3, + ) + injected_model = DistilBertForSequenceClassification(config) + injected_model.save_pretrained(tmp_path) + module = LtCustomDistilBert( + training_config=_training_config(), + pretrained_model=str(tmp_path), + num_labels=3, + model=injected_model, + ) + batch = { + "input_ids": torch.tensor([[1, 2, 3], [4, 5, 0]]), + "attention_mask": torch.tensor([[1, 1, 1], [1, 1, 0]]), + "labels": torch.tensor([0, 2]), + } + + logits = module(batch["input_ids"], batch["attention_mask"]) + attention_logits, attentions = module( + batch["input_ids"], + batch["attention_mask"], + output_attentions=True, + ) + loss = module.training_step(batch, 0) + loss.backward() + + assert module.model is injected_model + assert logits.shape == (2, 3) + assert attention_logits.shape == (2, 3) + assert len(attentions) == config.n_layers + assert injected_model.classifier.weight.grad is not None + + +def test_fastbert_inference_uses_configured_label_count() -> None: + graph = FastBertGraph(_bert_config(num_hidden_layers=2, num_labels=3)) + embeddings = torch.randn(2, 4, 16) + + probabilities, _ = graph( + embeddings=embeddings, + attention_mask=torch.zeros(2, 1, 1, 4), + device="cpu", + inference=True, + inference_speed=1.1, + ) + + assert probabilities.shape == (2, 3) + + +def test_fastbert_respects_requested_training_stage(tmp_path) -> None: + config = _bert_config(num_hidden_layers=2) + injected_model = BertForSequenceClassification(config) + injected_model.save_pretrained(tmp_path) + + module = LtFastBert( + training_config=_training_config(), + pretrained_model=str(tmp_path), + num_labels=2, + model=injected_model, + training_stage=1, + ) + + assert module.training_stage == 1 def test_theseus_loads_and_uses_a_custom_pretrained_encoder(tmp_path):