From 9a740824b0c5b56309fab020f3dc583c0b373cdc Mon Sep 17 00:00:00 2001 From: JulesBelveze Date: Sun, 18 Jan 2026 16:15:56 +0100 Subject: [PATCH 1/2] [bert_squeeze] - feature: improve configuration and tokenization for seq2seq models - Enable beam search with early stopping during the generation process in seq2seq model configurations - Introduce source and target prefixes in data module configurations to support structured input and output for seq2seq tasks - Update the parameter grouping for optimizer to include additional no-decay layer norms and biases across various model modules --- .../assistants/configs/distil_seq2seq.yaml | 8 +++++++ bert_squeeze/assistants/configs/train_t5.yaml | 5 ++++- .../data/modules/transformer_module.py | 21 +++++++++++++++---- bert_squeeze/distillation/base_distiller.py | 14 ++++++++++++- bert_squeeze/models/base_lt_module.py | 14 ++++++++++++- bert_squeeze/models/lt_berxit.py | 14 ++++++++++++- bert_squeeze/models/lt_deebert.py | 14 ++++++++++++- 7 files changed, 81 insertions(+), 9 deletions(-) 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..1c7b395 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.get("source_prefix") + self.target_prefix = dataset_config.get("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 or "" + target_prefix = self.target_prefix or "" + 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 ( From b51631e2ef24bafffeaaae98cf4a1e68ea5fa31a Mon Sep 17 00:00:00 2001 From: JulesBelveze Date: Sun, 18 Jan 2026 21:23:17 +0100 Subject: [PATCH 2/2] [bert_squeeze/data/modules] - fix: ensure proper handling of dataset prefix configuration - Changed the retrieval of `source_prefix` and `target_prefix` from using a `get` method to direct attribute access to avoid errors related to missing configuration - Removed fallback to empty strings for `source_prefix` and `target_prefix` to enforce explicit configuration values --- bert_squeeze/data/modules/transformer_module.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/bert_squeeze/data/modules/transformer_module.py b/bert_squeeze/data/modules/transformer_module.py index 1c7b395..90ffdde 100644 --- a/bert_squeeze/data/modules/transformer_module.py +++ b/bert_squeeze/data/modules/transformer_module.py @@ -209,8 +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.get("source_prefix") - self.target_prefix = dataset_config.get("target_prefix") + 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 @@ -364,8 +364,8 @@ def featurize(self) -> datasets.DatasetDict: Returns: DatasetDict: featurized dataset """ - source_prefix = self.source_prefix or "" - target_prefix = self.target_prefix or "" + source_prefix = self.source_prefix + target_prefix = self.target_prefix source_col = self.source_col target_col = self.target_col