diff --git a/bert_squeeze/assistants/configs/distil_seq2seq.yaml b/bert_squeeze/assistants/configs/distil_seq2seq.yaml index 0d9f2a2..22c14d1 100644 --- a/bert_squeeze/assistants/configs/distil_seq2seq.yaml +++ b/bert_squeeze/assistants/configs/distil_seq2seq.yaml @@ -35,6 +35,8 @@ model: training_config: ${train} generate_kwargs: do_sample: false + num_beams: 4 + early_stopping: true student: _target_: bert_squeeze.models.lt_t5.SimpleT5Model task: "summarization" @@ -42,6 +44,8 @@ model: training_config: ${train} generate_kwargs: do_sample: false + num_beams: 4 + early_stopping: true training_config: ${train} data: @@ -50,6 +54,8 @@ data: _target_: bert_squeeze.data.modules.transformer_module.Seq2SeqTransformerDataModule dataset_config: is_local: false + source_prefix: + target_prefix: target_col: path: split: @@ -65,6 +71,8 @@ data: split: ${data.teacher_module.dataset_config.split} source_col: ${data.teacher_module.dataset_config.source_col} target_col: ${data.teacher_module.dataset_config.target_col} + source_prefix: ${data.teacher_module.dataset_config.source_prefix} + target_prefix: ${data.teacher_module.dataset_config.target_prefix} tokenizer_name: ${model.student.pretrained_model} max_target_length: ${data.teacher_module.max_target_length} max_source_length: ${data.teacher_module.max_source_length} diff --git a/bert_squeeze/assistants/configs/train_t5.yaml b/bert_squeeze/assistants/configs/train_t5.yaml index 9ae39c3..a573b43 100644 --- a/bert_squeeze/assistants/configs/train_t5.yaml +++ b/bert_squeeze/assistants/configs/train_t5.yaml @@ -31,6 +31,8 @@ model: training_config: ${train} generate_kwargs: do_sample: false + num_beams: 4 + early_stopping: true scorer: _target_: bert_squeeze.utils.scorers.lm_scorer.LMScorer tokenizer_name: ${model.pretrained_model} @@ -39,6 +41,8 @@ data: _target_: bert_squeeze.data.modules.transformer_module.Seq2SeqTransformerDataModule dataset_config: is_local: false + source_prefix: + target_prefix: target_col: path: split: @@ -47,4 +51,3 @@ data: max_target_length: 64 max_source_length: 256 tokenizer_name: ${model.pretrained_model} - diff --git a/bert_squeeze/data/modules/transformer_module.py b/bert_squeeze/data/modules/transformer_module.py index ea63ea2..90ffdde 100644 --- a/bert_squeeze/data/modules/transformer_module.py +++ b/bert_squeeze/data/modules/transformer_module.py @@ -1,5 +1,5 @@ from pathlib import Path -from typing import List, Optional, Sequence +from typing import List, Mapping, Optional, Sequence import datasets from omegaconf import DictConfig @@ -209,6 +209,8 @@ def __init__( self.dataset_config = dataset_config self.source_col = dataset_config.source_col self.target_col = dataset_config.target_col + self.source_prefix = dataset_config.source_prefix + self.target_prefix = dataset_config.target_prefix raw_paths = dataset_config.get("data_path", None) self._uses_data_path = raw_paths is not None @@ -362,9 +364,20 @@ def featurize(self) -> datasets.DatasetDict: Returns: DatasetDict: featurized dataset """ + source_prefix = self.source_prefix + target_prefix = self.target_prefix + source_col = self.source_col + target_col = self.target_col + + def _format_source(example: Mapping[str, object]) -> str: + return f"{source_prefix}{example[source_col]}" + + def _format_target(example: Mapping[str, object]) -> str: + return f"{target_prefix}{example[target_col]}" + tokenized_dataset = self.dataset.map( - lambda x: self.tokenizer( - x[self.source_col], + lambda example: self.tokenizer( + _format_source(example), padding=False, max_length=self.max_source_length, truncation=True, @@ -374,7 +387,7 @@ def featurize(self) -> datasets.DatasetDict: tokenized_dataset = tokenized_dataset.map( lambda x: { "labels": self.tokenizer( - x[self.target_col], + _format_target(x), padding=False, max_length=self.max_target_length, truncation=True, diff --git a/bert_squeeze/distillation/base_distiller.py b/bert_squeeze/distillation/base_distiller.py index 0e1cad9..1e903f1 100644 --- a/bert_squeeze/distillation/base_distiller.py +++ b/bert_squeeze/distillation/base_distiller.py @@ -59,7 +59,19 @@ def _get_student_parameters(self) -> List[Dict]: Returns: List[Dict]: group of parameters to optimize """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight'] + no_decay = [ + 'bias', + 'gamma', + 'beta', + 'LayerNorm.weight', + 'LayerNorm.bias', + 'layer_norm.weight', + 'layer_norm.bias', + 'layernorm.weight', + 'layernorm.bias', + 'ln_f.weight', + 'ln_f.bias', + ] if self.params.discriminative_learning: if ( diff --git a/bert_squeeze/models/base_lt_module.py b/bert_squeeze/models/base_lt_module.py index f1085fe..731e9f3 100644 --- a/bert_squeeze/models/base_lt_module.py +++ b/bert_squeeze/models/base_lt_module.py @@ -172,7 +172,19 @@ def _get_optimizer_parameters(self) -> List[Dict]: Returns: List[Dict]: group of parameters to optimize """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight'] + no_decay = [ + 'bias', + 'gamma', + 'beta', + 'LayerNorm.weight', + 'LayerNorm.bias', + 'layer_norm.weight', + 'layer_norm.bias', + 'layernorm.weight', + 'layernorm.bias', + 'ln_f.weight', + 'ln_f.bias', + ] if self.config.discriminative_learning: if ( diff --git a/bert_squeeze/models/lt_berxit.py b/bert_squeeze/models/lt_berxit.py index 756035a..4b8c72b 100644 --- a/bert_squeeze/models/lt_berxit.py +++ b/bert_squeeze/models/lt_berxit.py @@ -214,7 +214,19 @@ def _maybe_switch_stage(self) -> None: 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'] + no_decay = [ + 'bias', + 'gamma', + 'beta', + 'LayerNorm.weight', + 'LayerNorm.bias', + 'layer_norm.weight', + 'layer_norm.bias', + 'layernorm.weight', + 'layernorm.bias', + 'ln_f.weight', + 'ln_f.bias', + ] # Gate-only training stage: optimize only gate parameters if getattr(self, "train_stage", "backbone") == "gates": diff --git a/bert_squeeze/models/lt_deebert.py b/bert_squeeze/models/lt_deebert.py index a0c0628..fd42444 100644 --- a/bert_squeeze/models/lt_deebert.py +++ b/bert_squeeze/models/lt_deebert.py @@ -212,7 +212,19 @@ def _get_optimizer_parameters(self) -> List[Dict]: Returns: List[Dict]: group of parameters to optimize """ - no_decay = ['bias', 'gamma', 'beta', 'LayerNorm.weight'] + no_decay = [ + 'bias', + 'gamma', + 'beta', + 'LayerNorm.weight', + 'LayerNorm.bias', + 'layer_norm.weight', + 'layer_norm.bias', + 'layernorm.weight', + 'layernorm.bias', + 'ln_f.weight', + 'ln_f.bias', + ] if self.config.discriminative_learning: if (