Skip to content
Closed
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
8 changes: 8 additions & 0 deletions bert_squeeze/assistants/configs/distil_seq2seq.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -35,13 +35,17 @@ 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"
pretrained_model: "t5-small"
training_config: ${train}
generate_kwargs:
do_sample: false
num_beams: 4
early_stopping: true
training_config: ${train}

data:
Expand All @@ -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:
Expand All @@ -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}
Expand Down
5 changes: 4 additions & 1 deletion bert_squeeze/assistants/configs/train_t5.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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:
Expand All @@ -47,4 +51,3 @@ data:
max_target_length: 64
max_source_length: 256
tokenizer_name: ${model.pretrained_model}

21 changes: 17 additions & 4 deletions bert_squeeze/data/modules/transformer_module.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
14 changes: 13 additions & 1 deletion bert_squeeze/distillation/base_distiller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
14 changes: 13 additions & 1 deletion bert_squeeze/models/base_lt_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
14 changes: 13 additions & 1 deletion bert_squeeze/models/lt_berxit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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":
Expand Down
14 changes: 13 additions & 1 deletion bert_squeeze/models/lt_deebert.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
Loading