Skip to content
This repository was archived by the owner on Apr 8, 2025. It is now read-only.

Commit 19fef0e

Browse files
authored
DPR - enable loading of custom BERT models (#746)
* Add loading of other bert type models * Fix loading and saving, add test
1 parent 328f2be commit 19fef0e

4 files changed

Lines changed: 63 additions & 14 deletions

File tree

‎farm/data_handler/processor.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2868,6 +2868,9 @@ def file_to_dicts(self, file: str) -> [dict]:
28682868
...]}
28692869
"""
28702870
dicts = read_dpr_json(file, max_samples=self.max_samples, num_hard_negatives=self.num_hard_negatives, num_positives=self.num_positives, shuffle_negatives=self.shuffle_negatives, shuffle_positives=self.shuffle_positives)
2871+
2872+
# shuffle dicts to make sure that similar positive passages do not end up in one batch
2873+
dicts = random.sample(dicts, len(dicts))
28712874
return dicts
28722875

28732876
def dataset_from_dicts(self, dicts, indices=None, return_baskets = False):

‎farm/modeling/language_model.py‎

Lines changed: 24 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1447,16 +1447,23 @@ def load(cls, pretrained_model_name_or_path, language=None, **kwargs):
14471447
dpr_question_encoder.model = transformers.DPRQuestionEncoder.from_pretrained(farm_lm_model, config=dpr_config, **kwargs)
14481448
dpr_question_encoder.language = dpr_question_encoder.model.config.language
14491449
else:
1450-
model_type = AutoConfig.from_pretrained(pretrained_model_name_or_path).model_type
1451-
if model_type == "dpr":
1450+
original_model_config = AutoConfig.from_pretrained(pretrained_model_name_or_path)
1451+
if original_model_config.model_type == "dpr":
14521452
# "pretrained dpr model": load existing pretrained DPRQuestionEncoder model
14531453
dpr_question_encoder.model = transformers.DPRQuestionEncoder.from_pretrained(
14541454
str(pretrained_model_name_or_path), **kwargs)
14551455
else:
14561456
# "from scratch": load weights from different architecture (e.g. bert) into DPRQuestionEncoder
1457-
dpr_question_encoder.model = transformers.DPRQuestionEncoder(config=transformers.DPRConfig(**kwargs))
1457+
# but keep config values from original architecture
1458+
# TODO test for architectures other than BERT, e.g. Electra
1459+
if original_model_config.model_type != "bert":
1460+
logger.warning(f"Using a model of type '{original_model_config.model_type}' which might be incompatible with DPR encoders."
1461+
f"Bert based encoders are supported that need input_ids,token_type_ids,attention_mask as input tensors.")
1462+
original_config_dict = vars(original_model_config)
1463+
original_config_dict.update(kwargs)
1464+
dpr_question_encoder.model = transformers.DPRQuestionEncoder(config=transformers.DPRConfig(**original_config_dict))
14581465
dpr_question_encoder.model.base_model.bert_model = AutoModel.from_pretrained(
1459-
str(pretrained_model_name_or_path), **kwargs)
1466+
str(pretrained_model_name_or_path), **original_config_dict)
14601467
dpr_question_encoder.language = cls._get_or_infer_language_from_name(language, pretrained_model_name_or_path)
14611468

14621469
return dpr_question_encoder
@@ -1541,16 +1548,25 @@ def load(cls, pretrained_model_name_or_path, language=None, **kwargs):
15411548
dpr_context_encoder.language = dpr_context_encoder.model.config.language
15421549
else:
15431550
# Pytorch-transformer Style
1544-
model_type = AutoConfig.from_pretrained(pretrained_model_name_or_path).model_type
1545-
if model_type == "dpr":
1551+
original_model_config = AutoConfig.from_pretrained(pretrained_model_name_or_path)
1552+
if original_model_config.model_type == "dpr":
15461553
# "pretrained dpr model": load existing pretrained DPRContextEncoder model
15471554
dpr_context_encoder.model = transformers.DPRContextEncoder.from_pretrained(
15481555
str(pretrained_model_name_or_path), **kwargs)
15491556
else:
15501557
# "from scratch": load weights from different architecture (e.g. bert) into DPRContextEncoder
1551-
dpr_context_encoder.model = transformers.DPRContextEncoder(config=transformers.DPRConfig(**kwargs))
1558+
# but keep config values from original architecture
1559+
# TODO test for architectures other than BERT, e.g. Electra
1560+
if original_model_config.model_type != "bert":
1561+
logger.warning(
1562+
f"Using a model of type '{original_model_config.model_type}' which might be incompatible with DPR encoders."
1563+
f"Bert based encoders are supported that need input_ids,token_type_ids,attention_mask as input tensors.")
1564+
original_config_dict = vars(original_model_config)
1565+
original_config_dict.update(kwargs)
1566+
dpr_context_encoder.model = transformers.DPRContextEncoder(
1567+
config=transformers.DPRConfig(**original_config_dict))
15521568
dpr_context_encoder.model.base_model.bert_model = AutoModel.from_pretrained(
1553-
str(pretrained_model_name_or_path), **kwargs)
1569+
str(pretrained_model_name_or_path), **original_config_dict)
15541570
dpr_context_encoder.language = cls._get_or_infer_language_from_name(language, pretrained_model_name_or_path)
15551571

15561572
return dpr_context_encoder

‎farm/modeling/prediction_head.py‎

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1709,7 +1709,7 @@ def _embeddings_to_scores(self, query_vectors:torch.Tensor, passage_vectors:torc
17091709
softmax_scores = nn.functional.log_softmax(scores, dim=1)
17101710
return softmax_scores
17111711

1712-
def logits_to_loss(self, logits: Tuple[torch.Tensor, torch.Tensor], **kwargs):
1712+
def logits_to_loss(self, logits: Tuple[torch.Tensor, torch.Tensor], label_ids, **kwargs):
17131713
"""
17141714
Computes the loss (Default: NLLLoss) by applying a similarity function (Default: dot product) to the input
17151715
tuple of (query_vectors, passage_vectors) and afterwards applying the loss function on similarity scores.
@@ -1729,8 +1729,7 @@ def logits_to_loss(self, logits: Tuple[torch.Tensor, torch.Tensor], **kwargs):
17291729
query_vectors, passage_vectors = logits
17301730

