From d328ee439824de2f65cd36e670e44325a7a93a92 Mon Sep 17 00:00:00 2001 From: Jules Belveze Date: Wed, 22 Jul 2026 13:07:33 +0200 Subject: [PATCH 1/2] [bert_squeeze] - refactor: unify model initialization across classes - Standardizes the model initialization process in multiple classes to ensure consistency and reduce redundancy. - Simplifies the model building logic by directly using the provided model instead of re-instantiating it. --- bert_squeeze/models/lt_bert.py | 6 +++++- bert_squeeze/models/lt_berxit.py | 12 ++++++++++-- bert_squeeze/models/lt_deebert.py | 12 ++++++++++-- bert_squeeze/models/lt_theseus_bert.py | 14 ++++++++++---- 4 files changed, 35 insertions(+), 9 deletions(-) diff --git a/bert_squeeze/models/lt_bert.py b/bert_squeeze/models/lt_bert.py index 68f2438..2bc52ae 100644 --- a/bert_squeeze/models/lt_bert.py +++ b/bert_squeeze/models/lt_bert.py @@ -41,9 +41,13 @@ def __init__( scorer: Scorer = None, **kwargs, ): + if model is None: + model = CustomBertModel.from_pretrained(pretrained_model) + super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) + self._build_model() @overrides def forward( @@ -142,7 +146,7 @@ def test_step(self, batch, batch_idx, *args, **kwargs) -> torch.Tensor: def _build_model(self): """""" - self.encoder = CustomBertModel.from_pretrained(self.pretrained_model) + self.encoder = self.model 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), diff --git a/bert_squeeze/models/lt_berxit.py b/bert_squeeze/models/lt_berxit.py index beb4f2f..b66c4b0 100644 --- a/bert_squeeze/models/lt_berxit.py +++ b/bert_squeeze/models/lt_berxit.py @@ -8,6 +8,7 @@ from omegaconf import DictConfig, ListConfig from overrides import overrides from torch.nn import CrossEntropyLoss +from transformers import AutoConfig from bert_squeeze.utils.scorers import Scorer @@ -32,6 +33,14 @@ def __init__( scorer: Scorer = None, **kwargs, ): + if model is None: + model = BerxitModel.from_pretrained( + pretrained_model, + config=AutoConfig.from_pretrained( + pretrained_model, num_labels=num_labels + ), + ) + super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) @@ -406,7 +415,7 @@ def _build_model(self): self.model_config.gate_hidden_dim = getattr( self.config, "gate_hidden_dim", 32 ) - self.bert = BerxitModel(self.model_config) + 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( @@ -417,7 +426,6 @@ def _build_model(self): torch.nn.Linear(self.model_config.hidden_size, self.model_config.num_labels), ) - self.bert.init_weights() 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"): diff --git a/bert_squeeze/models/lt_deebert.py b/bert_squeeze/models/lt_deebert.py index ce7aa69..a48c3a6 100644 --- a/bert_squeeze/models/lt_deebert.py +++ b/bert_squeeze/models/lt_deebert.py @@ -8,6 +8,7 @@ from omegaconf import DictConfig, ListConfig from overrides import overrides from torch.nn import CrossEntropyLoss +from transformers import AutoConfig from bert_squeeze.utils.scorers import Scorer @@ -42,6 +43,14 @@ def __init__( scorer: Scorer = None, **kwargs, ): + if model is None: + model = DeeBertModel.from_pretrained( + pretrained_model, + config=AutoConfig.from_pretrained( + pretrained_model, num_labels=num_labels + ), + ) + super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) @@ -384,7 +393,7 @@ def loss( def _build_model(self): """""" - self.bert = DeeBertModel(self.model_config) + 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( @@ -395,6 +404,5 @@ def _build_model(self): torch.nn.Linear(self.model_config.hidden_size, self.model_config.num_labels), ) - self.bert.init_weights() self.bert.encoder.set_early_exit_entropy(self.config.early_exit_entropy) self.bert.init_highway_pooler() diff --git a/bert_squeeze/models/lt_theseus_bert.py b/bert_squeeze/models/lt_theseus_bert.py index cfd9ec3..9308d26 100644 --- a/bert_squeeze/models/lt_theseus_bert.py +++ b/bert_squeeze/models/lt_theseus_bert.py @@ -47,6 +47,14 @@ def __init__( scorer: Scorer = None, **kwargs, ): + if model is None: + model = TheseusBertModel.from_pretrained( + pretrained_model, + config=AutoConfig.from_pretrained( + pretrained_model, num_labels=num_labels + ), + ) + super().__init__( training_config, pretrained_model, num_labels, model, scorer, **kwargs ) @@ -158,10 +166,8 @@ def test_step(self, batch, batch_idx, *args, **kwargs) -> None: def _build_model(self): """""" - encoder = TheseusBertModel(AutoConfig.from_pretrained(self.pretrained_model)) - encoder.from_pretrained(self.pretrained_model) - encoder.encoder.init_successor_layers() - self.encoder = encoder + self.encoder = self.model + self.encoder.encoder.init_successor_layers() self.classifier = torch.nn.Sequential( torch.nn.Dropout(self.model_config.hidden_dropout_prob), From 7ecbfcf67093374a2a9106e5e471cbf69d3551cd Mon Sep 17 00:00:00 2001 From: Jules Belveze Date: Wed, 22 Jul 2026 13:07:42 +0200 Subject: [PATCH 2/2] [tests] - test: add tests for custom model initialization - Implement tests to verify that custom transformer models load and utilize pretrained encoders correctly. - Ensure that the model configurations and outputs are validated for different model types. --- tests/test_custom_model_initialization.py | 100 ++++++++++++++++++++++ 1 file changed, 100 insertions(+) create mode 100644 tests/test_custom_model_initialization.py diff --git a/tests/test_custom_model_initialization.py b/tests/test_custom_model_initialization.py new file mode 100644 index 0000000..29576dc --- /dev/null +++ b/tests/test_custom_model_initialization.py @@ -0,0 +1,100 @@ +import torch +from omegaconf import OmegaConf +from transformers import BertConfig + +from bert_squeeze.models.custom_transformers.berxit import BerxitModel +from bert_squeeze.models.custom_transformers.deebert import DeeBertModel +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_theseus_bert import LtTheseusBert + + +def _bert_config(num_hidden_layers: int = 1) -> BertConfig: + return BertConfig( + hidden_size=16, + intermediate_size=32, + num_attention_heads=2, + num_hidden_layers=num_hidden_layers, + num_labels=2, + vocab_size=32, + ) + + +def _training_config(**overrides): + config = { + "logging_steps": 2, + "accumulation_steps": 1, + "objective": "ce", + "lr_scheduler": False, + } + config.update(overrides) + return OmegaConf.create(config) + + +def _inputs(): + return { + "input_ids": torch.tensor([[1, 2, 3, 0], [4, 5, 0, 0]]), + "attention_mask": torch.tensor([[1, 1, 1, 0], [1, 1, 0, 0]]), + "token_type_ids": torch.zeros((2, 4), dtype=torch.long), + } + + +def test_deebert_loads_and_uses_a_custom_pretrained_encoder(tmp_path): + source_encoder = DeeBertModel(_bert_config()) + source_encoder.save_pretrained(tmp_path) + + module = LtDeeBert( + training_config=_training_config(train_highway=False, early_exit_entropy=-1.0), + pretrained_model=str(tmp_path), + num_labels=2, + ) + logits, _, _ = module(**_inputs()) + + assert module.model is module.bert + assert torch.equal( + module.bert.embeddings.word_embeddings.weight, + source_encoder.embeddings.word_embeddings.weight, + ) + assert logits.shape == (2, 2) + + +def test_berxit_loads_and_uses_a_custom_pretrained_encoder(tmp_path): + source_encoder = BerxitModel(_bert_config()) + source_encoder.save_pretrained(tmp_path) + + module = LtBerxit( + training_config=_training_config(train_highway=False, early_exit_entropy=-1.0), + pretrained_model=str(tmp_path), + num_labels=2, + ) + logits, _, _, _ = module(**_inputs()) + + assert module.model is module.bert + assert torch.equal( + module.bert.embeddings.word_embeddings.weight, + source_encoder.embeddings.word_embeddings.weight, + ) + assert logits.shape == (2, 2) + + +def test_theseus_loads_and_uses_a_custom_pretrained_encoder(tmp_path): + source_encoder = TheseusBertModel(_bert_config(num_hidden_layers=6)) + source_encoder.save_pretrained(tmp_path) + + module = LtTheseusBert( + training_config=_training_config(), + pretrained_model=str(tmp_path), + num_labels=2, + replacement_scheduler=OmegaConf.create( + {"type": "constant", "replacing_rate": 1.0} + ), + ) + logits = module(**_inputs()) + + assert module.model is module.encoder + assert torch.equal( + module.encoder.embeddings.word_embeddings.weight, + source_encoder.embeddings.word_embeddings.weight, + ) + assert logits.shape == (2, 2)