Compare commits
275
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a07fcff087 | ||
|
|
243f2d98f4 | ||
|
|
7c6802d216 | ||
|
|
6dd3122ad4 | ||
|
|
3a8f32d919 | ||
|
|
29d1eb57ac | ||
|
|
02418a591b | ||
|
|
be95182ee1 | ||
|
|
970fe3d833 | ||
|
|
97e29a9cf3 | ||
|
|
785a1b6076 | ||
|
|
9ef5581d5e | ||
|
|
657b7fdada | ||
|
|
3a5999f5d4 | ||
|
|
0efb2830f4 | ||
|
|
b5099747f9 | ||
|
|
34fd62a5ea | ||
|
|
e167e9fc74 | ||
|
|
1df5e7af4a | ||
|
|
544f543861 | ||
|
|
10970ec650 | ||
|
|
78b1c850f2 | ||
|
|
ac8d301786 | ||
|
|
9f32be4fbc | ||
|
|
74c8ac0f84 | ||
|
|
1d5e77d87b | ||
|
|
47b3638593 | ||
|
|
f7c95805cf | ||
|
|
0169cbbe49 | ||
|
|
8e4a74b1cb | ||
|
|
86034ae640 | ||
|
|
5283878c9f | ||
|
|
6bdfb14664 | ||
|
|
99de2c3c3c | ||
|
|
9e95429843 | ||
|
|
c9597fef42 | ||
|
|
c34b886cd1 | ||
|
|
f179e7b37c | ||
|
|
969c271bd7 | ||
|
|
9f98a6ac15 | ||
|
|
4606bda735 | ||
|
|
eb84b9d5ee | ||
|
|
5526be3499 | ||
|
|
6bf996f58d | ||
|
|
35a82ee87f | ||
|
|
379e8c7dba | ||
|
|
3e62d96b14 | ||
|
|
350eeb7cd2 | ||
|
|
c06e38add3 | ||
|
|
d94b09d285 | ||
|
|
eab9779732 | ||
|
|
4a6b23ee40 | ||
|
|
e2d45445c3 | ||
|
|
de7035d9f6 | ||
|
|
abee3e2ecd | ||
|
|
0274338aff | ||
|
|
e67ad41ae7 | ||
|
|
c5a1de5b75 | ||
|
|
af3f86fd6a | ||
|
|
fb6eb096ec | ||
|
|
bdf0b1ab88 | ||
|
|
9d99761b1b | ||
|
|
3f3040ae12 | ||
|
|
4224ac999d | ||
|
|
c7d05f70db | ||
|
|
ac47dbdfdf | ||
|
|
10b092ea27 | ||
|
|
e5728e4ddb | ||
|
|
5a281449df | ||
|
|
585bb576ba | ||
|
|
6aa28bd1f5 | ||
|
|
a2f2d5da31 | ||
|
|
b2b222d0db | ||
|
|
574ffbb1eb | ||
|
|
a549aaac3e | ||
|
|
bad1b75a91 | ||
|
|
6624b63c24 | ||
|
|
7f3f9fae22 | ||
|
|
c553e1214c | ||
|
|
d26f5dac45 | ||
|
|
913bb8f770 | ||
|
|
af05c0e004 | ||
|
|
dc785f0d10 | ||
|
|
d3b8694cd8 | ||
|
|
b5adf48404 | ||
|
|
315ab6be60 | ||
|
|
76377c9854 | ||
|
|
a387c057db | ||
|
|
0098872c04 | ||
|
|
a79f5fb77e | ||
|
|
0dc32e6f27 | ||
|
|
fa7f8185e8 | ||
|
|
31b4c8cc08 | ||
|
|
318a585f2b | ||
|
|
630af0c3fa | ||
|
|
3c384445a7 | ||
|
|
9f6596ebc9 | ||
|
|
582efc0319 | ||
|
|
3467f1c7f6 | ||
|
|
bcb4624d42 | ||
|
|
e46e10150d | ||
|
|
ecff6f92fa | ||
|
|
c59be9b2e0 | ||
|
|
b32d283d26 | ||
|
|
7ccd2366de | ||
|
|
c294e5e331 | ||
|
|
3b304afdf7 | ||
|
|
09f199b790 | ||
|
|
bb1bffc89e | ||
|
|
47c226ea3f | ||
|
|
3abbf78a48 | ||
|
|
9fc692d9f3 | ||
|
|
e146713686 | ||
|
|
014c061f3d | ||
|
|
e446fd24ef | ||
|
|
055167ba12 | ||
|
|
8f68b91e81 | ||
|
|
8b98c4cfa8 | ||
|
|
e3a00a402b | ||
|
|
24f483cf06 | ||
|
|
8d4f075245 | ||
|
|
0b98750ecd | ||
|
|
cc83d34a7d | ||
|
|
a65ab52d74 | ||
|
|
3e47e61a1a | ||
|
|
20aedfbcf7 | ||
|
|
440744a111 | ||
|
|
a2ba031b92 | ||
|
|
8f614af5b7 | ||
|
|
df0b28a581 | ||
|
|
8c6769ee41 | ||
|
|
b11fc6977e | ||
|
|
fb6adb2a47 | ||
|
|
789f06e9d8 | ||
|
|
e17df7499c | ||
|
|
82a6b8345f | ||
|
|
60a90143ee | ||
|
|
7b9557fbf5 | ||
|
|
105d83d595 | ||
|
|
bd361eae11 | ||
|
|
8fb24ce3ce | ||
|
|
6689b594fc | ||
|
|
25edee17ed | ||
|
|
7c52e73dfe | ||
|
|
d53236c840 | ||
|
|
eb4b3cce7a | ||
|
|
324d4dcfce | ||
|
|
3f6463a895 | ||
|
|
3c34cae08a | ||
|
|
175efd864a | ||
|
|
516bbcabd0 | ||
|
|
0c3d6d8cbf | ||
|
|
84ea9fac23 | ||
|
|
44743e8dff | ||
|
|
7f1af24a05 | ||
|
|
99e7161eb2 | ||
|
|
d92a69f223 | ||
|
|
4f19747960 | ||
|
|
9b54356b0b | ||
|
|
f750d801ff | ||
|
|
298ec5361c | ||
|
|
ee18f81852 | ||
|
|
e7d9dc3d22 | ||
|
|
d2cd7b4d9d | ||
|
|
d3611c9844 | ||
|
|
c10bb9fb23 | ||
|
|
81e939fb54 | ||
|
|
0c3d22ea86 | ||
|
|
63521f19a9 | ||
|
|
56aedacc27 | ||
|
|
0933e118e0 | ||
|
|
12f2d9e6bc | ||
|
|
20a59685a9 | ||
|
|
f5e697bbe6 | ||
|
|
db52f3708f | ||
|
|
800dcdcce2 | ||
|
|
88f9f917a6 | ||
|
|
a42f09eec6 | ||
|
|
c154fb3f8b | ||
|
|
c7f6e62c3e | ||
|
|
af8b962edf | ||
|
|
6f1757d2b8 | ||
|
|
4d3617acc6 | ||
|
|
2443a7eefc | ||
|
|
6f03e39b30 | ||
|
|
9f214624f4 | ||
|
|
e48a8299b9 | ||
|
|
20abeec1f5 | ||
|
|
584c329ffe | ||
|
|
d32887160a | ||
|
|
ef45ee9a7e | ||
|
|
8fbdd7d205 | ||
|
|
1ca011ff75 | ||
|
|
a78603cee7 | ||
|
|
d4830f9d0a | ||
|
|
c779c38297 | ||
|
|
021d2b9f80 | ||
|
|
fa0eda43b6 | ||
|
|
c6f7e14d0b | ||
|
|
db3d85ad2c | ||
|
|
705ed2ffb4 | ||
|
|
29499c9ca5 | ||
|
|
bcb9996ae5 | ||
|
|
aa995f98a3 | ||
|
|
5d751729ec | ||
|
|
f5606df47b | ||
|
|
ca82946559 | ||
|
|
3461e24b8d | ||
|
|
15a2d4ee79 | ||
|
|
3f73cb76a7 | ||
|
|
88b4b970d3 | ||
|
|
ac72c7aadc | ||
|
|
dd418e8d1d | ||
|
|
07da0623ad | ||
|
|
21a2cc0eb1 | ||
|
|
2c2a10e5d7 | ||
|
|
910da0fb96 | ||
|
|
7a7dcd397f | ||
|
|
b39c39fe87 | ||
|
|
ce0de9073a | ||
|
|
cf47a6ff9d | ||
|
|
a6420767e9 | ||
|
|
fee71b4fef | ||
|
|
a2671a6eee | ||
|
|
18524d8396 | ||
|
|
a250d1c6aa | ||
|
|
b62bfe666d | ||
|
|
1fc653f626 | ||
|
|
f281724e92 | ||
|
|
67f9553d4b | ||
|
|
cd461e13ff | ||
|
|
05bfb1eb1d | ||
|
|
c79351a64a | ||
|
|
c774160bd2 | ||
|
|
beda65dd17 | ||
|
|
59c658de2f | ||
|
|
5947812374 | ||
|
|
09cc1e64cf | ||
|
|
5fbf42258c | ||
|
|
af315c157e | ||
|
|
2c7294830e | ||
|
|
2b7132c25e | ||
|
|
bef77fc211 | ||
|
|
abb81df15f | ||
|
|
51221fbbf8 | ||
|
|
2ee9388492 | ||
|
|
c4530f91ab | ||
|
|
bbc4e52d3b | ||
|
|
74704cedab | ||
|
|
bf5782e080 | ||
|
|
3cebd56848 | ||
|
|
696c8c28f2 | ||
|
|
04a8ace0df | ||
|
|
0ed41562eb | ||
|
|
7081d0f3fd | ||
|
|
ca9c685453 | ||
|
|
70cf536ae2 | ||
|
|
5a3ed998f0 | ||
|
|
6302fb02ef | ||
|
|
5a35811635 | ||
|
|
174aaf374d | ||
|
|
1edc50f670 | ||
|
|
d2cc12b182 | ||
|
|
8afb88ce24 | ||
|
|
4be1287671 | ||
|
|
4f5790f652 | ||
|
|
fcc49a0c5c | ||
|
|
937d3d6ae8 | ||
|
|
b01f1d594c | ||
|
|
71850ad26b | ||
|
|
1368969930 | ||
|
|
f2de528d04 | ||
|
|
eabb3c5f9a | ||
|
|
762dd832f1 | ||
|
|
8280dad80e |
@@ -116,6 +116,19 @@ class BaseTransformer(pl.LightningModule):
|
||||
self.opt = optimizer
|
||||
return [optimizer]
|
||||
|
||||
def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, second_order_closure=None):
|
||||
if self.trainer.use_tpu:
|
||||
xm.optimizer_step(optimizer)
|
||||
else:
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
def get_tqdm_dict(self):
|
||||
avg_loss = getattr(self.trainer, "avg_loss", 0.0)
|
||||
tqdm_dict = {"loss": "{:.3f}".format(avg_loss), "lr": self.lr_scheduler.get_last_lr()[-1]}
|
||||
return tqdm_dict
|
||||
|
||||
def test_step(self, batch, batch_nb):
|
||||
return self.validation_step(batch, batch_nb)
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ export me=`git config user.name`
|
||||
```
|
||||
|
||||
Tips:
|
||||
- 1 epoch at batch size 1 for bart-large takes 24 hours, requires 13GB GPU RAM with fp16 on an NVIDIA-V100.
|
||||
- 1 epoch at batch size 1 for bart-large takes 24 hours, requires 13GB GPU RAM with fp16 on an NVIDIA-V100.
|
||||
- try `bart-base`, `--freeze_encoder` or `--freeze_embeds` for faster training/larger batch size. (3hr/epoch with bs=8, see below)
|
||||
- `fp16_opt_level=O1` (the default works best).
|
||||
- If you are finetuning on your own dataset, start from `bart-large-cnn` if you want long summaries and `bart-large-xsum` if you want short summaries.
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
|
||||
from durbango import lmap, tqdm_nice
|
||||
from transformers import BartTokenizer
|
||||
|
||||
|
||||
try:
|
||||
from finetune import calculate_rouge
|
||||
except ImportError:
|
||||
from .finetune import calculate_rouge
|
||||
|
||||
|
||||
LOGDIR = Path("examples/summarization/dbart/logs/").absolute()
|
||||
DATA_DIR = Path("examples/summarization/dbart/cnn_dm").absolute()
|
||||
|
||||
|
||||
def rouge_files(src_file: Path, tgt_file: Path):
|
||||
src = lmap(str.strip, list(src_file.open().readlines()))
|
||||
tgt = lmap(str.strip, list(tgt_file.open().readlines()))
|
||||
return calculate_rouge(src, tgt)
|
||||
|
||||
|
||||
def read_gens(exp_name, split="test", n=None):
|
||||
if Path(exp_name).exists():
|
||||
expdir = exp_name
|
||||
else:
|
||||
expdir = LOGDIR / exp_name
|
||||
assert expdir.exists(), expdir
|
||||
paths = list(expdir.glob(f"{split}_generations*.txt"))
|
||||
assert paths
|
||||
path = paths[0]
|
||||
lns = lmap(str.strip, list(path.open().readlines()))
|
||||
if n is not None:
|
||||
return lns[:n]
|
||||
return lns
|
||||
|
||||
|
||||
def load_cpu(p):
|
||||
return torch.load(p, map_location="cpu")
|
||||
|
||||
|
||||
class RougeTracker:
|
||||
def __init__(self, csv_path="rouge_test_df.csv", logdir=LOGDIR, data_dir=DATA_DIR):
|
||||
try:
|
||||
self.df = pd.read_csv(csv_path, index_col=0)
|
||||
except FileNotFoundError:
|
||||
self.df = pd.DataFrame()
|
||||
self.logdir = logdir
|
||||
test_gt = lmap(str.strip, Path(data_dir / "test.target").open().readlines())
|
||||
test_gt = lmap(str.strip, test_gt)
|
||||
self.gt = test_gt # {'test': test_gt}
|
||||
self.tokenizer = BartTokenizer.from_pretrained("facebook/bart-large-cnn")
|
||||
self.csv_path = csv_path
|
||||
self.new_results = pd.DataFrame()
|
||||
|
||||
@property
|
||||
def finished_experiments(self):
|
||||
return [p.parent.name for p in list(self.logdir.glob("*/test_generations*.txt"))]
|
||||
|
||||
def tok_len(self, strang):
|
||||
return len(self.tokenizer.encode(strang))
|
||||
|
||||
def read_all_gens(self,):
|
||||
GENS = {}
|
||||
for f in self.finished_experiments:
|
||||
GENS[f] = read_gens(f)
|
||||
return GENS
|
||||
|
||||
def score(self, gens, k, all_stats=False):
|
||||
rouge_raw = calculate_rouge(gens, self.gt, all_stats=all_stats)
|
||||
lens = np.mean(lmap(self.tok_len, gens))
|
||||
return dict(avg_len=lens, exp_name=k, **rouge_raw)
|
||||
|
||||
@property
|
||||
def new_experiments(self):
|
||||
possible = set(self.finished_experiments).difference(self.df.index)
|
||||
to_score = set()
|
||||
for p in possible:
|
||||
gens = read_gens(p)
|
||||
if len(gens) == len(self.gt):
|
||||
to_score.add(p)
|
||||
return to_score
|
||||
|
||||
def update(self):
|
||||
records = []
|
||||
to_score = self.new_experiments
|
||||
for exp_name in tqdm_nice(to_score, desc="Rouge Update"):
|
||||
gens = read_gens(exp_name)
|
||||
if len(gens) != len(self.gt):
|
||||
continue
|
||||
records.append(self.score(gens, exp_name))
|
||||
if not records:
|
||||
return self.df
|
||||
new_df = (
|
||||
pd.DataFrame(records).rename(columns=lambda x: x.replace("rouge", "R")).set_index("exp_name").astype(float)
|
||||
)
|
||||
self.new_results = new_df
|
||||
self.df = pd.concat([self.df, new_df]).dsort("R2")
|
||||
return self.df
|
||||
@@ -25,6 +25,8 @@ try:
|
||||
any_requires_grad,
|
||||
)
|
||||
from .finetune import main as ft_main
|
||||
from .replacement_scheduler import LinearReplacementScheduler
|
||||
|
||||
except ImportError:
|
||||
from finetune import SummarizationModule
|
||||
from finetune import main as ft_main
|
||||
@@ -37,6 +39,62 @@ except ImportError:
|
||||
assert_all_frozen,
|
||||
any_requires_grad,
|
||||
)
|
||||
from replacement_scheduler import LinearReplacementScheduler
|
||||
|
||||
|
||||
class TheseusDistiller(SummarizationModule):
|
||||
def __init__(self, hparams):
|
||||
|
||||
assert Path(hparams.data_dir).exists()
|
||||
# config = BartConfig.from_pretrained(hparams.model_name_or_path, student_encoder_layers=hparams.student_encoder_layers, student_decoder_layers=hparams.student_decoder_layers, replacing_rate=hparams.theseus_replace_rate)
|
||||
model = BartForConditionalGeneration.from_pretrained(
|
||||
hparams.model_name_or_path,
|
||||
student_encoder_layers=hparams.student_encoder_layers,
|
||||
student_decoder_layers=hparams.student_decoder_layers,
|
||||
replacing_rate=hparams.theseus_replace_rate,
|
||||
)
|
||||
super().__init__(hparams, model=model)
|
||||
self.different_encoder: bool = hparams.student_encoder_layers != self.model.config.encoder_layers
|
||||
self.different_decoder: bool = hparams.student_decoder_layers != self.model.config.decoder_layers
|
||||
|
||||
if hparams.theseus_init_copy:
|
||||
hparams.d_layers_to_copy = get_layers_to_copy(
|
||||
hparams.student_decoder_layers, self.model.config.decoder_layers, strategy=hparams.init_strategy,
|
||||
)
|
||||
hparams.e_layers_to_copy: List = get_layers_to_copy(
|
||||
hparams.student_encoder_layers, self.model.config.encoder_layers, strategy=hparams.init_strategy
|
||||
)
|
||||
if self.different_decoder:
|
||||
copy_layers(
|
||||
self.model.model.decoder.layers, self.model.model.decoder.scc_layers, hparams.d_layers_to_copy
|
||||
)
|
||||
if self.different_encoder:
|
||||
copy_layers(
|
||||
self.model.model.encoder.layers, self.model.model.encoder.scc_layers, hparams.e_layers_to_copy
|
||||
)
|
||||
else:
|
||||
hparams.e_layers_to_copy, hparams.d_layers_to_copy = None, None
|
||||
|
||||
self.replace_scheduler_encoder = LinearReplacementScheduler(self.model.model.encoder, 0.6)
|
||||
self.replace_scheduler_decoder = LinearReplacementScheduler(self.model.model.decoder, 0.6)
|
||||
freeze_params(self.model.model.encoder.layers) # Test
|
||||
freeze_params(self.model.model.decoder.layers)
|
||||
|
||||
def optimizer_step(self, *args, **kwargs) -> None:
|
||||
self.replace_scheduler_encoder.step()
|
||||
replace_rate = self.replace_scheduler_decoder.step()
|
||||
self.logger.log_metrics({"replace_rate": replace_rate})
|
||||
super().optimizer_step(*args, **kwargs)
|
||||
|
||||
def copy_to_student(self, d_layers_to_copy, e_layers_to_copy, hparams, student, teacher):
|
||||
if teacher.config.model_type == "t5":
|
||||
return self.copy_t5_to_student(d_layers_to_copy, e_layers_to_copy, hparams, student, teacher)
|
||||
self.different_encoder: bool = hparams.student_encoder_layers != teacher.config.encoder_layers
|
||||
self.different_decoder = hparams.student_decoder_layers != teacher.config.decoder_layers
|
||||
if self.different_decoder:
|
||||
copy_layers(teacher.model.decoder.layers, student.model.decoder.layers, d_layers_to_copy)
|
||||
if self.different_encoder:
|
||||
copy_layers(teacher.model.encoder.layers, student.model.encoder.layers, e_layers_to_copy)
|
||||
|
||||
|
||||
class SummarizationDistiller(SummarizationModule):
|
||||
@@ -49,6 +107,9 @@ class SummarizationDistiller(SummarizationModule):
|
||||
|
||||
super().__init__(hparams, model=student, config=student_cfg)
|
||||
self.teacher = teacher
|
||||
if isinstance(self.teacher, BartForConditionalGeneration):
|
||||
assert teacher.model.encoder.scc_layers is None
|
||||
assert self.model.model.encoder.scc_layers is None
|
||||
use_task_specific_params(self.teacher, "summarization")
|
||||
freeze_params(self.teacher)
|
||||
self.sanity_check_gradients()
|
||||
@@ -79,8 +140,14 @@ class SummarizationDistiller(SummarizationModule):
|
||||
"decoder_layers": hparams.student_decoder_layers,
|
||||
"encoder_layers": hparams.student_encoder_layers,
|
||||
}
|
||||
d_layers_to_copy = get_layers_to_copy(student_updates["decoder_layers"], teacher.config.decoder_layers)
|
||||
e_layers_to_copy: List = get_layers_to_copy(student_updates["encoder_layers"], teacher.config.encoder_layers)
|
||||
|
||||
d_layers_to_copy = get_layers_to_copy(
|
||||
student_updates["decoder_layers"], teacher.config.decoder_layers, strategy=hparams.init_strategy
|
||||
)
|
||||
e_layers_to_copy: List = get_layers_to_copy(
|
||||
student_updates["encoder_layers"], teacher.config.encoder_layers, strategy=hparams.init_strategy
|
||||
)
|
||||
|
||||
hparams.d_layer_to_copy = d_layers_to_copy
|
||||
hparams.e_layer_to_copy = e_layers_to_copy
|
||||
kw = teacher.config.to_diff_dict()
|
||||
@@ -180,18 +247,13 @@ class SummarizationDistiller(SummarizationModule):
|
||||
# parser.add_argument("--alpha_cos", default=0.0, type=float)
|
||||
parser.add_argument("--alpha_encoder_loss", default=0.0, type=float)
|
||||
parser.add_argument("--alpha_hid", default=0.0, type=float, required=False)
|
||||
parser.add_argument(
|
||||
"--student_decoder_layers", default=12, type=int, required=False,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--student_encoder_layers", default=12, type=int, required=False,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_teacher", action="store_true", default=False,
|
||||
)
|
||||
parser.add_argument( # TODO: remove
|
||||
"--enc_only", action="store_true", default=False,
|
||||
)
|
||||
|
||||
parser.add_argument("--student_decoder_layers", default=12, type=int, required=False)
|
||||
parser.add_argument("--student_encoder_layers", default=12, type=int, required=False)
|
||||
parser.add_argument("--no_teacher", action="store_true", default=False)
|
||||
parser.add_argument("--theseus_replace_rate", type=float, default=0.0)
|
||||
parser.add_argument("--theseus_init_copy", action="store_true")
|
||||
parser.add_argument("--init_strategy", type=str, default="alternate", choices=["alternate", "top", "bottom"])
|
||||
return parser
|
||||
|
||||
def _step(self, batch):
|
||||
@@ -378,13 +440,12 @@ class T5SummarizationDistiller(SummarizationDistiller):
|
||||
|
||||
def create_module(args):
|
||||
t5 = "t5" in args.model_name_or_path
|
||||
if args.no_teacher:
|
||||
assert not args.enc_only
|
||||
if args.no_teacher and args.theseus_replace_rate == 0:
|
||||
module_cls = SummarizationModule
|
||||
elif args.no_teacher and args.theseus_replace_rate > 0:
|
||||
module_cls = TheseusDistiller
|
||||
elif t5:
|
||||
module_cls = T5SummarizationDistiller
|
||||
elif args.enc_only:
|
||||
raise ValueError("Deleted that")
|
||||
else:
|
||||
module_cls = SummarizationDistiller
|
||||
args.setup_cls: str = module_cls.__name__
|
||||
@@ -415,21 +476,30 @@ def evaluate_checkpoint(ckpt_path: Path, dest_dir=None):
|
||||
trainer.test(model)
|
||||
|
||||
|
||||
def get_layers_to_copy(n_to_get, tot):
|
||||
all_layers = list(range(tot))
|
||||
if tot == 12: # Alternating for special cases
|
||||
layers_to_copy = { # maps # layers in student -> which teacher layers to copy
|
||||
6: [0, 2, 4, 7, 9, 11],
|
||||
1: [11],
|
||||
3: [0, 6, 11],
|
||||
2: [0, 11],
|
||||
4: [0, 4, 8, 11],
|
||||
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
|
||||
12: all_layers,
|
||||
}
|
||||
return layers_to_copy[n_to_get]
|
||||
DISTILBERT_ALTERNATE_PATTERN = { # maps # layers in student -> which teacher layers to copy
|
||||
6: [0, 2, 4, 7, 9, 11],
|
||||
1: [11],
|
||||
3: [0, 6, 11],
|
||||
2: [0, 11],
|
||||
4: [0, 4, 8, 11],
|
||||
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
|
||||
12: list(range(12)),
|
||||
}
|
||||
|
||||
|
||||
def get_layers_to_copy(n_student_layers: int, n_teacher_layers: int, strategy="alternate") -> List:
|
||||
all_layers = list(range(n_teacher_layers))
|
||||
if strategy == "alternate":
|
||||
if n_teacher_layers == 12:
|
||||
return DISTILBERT_ALTERNATE_PATTERN[n_student_layers]
|
||||
else:
|
||||
return all_layers[::2][:n_student_layers]
|
||||
elif strategy == "bottom":
|
||||
return all_layers[:n_student_layers]
|
||||
elif strategy == "top":
|
||||
return all_layers[-n_student_layers:]
|
||||
else:
|
||||
return all_layers[:n_to_get]
|
||||
raise ValueError(f"layer copy strategy {strategy} not supported")
|
||||
|
||||
|
||||
def distill_main(args):
|
||||
|
||||
@@ -9,12 +9,14 @@ from typing import Dict, List, Tuple
|
||||
import numpy as np
|
||||
import pytorch_lightning as pl
|
||||
import torch
|
||||
from pytorch_lightning.loggers import WandbLogger
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from lightning_base import BaseTransformer, add_generic_args, generic_train
|
||||
from transformers import get_linear_schedule_with_warmup
|
||||
|
||||
|
||||
WANDB_PROJ_NAME = "transformers_fork-examples_summarization_bart"
|
||||
try:
|
||||
from .utils import (
|
||||
use_task_specific_params,
|
||||
@@ -52,6 +54,7 @@ class SummarizationModule(BaseTransformer):
|
||||
loss_names = ["loss"]
|
||||
|
||||
def __init__(self, hparams, **kwargs):
|
||||
assert Path(hparams.data_dir).exists()
|
||||
super().__init__(hparams, num_labels=None, mode=self.mode, **kwargs)
|
||||
use_task_specific_params(self.model, "summarization")
|
||||
save_git_info(self.hparams.output_dir)
|
||||
@@ -79,8 +82,7 @@ class SummarizationModule(BaseTransformer):
|
||||
}
|
||||
assert self.target_lens["train"] <= self.target_lens["val"], f"target_lens: {self.target_lens}"
|
||||
assert self.target_lens["train"] <= self.target_lens["test"], f"target_lens: {self.target_lens}"
|
||||
|
||||
if self.hparams.freeze_embeds:
|
||||
if not self.hparams.unfreeze_embeds:
|
||||
self.freeze_embeds()
|
||||
if self.hparams.freeze_encoder:
|
||||
freeze_params(self.model.model.encoder) # TODO: this will break for t5
|
||||
@@ -253,9 +255,9 @@ class SummarizationModule(BaseTransformer):
|
||||
help="The input data dir. Should contain train.source, train.target, val.source, val.target, test.source, test.target",
|
||||
)
|
||||
parser.add_argument("--freeze_encoder", action="store_true")
|
||||
parser.add_argument("--freeze_embeds", action="store_true")
|
||||
parser.add_argument("--unfreeze_embeds", action="store_true")
|
||||
parser.add_argument("--sortish_sampler", action="store_true", default=False)
|
||||
parser.add_argument("--logger", type=str, choices=["default", "wandb", "wandb_shared"], default="default")
|
||||
parser.add_argument("--logger", type=str, choices=["default", "wandb", "wandb_shared"], default="wandb")
|
||||
parser.add_argument("--n_train", type=int, default=-1, required=False, help="# examples. -1 means use all.")
|
||||
parser.add_argument("--n_val", type=int, default=500, required=False, help="# examples. -1 means use all.")
|
||||
parser.add_argument("--n_test", type=int, default=-1, required=False, help="# examples. -1 means use all.")
|
||||
@@ -278,6 +280,10 @@ def main(args, model=None) -> SummarizationModule:
|
||||
elif args.logger == "wandb":
|
||||
from pytorch_lightning.loggers import WandbLogger
|
||||
|
||||
logger = WandbLogger(name=model.output_dir.name, project=WANDB_PROJ_NAME)
|
||||
elif args.logger == "wandb_shared":
|
||||
from pytorch_lightning.loggers import WandbLogger
|
||||
|
||||
logger = WandbLogger(name=model.output_dir.name)
|
||||
elif args.logger == "wandb_shared":
|
||||
from pytorch_lightning.loggers import WandbLogger
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
|
||||
# Add parent directory to python path to access lightning_base.py
|
||||
export PYTHONPATH="../":"${PYTHONPATH}"
|
||||
|
||||
@@ -6,6 +5,7 @@ export PYTHONPATH="../":"${PYTHONPATH}"
|
||||
# --model_name_or_path=t5-base for t5
|
||||
|
||||
# the proper usage is documented in the README
|
||||
|
||||
python finetune.py \
|
||||
--model_name_or_path=facebook/bart-large \
|
||||
--learning_rate=3e-5 \
|
||||
|
||||
@@ -0,0 +1,639 @@
|
||||
"""PyTorch BERT-of-Theseus model. """
|
||||
|
||||
from __future__ import absolute_import, division, print_function, unicode_literals
|
||||
|
||||
import logging
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributions.bernoulli import Bernoulli
|
||||
from torch.nn import CrossEntropyLoss, MSELoss
|
||||
|
||||
from transformers.configuration_bert import BertConfig
|
||||
from transformers.modeling_bert import (
|
||||
ACT2FN,
|
||||
BERT_PRETRAINED_MODEL_ARCHIVE_MAP,
|
||||
BertAttention,
|
||||
BertEmbeddings,
|
||||
BertIntermediate,
|
||||
BertLayer,
|
||||
BertLayerNorm,
|
||||
BertLMPredictionHead,
|
||||
BertOnlyMLMHead,
|
||||
BertOnlyNSPHead,
|
||||
BertOutput,
|
||||
BertPooler,
|
||||
BertPredictionHeadTransform,
|
||||
BertPreTrainingHeads,
|
||||
BertSelfAttention,
|
||||
BertSelfOutput,
|
||||
gelu,
|
||||
gelu_new,
|
||||
load_tf_weights_in_bert,
|
||||
mish,
|
||||
swish,
|
||||
)
|
||||
from transformers.modeling_utils import PreTrainedModel, prune_linear_layer
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BertEncoder(nn.Module):
|
||||
def __init__(self, config, scc_n_layer=6):
|
||||
super(BertEncoder, self).__init__()
|
||||
self.prd_n_layer = config.num_hidden_layers
|
||||
self.scc_n_layer = scc_n_layer
|
||||
assert self.prd_n_layer % self.scc_n_layer == 0
|
||||
self.compress_ratio = self.prd_n_layer // self.scc_n_layer
|
||||
self.bernoulli = None
|
||||
self.output_attentions = config.output_attentions
|
||||
self.output_hidden_states = config.output_hidden_states
|
||||
self.layer = nn.ModuleList([BertLayer(config) for _ in range(self.prd_n_layer)])
|
||||
self.scc_layer = nn.ModuleList([BertLayer(config) for _ in range(self.scc_n_layer)])
|
||||
|
||||
def set_replacing_rate(self, replacing_rate):
|
||||
if not 0 < replacing_rate <= 1:
|
||||
raise Exception("Replace rate must be in the range (0, 1]!")
|
||||
self.bernoulli = Bernoulli(torch.tensor([replacing_rate]))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
attention_mask=None,
|
||||
head_mask=None,
|
||||
encoder_hidden_states=None,
|
||||
encoder_attention_mask=None,
|
||||
):
|
||||
all_hidden_states = ()
|
||||
all_attentions = ()
|
||||
if self.training:
|
||||
inference_layers = []
|
||||
for i in range(self.scc_n_layer):
|
||||
if self.bernoulli.sample() == 1: # REPLACE
|
||||
inference_layers.append(self.scc_layer[i])
|
||||
else: # KEEP the original
|
||||
for offset in range(self.compress_ratio):
|
||||
inference_layers.append(self.layer[i * self.compress_ratio + offset])
|
||||
|
||||
else: # inference with compressed model
|
||||
inference_layers = self.scc_layer
|
||||
|
||||
for i, layer_module in enumerate(inference_layers):
|
||||
if self.output_hidden_states:
|
||||
all_hidden_states = all_hidden_states + (hidden_states,)
|
||||
|
||||
layer_outputs = layer_module(
|
||||
hidden_states, attention_mask, head_mask[i], encoder_hidden_states, encoder_attention_mask
|
||||
)
|
||||
hidden_states = layer_outputs[0]
|
||||
|
||||
if self.output_attentions:
|
||||
all_attentions = all_attentions + (layer_outputs[1],)
|
||||
|
||||
# Add last layer
|
||||
if self.output_hidden_states:
|
||||
all_hidden_states = all_hidden_states + (hidden_states,)
|
||||
|
||||
outputs = (hidden_states,)
|
||||
if self.output_hidden_states:
|
||||
outputs = outputs + (all_hidden_states,)
|
||||
if self.output_attentions:
|
||||
outputs = outputs + (all_attentions,)
|
||||
return outputs # last-layer hidden state, (all hidden states), (all attentions)
|
||||
|
||||
|
||||
class BertPreTrainedModel(PreTrainedModel):
|
||||
config_class = BertConfig
|
||||
pretrained_model_archive_map = BERT_PRETRAINED_MODEL_ARCHIVE_MAP
|
||||
load_tf_weights = load_tf_weights_in_bert
|
||||
base_model_prefix = "bert"
|
||||
|
||||
def _init_weights(self, module):
|
||||
""" Initialize the weights """
|
||||
if isinstance(module, (nn.Linear, nn.Embedding)):
|
||||
# Slightly different from the TF version which uses truncated_normal for initialization
|
||||
# cf https://github.com/pytorch/pytorch/pull/5617
|
||||
module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)
|
||||
elif isinstance(module, BertLayerNorm):
|
||||
module.bias.data.zero_()
|
||||
module.weight.data.fill_(1.0)
|
||||
if isinstance(module, nn.Linear) and module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
|
||||
|
||||
class BertModel(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertModel, self).__init__(config)
|
||||
self.config = config
|
||||
|
||||
self.embeddings = BertEmbeddings(config)
|
||||
self.encoder = BertEncoder(config)
|
||||
self.pooler = BertPooler(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.embeddings.word_embeddings
|
||||
|
||||
def set_input_embeddings(self, value):
|
||||
self.embeddings.word_embeddings = value
|
||||
|
||||
def _prune_heads(self, heads_to_prune):
|
||||
""" Prunes heads of the model.
|
||||
heads_to_prune: dict of {layer_num: list of heads to prune in this layer}
|
||||
See base class PreTrainedModel
|
||||
"""
|
||||
for layer, heads in heads_to_prune.items():
|
||||
self.encoder.layer[layer].attention.prune_heads(heads)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
encoder_hidden_states=None,
|
||||
encoder_attention_mask=None,
|
||||
):
|
||||
if input_ids is not None and inputs_embeds is not None:
|
||||
raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
|
||||
elif input_ids is not None:
|
||||
input_shape = input_ids.size()
|
||||
elif inputs_embeds is not None:
|
||||
input_shape = inputs_embeds.size()[:-1]
|
||||
else:
|
||||
raise ValueError("You have to specify either input_ids or inputs_embeds")
|
||||
|
||||
device = input_ids.device if input_ids is not None else inputs_embeds.device
|
||||
|
||||
if attention_mask is None:
|
||||
attention_mask = torch.ones(input_shape, device=device)
|
||||
if token_type_ids is None:
|
||||
token_type_ids = torch.zeros(input_shape, dtype=torch.long, device=device)
|
||||
|
||||
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
|
||||
# ourselves in which case we just need to make it broadcastable to all heads.
|
||||
if attention_mask.dim() == 3:
|
||||
extended_attention_mask = attention_mask[:, None, :, :]
|
||||
elif attention_mask.dim() == 2:
|
||||
# Provided a padding mask of dimensions [batch_size, seq_length]
|
||||
# - if the model is a decoder, apply a causal mask in addition to the padding mask
|
||||
# - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
||||
if self.config.is_decoder:
|
||||
batch_size, seq_length = input_shape
|
||||
seq_ids = torch.arange(seq_length, device=device)
|
||||
causal_mask = seq_ids[None, None, :].repeat(batch_size, seq_length, 1) <= seq_ids[None, :, None]
|
||||
causal_mask = causal_mask.to(
|
||||
torch.long
|
||||
) # not converting to long will cause errors with pytorch version < 1.3
|
||||
extended_attention_mask = causal_mask[:, None, :, :] * attention_mask[:, None, None, :]
|
||||
else:
|
||||
extended_attention_mask = attention_mask[:, None, None, :]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Wrong shape for input_ids (shape {}) or attention_mask (shape {})".format(
|
||||
input_shape, attention_mask.shape
|
||||
)
|
||||
)
|
||||
|
||||
# Since attention_mask is 1.0 for positions we want to attend and 0.0 for
|
||||
# masked positions, this operation will create a tensor which is 0.0 for
|
||||
# positions we want to attend and -10000.0 for masked positions.
|
||||
# Since we are adding it to the raw scores before the softmax, this is
|
||||
# effectively the same as removing these entirely.
|
||||
extended_attention_mask = extended_attention_mask.to(dtype=next(self.parameters()).dtype) # fp16 compatibility
|
||||
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
|
||||
|
||||
# If a 2D ou 3D attention mask is provided for the cross-attention
|
||||
# we need to make broadcastabe to [batch_size, num_heads, seq_length, seq_length]
|
||||
if self.config.is_decoder and encoder_hidden_states is not None:
|
||||
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
|
||||
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
|
||||
if encoder_attention_mask is None:
|
||||
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
|
||||
|
||||
if encoder_attention_mask.dim() == 3:
|
||||
encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :]
|
||||
elif encoder_attention_mask.dim() == 2:
|
||||
encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Wrong shape for encoder_hidden_shape (shape {}) or encoder_attention_mask (shape {})".format(
|
||||
encoder_hidden_shape, encoder_attention_mask.shape
|
||||
)
|
||||
)
|
||||
|
||||
encoder_extended_attention_mask = encoder_extended_attention_mask.to(
|
||||
dtype=next(self.parameters()).dtype
|
||||
) # fp16 compatibility
|
||||
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -10000.0
|
||||
else:
|
||||
encoder_extended_attention_mask = None
|
||||
|
||||
# Prepare head mask if needed
|
||||
# 1.0 in head_mask indicate we keep the head
|
||||
# attention_probs has shape bsz x n_heads x N x N
|
||||
# input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
|
||||
# and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
|
||||
if head_mask is not None:
|
||||
if head_mask.dim() == 1:
|
||||
head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
head_mask = head_mask.expand(self.config.num_hidden_layers, -1, -1, -1, -1)
|
||||
elif head_mask.dim() == 2:
|
||||
head_mask = (
|
||||
head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1)
|
||||
) # We can specify head_mask for each layer
|
||||
head_mask = head_mask.to(
|
||||
dtype=next(self.parameters()).dtype
|
||||
) # switch to fload if need + fp16 compatibility
|
||||
else:
|
||||
head_mask = [None] * self.config.num_hidden_layers
|
||||
|
||||
embedding_output = self.embeddings(
|
||||
input_ids=input_ids, position_ids=position_ids, token_type_ids=token_type_ids, inputs_embeds=inputs_embeds
|
||||
)
|
||||
encoder_outputs = self.encoder(
|
||||
embedding_output,
|
||||
attention_mask=extended_attention_mask,
|
||||
head_mask=head_mask,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_extended_attention_mask,
|
||||
)
|
||||
sequence_output = encoder_outputs[0]
|
||||
pooled_output = self.pooler(sequence_output)
|
||||
|
||||
outputs = (sequence_output, pooled_output,) + encoder_outputs[
|
||||
1:
|
||||
] # add hidden_states and attentions if they are here
|
||||
return outputs # sequence_output, pooled_output, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForPreTraining(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForPreTraining, self).__init__(config)
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.cls = BertPreTrainingHeads(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.cls.predictions.decoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
masked_lm_labels=None,
|
||||
next_sentence_label=None,
|
||||
):
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
sequence_output, pooled_output = outputs[:2]
|
||||
prediction_scores, seq_relationship_score = self.cls(sequence_output, pooled_output)
|
||||
|
||||
outputs = (prediction_scores, seq_relationship_score,) + outputs[
|
||||
2:
|
||||
] # add hidden states and attention if they are here
|
||||
|
||||
if masked_lm_labels is not None and next_sentence_label is not None:
|
||||
loss_fct = CrossEntropyLoss(ignore_index=-1)
|
||||
masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), masked_lm_labels.view(-1))
|
||||
next_sentence_loss = loss_fct(seq_relationship_score.view(-1, 2), next_sentence_label.view(-1))
|
||||
total_loss = masked_lm_loss + next_sentence_loss
|
||||
outputs = (total_loss,) + outputs
|
||||
|
||||
return outputs # (loss), prediction_scores, seq_relationship_score, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForMaskedLM(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForMaskedLM, self).__init__(config)
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.cls = BertOnlyMLMHead(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_output_embeddings(self):
|
||||
return self.cls.predictions.decoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
masked_lm_labels=None,
|
||||
encoder_hidden_states=None,
|
||||
encoder_attention_mask=None,
|
||||
lm_labels=None,
|
||||
):
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
|
||||
sequence_output = outputs[0]
|
||||
prediction_scores = self.cls(sequence_output)
|
||||
|
||||
outputs = (prediction_scores,) + outputs[2:] # Add hidden states and attention if they are here
|
||||
|
||||
# Although this may seem awkward, BertForMaskedLM supports two scenarios:
|
||||
# 1. If a tensor that contains the indices of masked labels is provided,
|
||||
# the cross-entropy is the MLM cross-entropy that measures the likelihood
|
||||
# of predictions for masked words.
|
||||
# 2. If `lm_labels` is provided we are in a causal scenario where we
|
||||
# try to predict the next token for each input in the decoder.
|
||||
if masked_lm_labels is not None:
|
||||
loss_fct = CrossEntropyLoss(ignore_index=-1) # -1 index = padding token
|
||||
masked_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), masked_lm_labels.view(-1))
|
||||
outputs = (masked_lm_loss,) + outputs
|
||||
|
||||
if lm_labels is not None:
|
||||
# we are doing next-token prediction; shift prediction scores and input ids by one
|
||||
prediction_scores = prediction_scores[:, :-1, :].contiguous()
|
||||
lm_labels = lm_labels[:, 1:].contiguous()
|
||||
loss_fct = CrossEntropyLoss(ignore_index=-1)
|
||||
ltr_lm_loss = loss_fct(prediction_scores.view(-1, self.config.vocab_size), lm_labels.view(-1))
|
||||
outputs = (ltr_lm_loss,) + outputs
|
||||
|
||||
return outputs # (masked_lm_loss), (ltr_lm_loss), prediction_scores, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForNextSentencePrediction(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForNextSentencePrediction, self).__init__(config)
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.cls = BertOnlyNSPHead(config)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
next_sentence_label=None,
|
||||
):
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
pooled_output = outputs[1]
|
||||
|
||||
seq_relationship_score = self.cls(pooled_output)
|
||||
|
||||
outputs = (seq_relationship_score,) + outputs[2:] # add hidden states and attention if they are here
|
||||
if next_sentence_label is not None:
|
||||
loss_fct = CrossEntropyLoss(ignore_index=-1)
|
||||
next_sentence_loss = loss_fct(seq_relationship_score.view(-1, 2), next_sentence_label.view(-1))
|
||||
outputs = (next_sentence_loss,) + outputs
|
||||
|
||||
return outputs # (next_sentence_loss), seq_relationship_score, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForSequenceClassification(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForSequenceClassification, self).__init__(config)
|
||||
self.num_labels = config.num_labels
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
||||
self.classifier = nn.Linear(config.hidden_size, self.config.num_labels)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
labels=None,
|
||||
):
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
pooled_output = outputs[1]
|
||||
|
||||
pooled_output = self.dropout(pooled_output)
|
||||
logits = self.classifier(pooled_output)
|
||||
|
||||
outputs = (logits,) + outputs[2:] # add hidden states and attention if they are here
|
||||
|
||||
if labels is not None:
|
||||
if self.num_labels == 1:
|
||||
# We are doing regression
|
||||
loss_fct = MSELoss()
|
||||
loss = loss_fct(logits.view(-1), labels.view(-1))
|
||||
else:
|
||||
loss_fct = CrossEntropyLoss()
|
||||
loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
|
||||
outputs = (loss,) + outputs
|
||||
|
||||
return outputs # (loss), logits, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForMultipleChoice(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForMultipleChoice, self).__init__(config)
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
||||
self.classifier = nn.Linear(config.hidden_size, 1)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
labels=None,
|
||||
):
|
||||
num_choices = input_ids.shape[1]
|
||||
|
||||
input_ids = input_ids.view(-1, input_ids.size(-1))
|
||||
attention_mask = attention_mask.view(-1, attention_mask.size(-1)) if attention_mask is not None else None
|
||||
token_type_ids = token_type_ids.view(-1, token_type_ids.size(-1)) if token_type_ids is not None else None
|
||||
position_ids = position_ids.view(-1, position_ids.size(-1)) if position_ids is not None else None
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
pooled_output = outputs[1]
|
||||
|
||||
pooled_output = self.dropout(pooled_output)
|
||||
logits = self.classifier(pooled_output)
|
||||
reshaped_logits = logits.view(-1, num_choices)
|
||||
|
||||
outputs = (reshaped_logits,) + outputs[2:] # add hidden states and attention if they are here
|
||||
|
||||
if labels is not None:
|
||||
loss_fct = CrossEntropyLoss()
|
||||
loss = loss_fct(reshaped_logits, labels)
|
||||
outputs = (loss,) + outputs
|
||||
|
||||
return outputs # (loss), reshaped_logits, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForTokenClassification(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForTokenClassification, self).__init__(config)
|
||||
self.num_labels = config.num_labels
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
||||
self.classifier = nn.Linear(config.hidden_size, config.num_labels)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
labels=None,
|
||||
):
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
sequence_output = outputs[0]
|
||||
|
||||
sequence_output = self.dropout(sequence_output)
|
||||
logits = self.classifier(sequence_output)
|
||||
|
||||
outputs = (logits,) + outputs[2:] # add hidden states and attention if they are here
|
||||
if labels is not None:
|
||||
loss_fct = CrossEntropyLoss()
|
||||
# Only keep active parts of the loss
|
||||
if attention_mask is not None:
|
||||
active_loss = attention_mask.view(-1) == 1
|
||||
active_logits = logits.view(-1, self.num_labels)[active_loss]
|
||||
active_labels = labels.view(-1)[active_loss]
|
||||
loss = loss_fct(active_logits, active_labels)
|
||||
else:
|
||||
loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
|
||||
outputs = (loss,) + outputs
|
||||
|
||||
return outputs # (loss), scores, (hidden_states), (attentions)
|
||||
|
||||
|
||||
class BertForQuestionAnswering(BertPreTrainedModel):
|
||||
def __init__(self, config):
|
||||
super(BertForQuestionAnswering, self).__init__(config)
|
||||
self.num_labels = config.num_labels
|
||||
|
||||
self.bert = BertModel(config)
|
||||
self.qa_outputs = nn.Linear(config.hidden_size, config.num_labels)
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids=None,
|
||||
attention_mask=None,
|
||||
token_type_ids=None,
|
||||
position_ids=None,
|
||||
head_mask=None,
|
||||
inputs_embeds=None,
|
||||
start_positions=None,
|
||||
end_positions=None,
|
||||
):
|
||||
|
||||
outputs = self.bert(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
token_type_ids=token_type_ids,
|
||||
position_ids=position_ids,
|
||||
head_mask=head_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
)
|
||||
|
||||
sequence_output = outputs[0]
|
||||
|
||||
logits = self.qa_outputs(sequence_output)
|
||||
start_logits, end_logits = logits.split(1, dim=-1)
|
||||
start_logits = start_logits.squeeze(-1)
|
||||
end_logits = end_logits.squeeze(-1)
|
||||
|
||||
outputs = (start_logits, end_logits,) + outputs[2:]
|
||||
if start_positions is not None and end_positions is not None:
|
||||
# If we are on multi-GPU, split add a dimension
|
||||
if len(start_positions.size()) > 1:
|
||||
start_positions = start_positions.squeeze(-1)
|
||||
if len(end_positions.size()) > 1:
|
||||
end_positions = end_positions.squeeze(-1)
|
||||
# sometimes the start/end positions are outside our model inputs, we ignore these terms
|
||||
ignored_index = start_logits.size(1)
|
||||
start_positions.clamp_(0, ignored_index)
|
||||
end_positions.clamp_(0, ignored_index)
|
||||
|
||||
loss_fct = CrossEntropyLoss(ignore_index=ignored_index)
|
||||
start_loss = loss_fct(start_logits, start_positions)
|
||||
end_loss = loss_fct(end_logits, end_positions)
|
||||
total_loss = (start_loss + end_loss) / 2
|
||||
outputs = (total_loss,) + outputs
|
||||
|
||||
return outputs # (loss), start_logits, end_logits, (hidden_states), (attentions)
|
||||
@@ -0,0 +1,32 @@
|
||||
class ConstantReplacementScheduler:
|
||||
def __init__(self, module, replacing_rate, replacing_steps=None):
|
||||
self.module = module
|
||||
self.replacing_rate = replacing_rate
|
||||
self.replacing_steps = replacing_steps
|
||||
self.step_counter = 0
|
||||
self.module.set_replacing_rate(replacing_rate)
|
||||
|
||||
def step(self):
|
||||
self.step_counter += 1
|
||||
if self.replacing_steps is None or self.replacing_rate == 1.0:
|
||||
return self.replacing_rate
|
||||
else:
|
||||
if self.step_counter >= self.replacing_steps:
|
||||
self.module.set_replacing_rate(1.0)
|
||||
self.replacing_rate = 1.0
|
||||
return self.replacing_rate
|
||||
|
||||
|
||||
class LinearReplacementScheduler:
|
||||
def __init__(self, module, base_replacing_rate, k=1e-4):
|
||||
self.module = module
|
||||
self.base_replacing_rate = base_replacing_rate
|
||||
self.step_counter = 0
|
||||
self.k = k
|
||||
self.module.set_replacing_rate(base_replacing_rate)
|
||||
|
||||
def step(self):
|
||||
self.step_counter += 1
|
||||
current_replacing_rate = min(self.k * self.step_counter + self.base_replacing_rate, 1.0)
|
||||
self.module.set_replacing_rate(current_replacing_rate)
|
||||
return current_replacing_rate
|
||||
@@ -7,5 +7,6 @@ python distillation.py \
|
||||
--learning_rate=3e-4 \
|
||||
--do_train \
|
||||
--do_predict \
|
||||
--fp16 \
|
||||
--val_check_interval 0.1 \
|
||||
$@
|
||||
|
||||
@@ -5,7 +5,7 @@ from pathlib import Path
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
||||
from transformers import AutoModelWithLMHead, AutoTokenizer
|
||||
|
||||
|
||||
try:
|
||||
@@ -23,13 +23,10 @@ def chunks(lst, n):
|
||||
|
||||
|
||||
def generate_summaries(
|
||||
examples: list, out_file: str, model_name: str, batch_size: int = 8, device: str = DEFAULT_DEVICE, fp16=False,
|
||||
) -> None:
|
||||
examples: list, out_file: str, model_name: str, batch_size: int = 8, device: str = DEFAULT_DEVICE
|
||||
):
|
||||
fout = Path(out_file).open("w", encoding="utf-8")
|
||||
model_name = str(model_name)
|
||||
model = AutoModelForSeq2SeqLM.from_pretrained(model_name).to(device)
|
||||
if fp16:
|
||||
model = model.half()
|
||||
model = AutoModelWithLMHead.from_pretrained(model_name).to(device)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
|
||||
@@ -52,20 +49,32 @@ def generate_summaries(
|
||||
|
||||
def run_generate():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("input_path", type=str, help="like cnn_dm/test.source")
|
||||
parser.add_argument("output_path", type=str, help="where to save summaries")
|
||||
parser.add_argument("model_name", type=str, help="like facebook/bart-large-cnn,t5-base, etc.")
|
||||
parser.add_argument(
|
||||
"input_path", type=str, help="like cnn_dm/test.source",
|
||||
)
|
||||
parser.add_argument(
|
||||
"output_path", type=str, help="where to save summaries",
|
||||
)
|
||||
parser.add_argument(
|
||||
"model_name",
|
||||
type=str,
|
||||
default="facebook/bart-large-cnn",
|
||||
help="like facebook/bart-large-cnn,'t5-small', 't5-base', 't5-large', 't5-3b', 't5-11b",
|
||||
)
|
||||
parser.add_argument("--reference_path", type=str, required=False, help="like cnn_dm/test_reference_summaries.txt")
|
||||
parser.add_argument("--score_path", type=str, required=False, help="where to save the rouge score in json format")
|
||||
parser.add_argument("--device", type=str, required=False, default=DEFAULT_DEVICE, help="cuda, cuda:1, cpu etc.")
|
||||
parser.add_argument("--bs", type=int, default=8, required=False, help="batch size")
|
||||
parser.add_argument("--fp16", action="store_true")
|
||||
parser.add_argument(
|
||||
"--score_path", type=str, required=False, help="where to save the rouge score in json format",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device", type=str, required=False, default=DEFAULT_DEVICE, help="cuda, cuda:1, cpu etc.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--bs", type=int, default=8, required=False, help="batch size: how many to summarize at a time",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
examples = [" " + x.rstrip() if "t5" in args.model_name else x.rstrip() for x in open(args.input_path).readlines()]
|
||||
|
||||
generate_summaries(
|
||||
examples, args.output_path, args.model_name, batch_size=args.bs, device=args.device, fp16=args.fp16
|
||||
)
|
||||
generate_summaries(examples, args.output_path, args.model_name, batch_size=args.bs, device=args.device)
|
||||
if args.score_path is not None:
|
||||
output_lns = [x.rstrip() for x in open(args.output_path).readlines()]
|
||||
reference_lns = [x.rstrip() for x in open(args.reference_path).readlines()]
|
||||
|
||||
@@ -23,10 +23,12 @@ logging.basicConfig(level=logging.DEBUG)
|
||||
logger = logging.getLogger()
|
||||
FP16_EVER = False
|
||||
CHEAP_ARGS = {
|
||||
"theseus_init_copy": False,
|
||||
"theseus_replace_rate": 0,
|
||||
"logger": "default",
|
||||
"num_workers": 2,
|
||||
"alpha_hid": 0,
|
||||
"freeze_embeds": True,
|
||||
"unfreeze_embeds": False,
|
||||
"enc_only": False,
|
||||
"tgt_suffix": "",
|
||||
"resume_from_checkpoint": None,
|
||||
@@ -72,6 +74,8 @@ CHEAP_ARGS = {
|
||||
"alpha_loss_encoder": 0.0,
|
||||
"freeze_encoder": False,
|
||||
"auto_scale_batch_size": False,
|
||||
"freeze_decoder": False,
|
||||
"init_strategy": "bottom",
|
||||
}
|
||||
|
||||
|
||||
@@ -80,7 +84,6 @@ def _dump_articles(path: Path, articles: list):
|
||||
f.write("\n".join(articles))
|
||||
|
||||
|
||||
MSG = "T5 is broken at the moment"
|
||||
T5_TINY = "patrickvonplaten/t5-tiny-random"
|
||||
|
||||
|
||||
@@ -109,8 +112,6 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
freeze_encoder=True,
|
||||
gpus=2,
|
||||
sortish_sampler=False,
|
||||
fp16_opt_level="O1",
|
||||
fp16=FP16_EVER,
|
||||
)
|
||||
self._bart_distiller_cli(updates)
|
||||
|
||||
@@ -132,10 +133,6 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
updates = dict(student_encoder_layers=2, student_decoder_layers=1, no_teacher=True,)
|
||||
self._bart_distiller_cli(updates)
|
||||
|
||||
def test_bdc_yes_teacher(self):
|
||||
updates = dict(student_encoder_layers=2, student_decoder_layers=1,)
|
||||
self._bart_distiller_cli(updates)
|
||||
|
||||
def test_bdc_checkpointing(self):
|
||||
updates = dict(
|
||||
student_encoder_layers=2,
|
||||
@@ -143,6 +140,7 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
num_train_epochs=4,
|
||||
val_check_interval=0.25,
|
||||
alpha_hid=2.0,
|
||||
init_strategy="alternate",
|
||||
)
|
||||
model = self._bart_distiller_cli(updates, check_contents=False)
|
||||
|
||||
@@ -154,11 +152,31 @@ class TestSummarizationDistiller(unittest.TestCase):
|
||||
self.assertEqual(len(new_transformer_ckpts), 1)
|
||||
examples = lmap(str.strip, model.hparams.data_dir.joinpath("test.source").open().readlines())
|
||||
out_path = tempfile.mktemp()
|
||||
generate_summaries(examples, out_path, new_transformer_ckpts[0].parent)
|
||||
generate_summaries(examples, out_path, model_name=str(new_transformer_ckpts[0].parent))
|
||||
self.assertTrue(Path(out_path).exists())
|
||||
|
||||
evaluate_checkpoint(ckpts[0], dest_dir=Path(tempfile.mkdtemp()))
|
||||
|
||||
def test_bdc_theseus(self):
|
||||
updates = dict(
|
||||
theseus_replace_rate=0.5,
|
||||
student_encoder_layers=1,
|
||||
student_decoder_layers=1,
|
||||
no_teacher=True,
|
||||
theseus_init_copy=True,
|
||||
)
|
||||
self._bart_distiller_cli(updates)
|
||||
|
||||
def test_bdc_frozen_theseus(self):
|
||||
updates = dict(
|
||||
theseus_replace_rate=0.5,
|
||||
student_encoder_layers=2,
|
||||
student_decoder_layers=1,
|
||||
no_teacher=True,
|
||||
freeze_encoder=True,
|
||||
theseus_init_copy=True,
|
||||
)
|
||||
self._bart_distiller_cli(updates)
|
||||
|
||||
def _bart_distiller_cli(self, updates, check_contents=True):
|
||||
default_updates = dict(
|
||||
train_batch_size=1,
|
||||
|
||||
@@ -13,6 +13,8 @@ from torch import nn
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from transformers import BartTokenizer
|
||||
|
||||
|
||||
def encode_file(
|
||||
tokenizer,
|
||||
@@ -30,7 +32,6 @@ def encode_file(
|
||||
examples = torch.load(cache_path)
|
||||
assert isinstance(examples, list)
|
||||
return examples
|
||||
|
||||
except Exception:
|
||||
print(f"failed to load from {cache_path}, retokenizing {data_path}")
|
||||
data_path = Path(data_path)
|
||||
@@ -85,7 +86,7 @@ class SummarizationDataset(Dataset):
|
||||
prefix="",
|
||||
):
|
||||
super().__init__()
|
||||
tok_name = tokenizer.__class__.__name__.lower().rstrip("tokenizer")
|
||||
tok_name = "T5" if not isinstance(tokenizer, BartTokenizer) else ""
|
||||
self.source = encode_file(
|
||||
tokenizer,
|
||||
os.path.join(data_dir, type_path + ".source"),
|
||||
@@ -98,6 +99,7 @@ class SummarizationDataset(Dataset):
|
||||
self.target = encode_file(
|
||||
tokenizer, tgt_path, max_target_length, overwrite_cache=overwrite_cache, tok_name=tok_name
|
||||
)
|
||||
|
||||
if n_obs is not None:
|
||||
self.source = self.source[:n_obs]
|
||||
self.target = self.target[:n_obs]
|
||||
@@ -212,7 +214,7 @@ def get_git_info():
|
||||
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL"]
|
||||
|
||||
|
||||
def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
|
||||
def calculate_rouge(output_lns: List[str], reference_lns: List[str], all_stats=False):
|
||||
scorer = rouge_scorer.RougeScorer(ROUGE_KEYS, use_stemmer=True)
|
||||
aggregator = scoring.BootstrapAggregator()
|
||||
|
||||
@@ -221,7 +223,11 @@ def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
|
||||
aggregator.add_scores(scores)
|
||||
|
||||
result = aggregator.aggregate()
|
||||
return {k: v.mid.fmeasure for k, v in result.items()}
|
||||
|
||||
if all_stats:
|
||||
return expanded_rouge_df(result)
|
||||
else:
|
||||
return {k: v.mid.fmeasure for k, v in result.items()}
|
||||
|
||||
|
||||
def freeze_params(model: nn.Module):
|
||||
@@ -241,10 +247,36 @@ def assert_all_frozen(model):
|
||||
model_grads: List[bool] = list(grad_status(model))
|
||||
n_require_grad = sum(lmap(int, model_grads))
|
||||
npars = len(model_grads)
|
||||
assert not any(model_grads), f"{n_require_grad/npars:.1%} of {npars} weights require grad"
|
||||
assert not any(model_grads), f"{n_require_grad / npars:.1%} of {npars} weights require grad"
|
||||
|
||||
|
||||
def assert_not_all_frozen(model):
|
||||
model_grads: List[bool] = list(grad_status(model))
|
||||
npars = len(model_grads)
|
||||
assert any(model_grads), f"none of {npars} weights require grad"
|
||||
|
||||
|
||||
def dictify(rouge_obj) -> List:
|
||||
records = []
|
||||
for k, rouge_measurement in rouge_obj.items():
|
||||
if k == "rouge1":
|
||||
continue
|
||||
for k1 in ["low", "mid", "high"]:
|
||||
if k1 != "mid":
|
||||
continue
|
||||
v1 = getattr(rouge_measurement, k1)
|
||||
for k2 in ["precision", "recall", "fmeasure"]:
|
||||
records.append([k, k1, k2, getattr(v1, k2)])
|
||||
|
||||
return records
|
||||
|
||||
|
||||
def expanded_rouge_df(rouge_all):
|
||||
import pandas as pd
|
||||
|
||||
return (
|
||||
pd.DataFrame(dictify(rouge_all), columns=["metric", "k1", "k2", "val"])
|
||||
.set_index(["metric", "k2"])["val"]
|
||||
.unstack("metric")
|
||||
.rename_axis(None)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
exp_name,avg_len,R2_56,R1,R2,RL
|
||||
dl6_no_teacher,89.17162750217581,,0.4426298740974814,0.21214627660372445,0.30352876002780776
|
||||
cnn_f12_9_noteach,93.90400348128807,,0.44021766057145817,0.2112684144707048,0.30203109862281485
|
||||
baseline_140,83.3504,0.2012,0.4405,0.2106,0.3063
|
||||
dl6_yes_teacher,65.3845953002611,,0.4358592269823179,0.2094872879851778,0.3055267315672064
|
||||
cnn_12_9_no_teacher,65.63707571801567,,0.4349178365295365,0.20912713468933475,0.30577661504959625
|
||||
cnn_f12_9,66.25691906005223,,0.4360757701741097,0.20874355812934306,0.30534340603938714
|
||||
dl6_ckpt,64.7923,0.211,0.4344,0.2087,0.3049
|
||||
run_6l,64.7923,0.2109,0.4344,0.2086,0.3049
|
||||
dl6_4ep,66.4003,0.2112,0.4351,0.2082,0.3047
|
||||
cnn_9_9_yes_teacher,65.75987815491732,,0.43407500508924857,0.2077387062275759,0.3041112175512638
|
||||
dl6_maxlen_140_gacc2,65.2891,0.2099,0.433,0.2072,0.3038
|
||||
brewer_cnn_12_6_v2,58.29869451697128,,0.4275596591964453,0.2061398841242274,0.3043299614540603
|
||||
cnn_enc_only_6_12_brutasse,78.37284595300261,,0.4352447538697767,0.2057028264435689,0.3015751874385552
|
||||
cnn_no_teacher_f12_3,84.8461270670148,,0.4341908103753934,0.20524584699497528,0.29823737701046504
|
||||
cnn_6_6_yes_teacher,66.56605744125326,,0.4271799135757925,0.2017168386739932,0.2971576748941705
|
||||
el6_12_fix,94.3361183637946,,0.4295203853115971,0.20114288940092984,0.29082690956646995
|
||||
dl3,56.8062,0.2047,0.4197,0.2002,0.2989
|
||||
dl3_4ep,56.585,0.2057,0.4192,0.1998,0.3
|
||||
pseudo_dd_140,102.0577,0.1864,0.4254,0.1997,0.2886
|
||||
dl6_4ep_sasha_ds,60.0613,0.2082,0.4217,0.1993,0.2965
|
||||
cnn_enc_only_3_12_brutasse,75.5237597911227,,0.4278349430460213,0.19820455870913525,0.2945791235929408
|
||||
pseudo_6l_140_fix,62.5668,0.2035,0.4197,0.1968,0.2941
|
||||
staged_cnn_6_6,89.18093994778067,,0.4255412775299003,0.19658951474224035,0.2885276829133131
|
||||
el9_12,66.26144473455179,,0.42086970025135134,0.1948512909973004,0.29011441089763124
|
||||
el9_9,65.4414273281114,,0.419385405970708,0.1941718068172067,0.2900372909854251
|
||||
dl2_4ep,56.013054830287196,,0.4097263081894015,0.1919076904112108,0.29258786794511826
|
||||
dl3_4ep_sasha_ds,56.9671,0.2007,0.4111,0.1912,0.2914
|
||||
dl2_4ep_sasha_ds,55.999129677980854,,0.4057876788402908,0.1871524671665716,0.28735289627716404
|
||||
brewer_cnn_12_1_v4,56.39904264577894,,0.4031071697939543,0.18464802550618464,0.2839190569931007
|
||||
blarge_12_6_no_teacher,98.89782419495216,,0.4056381963795483,0.1826793955605352,0.2700670904734034
|
||||
2l_96,56.0208,0.1808,0.395,0.1799,0.2799
|
||||
2l_140,56.0208,0.1808,0.3951,0.1799,0.2798
|
||||
2l_56,56.0336,0.1762,0.3895,0.1758,0.2762
|
||||
2l_eval_140,56.0336,0.1762,0.3897,0.1757,0.2762
|
||||
el6_12,66.68172323759791,,0.4013734110868821,0.1752702303482339,0.2711560926634892
|
||||
el6_6,57.5147954743255,,0.38798696359999946,0.1672810503629001,0.2648421408851355
|
||||
dl1,56.0179,0.1264,0.3411,0.1285,0.226
|
||||
last_layer,55.6623,0.0596,0.2654,0.0616,0.1764
|
||||
cnn_enc_only_9_12_brutasse,122.32297650130549,,0.12110269669223532,0.011040620587715645,0.09900079897379864
|
||||
|
@@ -0,0 +1,30 @@
|
||||
exp_name,avg_len,R1,R2,RL
|
||||
brewer_xsum_12_6,27.251213270978557,0.45257522039021847,0.2195833165680984,0.3679937489447971
|
||||
xsum_baseline,29.06997264625431,0.45234885716476203,0.21848808526257507,0.3650268517677788
|
||||
xsum_f12_9,29.35074561016501,0.451213614023561,0.2163378899615816,0.363539265208816
|
||||
brewer_xsum_9_6,27.35771640342363,0.4492537548158509,0.2162358630026469,0.3641728953722555
|
||||
xsum_no_teacher_12_6,27.75319862348893,0.4486508407826425,0.2149919319149121,0.36419730252620985
|
||||
xsum_clean_baseline,29.314656313420983,0.4477747607857255,0.21439354376130093,0.3593759295843251
|
||||
brewer_xsum_12_3,24.899232330362658,0.44368280928513437,0.21344371120165062,0.3630886297392266
|
||||
xsum_f12_9_noteach,29.46845495455749,0.4472052566240151,0.21274301655819566,0.3593763343512352
|
||||
xsum_12_9_yes_teacher_mlm_low,29.737845230742078,0.4465496602405639,0.2123133092772882,0.35962822466043964
|
||||
brewer_xsum_9_9_v8.2,27.76872849201447,0.4443257538168877,0.21127297599326367,0.3582479791521699
|
||||
xsum_theseus_12_6_copy,26.884320127062562,0.4424918692803468,0.21067481121631593,0.3583553330701445
|
||||
brewer_xsum_9_6_gaccum,27.142063001853,0.4425169154001685,0.2085901108525016,0.3575200377038006
|
||||
xsum_dl6,27.928086120180005,0.4428406843456235,0.20797826365160146,0.3570522063056493
|
||||
brewer_xsum_6_6_v3,27.265684284831906,0.43971173638675576,0.20748396024553475,0.3548038868729512
|
||||
xsum_no_teacher_f12_6,27.745874878672904,0.4410138016184357,0.20718399111770808,0.3554781824669484
|
||||
xsum_9_9_encloss10,28.9296744021883,0.4375314579000953,0.2042900182424192,0.3507282669577781
|
||||
xsum_enc_only_9_12_v8,27.90267360804729,0.4317712176227767,0.1975871213673145,0.3445507930002034
|
||||
xsum_no_teacher_f12_3,25.100238242301245,0.4212459339899797,0.19363224077581093,0.3425534541928968
|
||||
xsum_enc_only_6_12_brutasse,28.53339804111885,0.4268136827870216,0.19356708317121352,0.3391381862315658
|
||||
xsum_9_12_noteacher,29.28192005647225,0.4197414451047208,0.1873359194795431,0.33264079827956505
|
||||
eval_xsum_9_9_no_teacher,28.378540545310155,0.4154520215812441,0.1838431267358389,0.3296500161194308
|
||||
xsum_teacher_f12_3,22.863937174622784,0.4099647311777451,0.1838389871897501,0.3350523158599553
|
||||
xsum_9_9_no_teacher,28.378540545310155,0.4155138730166103,0.1837896883494328,0.3296699124599892
|
||||
xsum_6_6_yes_teacher,27.355245742521838,0.4156289447797342,0.1833851442061078,0.3312280244052153
|
||||
brewer_xsum_3_3_v2,24.55272213888644,0.4066228194707484,0.1820601762610753,0.3289002584850459
|
||||
xsum_theseus_6_6_v2,27.172240360010587,0.3843975566079777,0.16290567759393396,0.3068524775579614
|
||||
brewer_xsum_6_6_fix2,26.134651019147622,0.38678850503910334,0.15887838687468492,0.3076899937354824
|
||||
staged_xsum_9_9,27.46360187064325,0.3735212492621201,0.14991821061038266,0.2934425279687253
|
||||
brewer_xsum_6_6_newcode,59.600811788582014,0.24447308274539264,0.06887251729221451,0.17840798937780264
|
||||
|
@@ -69,6 +69,9 @@ class BartConfig(PretrainedConfig):
|
||||
normalize_embedding=True,
|
||||
static_position_embeddings=False,
|
||||
add_bias_logits=False,
|
||||
student_decoder_layers=None,
|
||||
student_encoder_layers=None,
|
||||
replacing_rate=0,
|
||||
**common_kwargs
|
||||
):
|
||||
r"""
|
||||
@@ -87,6 +90,7 @@ class BartConfig(PretrainedConfig):
|
||||
is_encoder_decoder=is_encoder_decoder,
|
||||
**common_kwargs,
|
||||
)
|
||||
self.replacing_rate = replacing_rate
|
||||
self.vocab_size = vocab_size
|
||||
self.d_model = d_model # encoder_embed_dim and decoder_embed_dim
|
||||
self.encoder_ffn_dim = encoder_ffn_dim
|
||||
@@ -119,9 +123,12 @@ class BartConfig(PretrainedConfig):
|
||||
# Classifier stuff
|
||||
self.classif_dropout = classifier_dropout
|
||||
|
||||
# pos embedding offset
|
||||
self.extra_pos_embeddings = self.pad_token_id + 1
|
||||
|
||||
# Theseus params
|
||||
self.student_encoder_layers = student_encoder_layers
|
||||
self.student_decoder_layers = student_decoder_layers
|
||||
|
||||
@property
|
||||
def num_attention_heads(self) -> int:
|
||||
return self.encoder_attention_heads
|
||||
|
||||
@@ -23,6 +23,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor, nn
|
||||
from torch.distributions.bernoulli import Bernoulli
|
||||
from torch.nn import CrossEntropyLoss
|
||||
|
||||
from .activations import ACT2FN
|
||||
@@ -44,6 +45,23 @@ BART_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
||||
]
|
||||
|
||||
|
||||
def get_layers_to_copy(n_to_get, tot):
|
||||
all_layers = list(range(tot))
|
||||
if tot == 12: # Alternating for special cases
|
||||
layers_to_copy = { # maps # layers in student -> which teacher layers to copy
|
||||
6: [0, 2, 4, 7, 9, 11],
|
||||
1: [11],
|
||||
3: [0, 6, 11],
|
||||
2: [0, 11],
|
||||
4: [0, 4, 8, 11],
|
||||
9: [0, 1, 2, 4, 5, 7, 9, 10, 11],
|
||||
12: all_layers,
|
||||
}
|
||||
return layers_to_copy[n_to_get]
|
||||
else:
|
||||
return all_layers[:n_to_get]
|
||||
|
||||
|
||||
BART_START_DOCSTRING = r"""
|
||||
|
||||
This model is a PyTorch `torch.nn.Module <https://pytorch.org/docs/stable/nn.html#torch.nn.Module>`_ sub-class. Use it as a regular PyTorch Module and
|
||||
@@ -62,6 +80,7 @@ BART_GENERATION_EXAMPLE = r"""
|
||||
# see ``examples/summarization/bart/run_eval.py`` for a longer example
|
||||
model = BartForConditionalGeneration.from_pretrained('facebook/bart-large-cnn')
|
||||
tokenizer = BartTokenizer.from_pretrained('facebook/bart-large-cnn')
|
||||
|
||||
ARTICLE_TO_SUMMARIZE = "My friends are cool but they eat too many carbs."
|
||||
inputs = tokenizer.batch_encode_plus([ARTICLE_TO_SUMMARIZE], max_length=1024, return_tensors='pt')
|
||||
# Generate Summary
|
||||
@@ -235,7 +254,37 @@ class EncoderLayer(nn.Module):
|
||||
return x, attn_weights
|
||||
|
||||
|
||||
class BartEncoder(nn.Module):
|
||||
class TheseusMixin:
|
||||
compress_ratio = 2
|
||||
|
||||
def set_replacing_rate(self, replacing_rate):
|
||||
if not 0 < replacing_rate <= 1:
|
||||
raise Exception("Replace rate must be in the range (0, 1]!")
|
||||
self.bernoulli = Bernoulli(torch.tensor([replacing_rate]))
|
||||
|
||||
def determine_inference_layers(self) -> nn.ModuleList:
|
||||
if self.scc_layers is None:
|
||||
return self.layers
|
||||
if self.training:
|
||||
inference_layers = []
|
||||
for i in range(len(self.scc_layers)):
|
||||
if self.bernoulli.sample() == 1: # REPLACE
|
||||
inference_layers.append(self.scc_layers[i])
|
||||
else: # KEEP the original
|
||||
for offset in range(self.compress_ratio):
|
||||
inference_layers.append(self.layers[i * self.compress_ratio + offset])
|
||||
|
||||
else: # inference with compressed model
|
||||
inference_layers = self.scc_layers
|
||||
|
||||
return inference_layers
|
||||
|
||||
def init_successor_layers(self, replacing_rate):
|
||||
if replacing_rate > 0 and replacing_rate < 1:
|
||||
self.set_replacing_rate(replacing_rate)
|
||||
|
||||
|
||||
class BartEncoder(nn.Module, TheseusMixin):
|
||||
"""
|
||||
Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer
|
||||
is a :class:`EncoderLayer`.
|
||||
@@ -265,9 +314,15 @@ class BartEncoder(nn.Module):
|
||||
config.max_position_embeddings, embed_dim, self.padding_idx, config.extra_pos_embeddings,
|
||||
)
|
||||
self.layers = nn.ModuleList([EncoderLayer(config) for _ in range(config.encoder_layers)])
|
||||
|
||||
self.layernorm_embedding = LayerNorm(embed_dim) if config.normalize_embedding else nn.Identity()
|
||||
# mbart has one extra layer_norm
|
||||
self.layer_norm = LayerNorm(config.d_model) if config.normalize_before else None
|
||||
self.scc_layers = None
|
||||
if config.student_encoder_layers is not None and config.student_encoder_layers < config.encoder_layers:
|
||||
self.scc_layers = nn.ModuleList([EncoderLayer(config) for _ in range(config.student_encoder_layers)])
|
||||
self.compress_ratio = len(self.scc_layers) // len(self.layers)
|
||||
self.init_successor_layers(config.replacing_rate)
|
||||
|
||||
def forward(self, input_ids, attention_mask=None, output_attentions=False, output_hidden_states=False):
|
||||
"""
|
||||
@@ -294,12 +349,13 @@ class BartEncoder(nn.Module):
|
||||
x = inputs_embeds + embed_pos
|
||||
x = self.layernorm_embedding(x)
|
||||
x = F.dropout(x, p=self.dropout, training=self.training)
|
||||
inference_layers = self.determine_inference_layers()
|
||||
|
||||
# B x T x C -> T x B x C
|
||||
x = x.transpose(0, 1)
|
||||
|
||||
encoder_states, all_attentions = [], []
|
||||
for encoder_layer in self.layers:
|
||||
for encoder_layer in inference_layers:
|
||||
if output_hidden_states:
|
||||
encoder_states.append(x)
|
||||
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
|
||||
@@ -413,7 +469,7 @@ class DecoderLayer(nn.Module):
|
||||
) # just self_attn weights for now, following t5, layer_state = cache for decoding
|
||||
|
||||
|
||||
class BartDecoder(nn.Module):
|
||||
class BartDecoder(nn.Module, TheseusMixin):
|
||||
"""
|
||||
Transformer decoder consisting of *config.decoder_layers* layers. Each layer
|
||||
is a :class:`DecoderLayer`.
|
||||
@@ -443,6 +499,11 @@ class BartDecoder(nn.Module):
|
||||
) # type: List[DecoderLayer]
|
||||
self.layernorm_embedding = LayerNorm(config.d_model) if config.normalize_embedding else nn.Identity()
|
||||
self.layer_norm = LayerNorm(config.d_model) if config.add_final_layer_norm else None
|
||||
self.scc_layers = None
|
||||
if config.student_decoder_layers is not None and config.student_decoder_layers < config.decoder_layers:
|
||||
self.scc_layers = nn.ModuleList([DecoderLayer(config) for _ in range(config.student_decoder_layers)])
|
||||
self.compress_ratio = len(self.scc_layers) // len(self.layers)
|
||||
self.init_successor_layers(config.replacing_rate)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -500,7 +561,9 @@ class BartDecoder(nn.Module):
|
||||
all_hidden_states = ()
|
||||
all_self_attns = ()
|
||||
next_decoder_cache = []
|
||||
for idx, decoder_layer in enumerate(self.layers):
|
||||
|
||||
inference_layers = self.determine_inference_layers()
|
||||
for idx, decoder_layer in enumerate(inference_layers):
|
||||
# add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (x,)
|
||||
@@ -509,7 +572,7 @@ class BartDecoder(nn.Module):
|
||||
continue
|
||||
|
||||
layer_state = decoder_cached_states[idx] if decoder_cached_states is not None else None
|
||||
|
||||
# DecoderLayer.forward()
|
||||
x, layer_self_attn, layer_past = decoder_layer(
|
||||
x,
|
||||
encoder_hidden_states,
|
||||
@@ -797,7 +860,6 @@ def _get_shape(t):
|
||||
class BartModel(PretrainedBartModel):
|
||||
def __init__(self, config: BartConfig):
|
||||
super().__init__(config)
|
||||
|
||||
padding_idx, vocab_size = config.pad_token_id, config.vocab_size
|
||||
self.shared = nn.Embedding(vocab_size, config.d_model, padding_idx)
|
||||
|
||||
@@ -976,8 +1038,8 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
decoder_cached_states=decoder_cached_states,
|
||||
use_cache=use_cache,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
output_attentions=output_attentions,
|
||||
)
|
||||
lm_logits = F.linear(outputs[0], self.model.shared.weight, bias=self.final_logits_bias)
|
||||
outputs = (lm_logits,) + outputs[1:] # Add cache, hidden states and attention if they are here
|
||||
@@ -991,7 +1053,6 @@ class BartForConditionalGeneration(PretrainedBartModel):
|
||||
|
||||
def prepare_inputs_for_generation(self, decoder_input_ids, past, attention_mask, use_cache, **kwargs):
|
||||
assert past is not None, "past has to be defined for encoder_outputs"
|
||||
|
||||
encoder_outputs, decoder_cached_states = past
|
||||
return {
|
||||
"input_ids": None, # encoder_outputs is defined. input_ids not needed
|
||||
@@ -1112,8 +1173,8 @@ class BartForSequenceClassification(PretrainedBartModel):
|
||||
decoder_input_ids=decoder_input_ids,
|
||||
decoder_attention_mask=decoder_attention_mask,
|
||||
encoder_outputs=encoder_outputs,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
output_attentions=output_attentions,
|
||||
)
|
||||
x = outputs[0] # last hidden state
|
||||
eos_mask = input_ids.eq(self.config.eos_token_id)
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
,length_penalty,max_length,min_length,num_beams,rouge1,rouge2,rougeL
|
||||
beam_search_xsum_base_500/generations_0.txt,0.5,62.0,11.0,2.0,0.4405564669595006,0.1983002888180257,0.3437812408347012
|
||||
beam_search_xsum_base_500/generations_1.txt,0.5,62.0,11.0,4.0,0.4448474889733566,0.20778679115689908,0.35315548966820265
|
||||
beam_search_xsum_base_500/generations_2.txt,0.5,62.0,11.0,6.0,0.44523636974022285,0.21033071404850712,0.3580167024803249
|
||||
beam_search_xsum_base_500/generations_3.txt,0.5,62.0,11.0,8.0,0.4468688924754079,0.21170992719882548,0.35932945274070693
|
||||
beam_search_xsum_base_500/generations_4.txt,0.5,62.0,11.0,10.0,0.44549478467300774,0.21288210214860998,0.36011253722710324
|
||||
beam_search_xsum_base_500/generations_5.txt,0.5,62.0,11.0,20.0,0.4459196355164726,0.21553443653166177,0.36066326922139236
|
||||
beam_search_xsum_base_500/generations_6.txt,1.0,62.0,11.0,2.0,0.44053092513361236,0.1993972468658327,0.3444322842619076
|
||||
beam_search_xsum_base_500/generations_7.txt,1.0,62.0,11.0,4.0,0.4417046134948439,0.20676749488131385,0.3499947350751297
|
||||
beam_search_xsum_base_500/generations_8.txt,1.0,62.0,11.0,6.0,0.44452835581665295,0.20896315664629483,0.3549064385172764
|
||||
beam_search_xsum_base_500/generations_9.txt,1.0,62.0,11.0,8.0,0.44358288615105523,0.2093555938220097,0.3549769613045229
|
||||
beam_search_xsum_base_500/generations_10.txt,1.0,62.0,11.0,10.0,0.44330114244232066,0.2119742308397891,0.3570212060039547
|
||||
beam_search_xsum_base_500/generations_11.txt,1.0,62.0,11.0,20.0,0.44368891403518096,0.21275693156665768,0.35774144153552306
|
||||
beam_search_xsum_base_500/generations_12.txt,2.0,62.0,11.0,2.0,0.4351172939026171,0.1946192131538515,0.338306269234114
|
||||
beam_search_xsum_base_500/generations_13.txt,2.0,62.0,11.0,4.0,0.4355928169879645,0.20028529354739374,0.3409900672696907
|
||||
beam_search_xsum_base_500/generations_14.txt,2.0,62.0,11.0,6.0,0.43613272840789125,0.20204391577628494,0.34330389423389684
|
||||
beam_search_xsum_base_500/generations_15.txt,2.0,62.0,11.0,8.0,0.43189181564990875,0.1985819991945317,0.33890611664212345
|
||||
beam_search_xsum_base_500/generations_16.txt,2.0,62.0,11.0,10.0,0.4312709248247222,0.19786145957678045,0.340228104937922
|
||||
beam_search_xsum_base_500/generations_17.txt,2.0,62.0,11.0,20.0,0.431958235502211,0.1997066827574804,0.3410722313510985
|
||||
|
@@ -0,0 +1,19 @@
|
||||
,length_penalty,max_length,min_length,num_beams,time,rouge1,rouge2,rougeL
|
||||
beam_search_xsum_d6_500/generations_5.txt,1.0,62.0,11.0,20.0,236.1429727077484,0.4305654553611711,0.19600039925703722,0.34201476194014246
|
||||
beam_search_xsum_d6_500/generations_4.txt,1.0,62.0,11.0,10.0,135.6617534160614,0.42671786797903843,0.1929509089620034,0.3357873765545949
|
||||
beam_search_xsum_d6_500/generations_3.txt,1.0,62.0,11.0,8.0,118.18162441253662,0.4244918905164582,0.19240953848063766,0.33592718733437477
|
||||
beam_search_xsum_d6_500/generations_2.txt,1.0,62.0,11.0,6.0,99.75846719741821,0.4241986674718551,0.1904178221737402,0.33554345459177803
|
||||
beam_search_xsum_d6_500/generations_1.txt,1.0,62.0,11.0,4.0,82.19014501571655,0.42552162916374947,0.1900029083958625,0.3348035451647859
|
||||
beam_search_xsum_d6_500/generations_0.txt,1.0,62.0,11.0,2.0,73.42616081237793,0.42185608960262033,0.188239694825568,0.3380120701696464
|
||||
beam_search_xsum_d6_500/generations_10.txt,2.0,62.0,11.0,10.0,134.8128764629364,0.41995279044460687,0.1873177328806112,0.32606274473224905
|
||||
beam_search_xsum_d6_500/generations_16.txt,3.0,62.0,11.0,10.0,135.02988576889038,0.41958442037714705,0.1862058575783302,0.3248733090856767
|
||||
beam_search_xsum_d6_500/generations_11.txt,2.0,62.0,11.0,20.0,235.85728216171265,0.4218916450932465,0.18532554772332022,0.3277383511244094
|
||||
beam_search_xsum_d6_500/generations_8.txt,2.0,62.0,11.0,6.0,97.2721335887909,0.41963315273987645,0.18522376636363672,0.3271210646772599
|
||||
beam_search_xsum_d6_500/generations_9.txt,2.0,62.0,11.0,8.0,115.15277361869812,0.417398394462531,0.1845339882603932,0.3254173997315589
|
||||
beam_search_xsum_d6_500/generations_15.txt,3.0,62.0,11.0,8.0,115.81080627441406,0.41703499851068504,0.18357496523072026,0.3242598639398694
|
||||
beam_search_xsum_d6_500/generations_6.txt,2.0,62.0,11.0,2.0,68.16163468360901,0.41780417025178973,0.18351703902053906,0.3313922817377364
|
||||
beam_search_xsum_d6_500/generations_14.txt,3.0,62.0,11.0,6.0,97.10926246643066,0.41681309969268543,0.18350802133208385,0.3248941386957547
|
||||
beam_search_xsum_d6_500/generations_7.txt,2.0,62.0,11.0,4.0,79.3750171661377,0.4192827070533588,0.18286781380405798,0.32295416688769185
|
||||
beam_search_xsum_d6_500/generations_12.txt,3.0,62.0,11.0,2.0,67.85268831253052,0.4173974464414504,0.18261800246898374,0.33032670326211594
|
||||
beam_search_xsum_d6_500/generations_13.txt,3.0,62.0,11.0,4.0,78.7077898979187,0.41858355068804654,0.18122511913031086,0.32144676793629146
|
||||
beam_search_xsum_d6_500/generations_17.txt,3.0,62.0,11.0,20.0,236.6838607788086,0.25952276985447975,0.100359902189957,0.19233054575018543
|
||||
|
Reference in New Issue
Block a user