17311731
# Prepare Labels
1732-
lm_label_ids = kwargs.get(self.label_tensor_name)
1733-
positive_idx_per_question = torch.nonzero((lm_label_ids.view(-1) == 1), as_tuple=False)
1732+
positive_idx_per_question = torch.nonzero((label_ids.view(-1) == 1), as_tuple=False)
17341733

17351734
# Gather global embeddings from all distributed nodes (DDP)
17361735
if rank != -1:
@@ -1788,13 +1787,12 @@ def logits_to_preds(self, logits: Tuple[torch.Tensor, torch.Tensor], **kwargs):
17881787
_, sorted_scores = torch.sort(softmax_scores, dim=1, descending=True)
17891788
return sorted_scores
17901789

1791-
def prepare_labels(self, **kwargs):
1790+
def prepare_labels(self, label_ids, **kwargs):
17921791
"""
17931792
Returns a tensor with passage labels(0:hard_negative/1:positive) for each query
17941793
17951794
:return: passage labels(0:hard_negative/1:positive) for each query
17961795
"""
1797-
label_ids = kwargs.get(self.label_tensor_name)
17981796
labels = torch.zeros(label_ids.size(0), label_ids.numel())
17991797

18001798
positive_indices = torch.nonzero(label_ids.view(-1) == 1, as_tuple=False)

‎test/test_dpr.py‎

Lines changed: 33 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import torch
33
import logging
44
import numpy as np
5+
from pathlib import Path
56
from farm.data_handler.processor import TextSimilarityProcessor
67
from farm.data_handler.data_silo import DataSilo
78
from farm.train import Trainer
@@ -402,7 +403,7 @@ def test_dpr_context_only():
402403
assert tensor_names == ['passage_input_ids', 'passage_segment_ids', 'passage_attention_mask', 'label_ids']
403404

404405

405-
def test_dpr_save_load():
406+
def test_dpr_processor_save_load():
406407
d = {'query': 'big little lies season 2 how many episodes',
407408
'passages': [
408409
{'title': 'Big Little Lies (TV series)',
@@ -443,6 +444,7 @@ def test_dpr_save_load():
443444
assert np.array_equal(dataset.tensors[0],dataset2.tensors[0])
444445

445446

447+
446448
def test_dpr_training():
447449
batch_size = 1
448450
n_epochs = 1
@@ -524,6 +526,36 @@ def test_dpr_training():
524526

525527
trainer.train()
526528

529+
######## save and load model again
530+
save_dir = Path("testsave/dpr-model")
531+
model.save(save_dir)
532+
del model
533+
534+
model2 = BiAdaptiveModel.load(save_dir, device=device)
535+
model2, optimizer2, lr_schedule = initialize_optimizer(
536+
model=model2,
537+
learning_rate=1e-5,
538+
optimizer_opts={"name": "TransformersAdamW", "correct_bias": True, "weight_decay": 0.0, \
539+
"eps": 1e-08},
540+
schedule_opts={"name": "LinearWarmup", "num_warmup_steps": 100},
541+
n_batches=len(data_silo.loaders["train"]),
542+
n_epochs=n_epochs,
543+
grad_acc_steps=1,
544+
device=device,
545+
distributed=distributed
546+
)
547+
trainer2 = Trainer(
548+
model=model2,
549+
optimizer=optimizer,
550+
data_silo=data_silo,
551+
epochs=n_epochs,
552+
n_gpu=n_gpu,
553+
lr_schedule=lr_schedule,
554+
evaluate_every=evaluate_every,
555+
device=device,
556+
)
557+
558+
trainer2.train()
527559

528560

529561
if __name__=="__main__":

0 commit comments

Comments
 (0)