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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion bert_squeeze/models/lt_bert.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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),
Expand Down
12 changes: 10 additions & 2 deletions bert_squeeze/models/lt_berxit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
)
Expand Down Expand Up @@ -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(
Expand All @@ -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"):
Expand Down
12 changes: 10 additions & 2 deletions bert_squeeze/models/lt_deebert.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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(
Expand All @@ -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()
14 changes: 10 additions & 4 deletions bert_squeeze/models/lt_theseus_bert.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down Expand Up @@ -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),
Expand Down
100 changes: 100 additions & 0 deletions tests/test_custom_model_initialization.py
Original file line number Diff line number Diff line change
@@ -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)
Loading