Merge branch 'distilbart-clean' into theseus

This commit is contained in:
sshleifer
2020-06-14 19:51:25 -04:00
13 changed files with 260 additions and 213 deletions
+12 -13
View File
@@ -14,6 +14,7 @@ from transformers import (
AutoModel,
AutoModelForPreTraining,
AutoModelForQuestionAnswering,
AutoModelForSeq2SeqLM,
AutoModelForSequenceClassification,
AutoModelForTokenClassification,
AutoModelWithLMHead,
@@ -34,6 +35,8 @@ MODEL_MODES = {
"pretraining": AutoModelForPreTraining,
"token-classification": AutoModelForTokenClassification,
"language-modeling": AutoModelWithLMHead,
"summarization": AutoModelForSeq2SeqLM,
"translation": AutoModelForSeq2SeqLM,
}
@@ -45,12 +48,6 @@ def set_seed(args: argparse.Namespace):
torch.cuda.manual_seed_all(args.seed)
def count_trainable_parameters(model):
model_parameters = filter(lambda p: p.requires_grad, model.parameters())
params = sum([np.prod(p.size()) for p in model_parameters])
return params
class BaseTransformer(pl.LightningModule):
def __init__(
self,
@@ -94,7 +91,7 @@ class BaseTransformer(pl.LightningModule):
cache_dir=cache_dir,
)
else:
self.model_type = type(model)
self.model_type = None
self.model = model
def load_hf_checkpoint(self, *args, **kwargs):
@@ -186,7 +183,7 @@ class BaseTransformer(pl.LightningModule):
)
parser.add_argument(
"--tokenizer_name",
default="facebook/bart-large",
default=None,
type=str,
help="Pretrained tokenizer name or path if not the same as model_name",
)
@@ -233,8 +230,8 @@ class LoggingCallback(pl.Callback):
writer.write("{} = {}\n".format(key, str(metrics[key])))
def add_generic_args(parser, root_dir):
parser = pl.Trainer.add_argparse_args(parser)
def add_generic_args(parser, root_dir) -> None:
# TODO(SS): allow all pl args? parser = pl.Trainer.add_argparse_args(parser)
parser.add_argument(
"--output_dir",
default=None,
@@ -271,13 +268,14 @@ def add_generic_args(parser, root_dir):
parser.add_argument("--seed", type=int, default=42, help="random seed for initialization")
parser.add_argument("--resume_from_checkpoint", type=str, default=None)
parser.add_argument("--val_check_interval", default=1.0, type=float)
def generic_train(
model: BaseTransformer,
args: argparse.Namespace,
early_stopping_callback=False,
logger=True,
logger=True, # can pass WandbLogger() here
extra_callbacks=[],
checkpoint_callback=None,
logging_callback=None,
@@ -302,6 +300,8 @@ def generic_train(
if args.n_tpu_cores > 0:
global xm
import torch_xla.core.xla_model as xm
train_params["num_tpu_cores"] = args.n_tpu_cores
train_params["gpus"] = 0
@@ -309,7 +309,7 @@ def generic_train(
train_params["distributed_backend"] = "ddp"
trainer = pl.Trainer(
logger=True,
logger=logger,
accumulate_grad_batches=args.gradient_accumulation_steps,
gpus=args.gpus,
max_epochs=args.num_train_epochs,
@@ -321,7 +321,6 @@ def generic_train(
val_check_interval=args.val_check_interval,
weights_summary=None,
resume_from_checkpoint=args.resume_from_checkpoint,
auto_scale_batch_size=args.auto_scale_batch_size,
**train_params,
)
+1
View File
@@ -7,3 +7,4 @@ rouge-score
tensorflow_datasets
pytorch-lightning==0.7.6 # April 10, 2020 release
matplotlib
git-python==1.0.3
+43 -26
View File
@@ -1,47 +1,64 @@
### Get CNN Data
To be able to reproduce the authors' results on the CNN/Daily Mail dataset you first need to download both CNN and Daily Mail datasets [from Kyunghyun Cho's website](https://cs.nyu.edu/~kcho/DMQA/) (the links next to "Stories") in the same folder. Then uncompress the archives by running:
### Data
CNN/DailyMail data
```bash
cd examples/summarization
wget https://s3.amazonaws.com/datasets.huggingface.co/summarization/cnn_dm.tgz
tar -xzvf cnn_dm.tgz
export CNN_DIR=${PWD}/cnn_dm
```
this should make a directory called cnn_dm/ with files like `test.source`.
To use your own data, copy that files format. Each article to be summarized is on its own line.
XSUM Data:
```bash
cd examples/summarization
wget https://s3.amazonaws.com/datasets.huggingface.co/summarization/xsum.tar.gz
tar -xzvf xsum.tar.gz
export XSUM_DIR=${PWD}/xsum
```
### Evaluation
To create summaries for each article in dataset, run:
```bash
python evaluate_cnn.py <path_to_test.source> test_generations.txt <model-name> --score_path rouge_scores.txt
python run_eval.py <path_to_test.source> test_generations.txt <model-name> --score_path rouge_scores.txt
```
The default batch size, 8, fits in 16GB GPU memory, but may need to be adjusted to fit your system.
The default batch size, 4, fits in 16GB GPU memory, but may need to be adjusted to fit your system.
### Training
Run/modify `finetune_bart.sh` or `finetune_t5.sh`
Run/modify `finetune.sh`
### Stanford CoreNLP Setup
The following command should work on a 16GB GPU:
```bash
export me=`git config user.name`
./finetune.sh \
--data_dir $XSUM_DIR \
--train_batch_size=1 \
--eval_batch_size=1 \
--output_dir="$me"_xsum_results \
--num_train_epochs 1
```
ptb_tokenize () {
cat $1 | java edu.stanford.nlp.process.PTBTokenizer -ioFileList -preserveLines > $2
}
sudo apt install openjdk-8-jre-headless
sudo apt-get install ant
wget http://nlp.stanford.edu/software/stanford-corenlp-full-2018-10-05.zip
unzip stanford-corenlp-full-2018-10-05.zip
cd stanford-corenlp-full-2018-10-05
export CLASSPATH=stanford-corenlp-3.9.2.jar:stanford-corenlp-3.9.2-models.jar
```
Then run `ptb_tokenize` on `test.target` and your generated hypotheses.
### Rouge Setup
Install `files2rouge` following the instructions at [here](https://github.com/pltrdy/files2rouge).
I also needed to run `sudo apt-get install libxml-parser-perl`
Tips:
- 1 epoch at batch size 1 for bart-large takes 24 hours, requires 13GB GPU RAM with fp16.
- try --freeze_encoder or --freeze_embeds for faster training/larger batch size. (3hr/epoch with bs=8, see below)
- fp16 opt level O1 is best
```python
from files2rouge import files2rouge
from files2rouge import settings
files2rouge.run(<path_to_tokenized_hypo>,
<path_to_tokenized_target>,
saveto='rouge_output.txt')
### Shared Task
Compare XSUM results with others by using `--logger wandb_shared`. This requires `wandb` registration.
Here is an example command
```bash
export me=`git config user.name`
./finetune.sh \
--data_dir $XSUM_DIR \
--output_dir "$me"_xsum_frozen_embs \
--logger wandb_shared \
--train_batch_size 16 --eval_batch_size 16 --freeze_embeds --freeze_encoder \
--num_train_epochs 6
```
Results can be viewed [here](https://app.wandb.ai/sshleifer/hf_summarization/table?workspace=user-)
+6 -1
View File
@@ -2,12 +2,17 @@ import logging
import os
from pathlib import Path
import numpy as np
import pytorch_lightning as pl
import torch
from pytorch_lightning.callbacks import ModelCheckpoint
from pytorch_lightning.utilities import rank_zero_only
from lightning_base import count_trainable_parameters
def count_trainable_parameters(model):
model_parameters = filter(lambda p: p.requires_grad, model.parameters())
params = sum([np.prod(p.size()) for p in model_parameters])
return params
logger = logging.getLogger(__name__)
+34 -31
View File
@@ -11,29 +11,32 @@ from torch.nn import functional as F
from lightning_base import generic_train
from transformers import AdamW, BartConfig, BartForConditionalGeneration, T5Config, T5ForConditionalGeneration
from transformers.modeling_bart import invert_mask
try:
from .finetune import (
SummarizationTrainer,
freeze_params,
assert_all_frozen,
any_requires_grad,
)
from .finetune import SummarizationTrainer
from .initialization_utils import init_student, copy_layers
from .utils import SummarizationDataset, pickle_load
from .finetune import main as ft_main
except ModuleNotFoundError:
from finetune import (
SummarizationTrainer,
from .utils import (
use_task_specific_params,
SummarizationDataset,
pickle_load,
freeze_params,
assert_all_frozen,
any_requires_grad,
)
from .finetune import main as ft_main
except ImportError:
from finetune import SummarizationTrainer
from finetune import main as ft_main
from initialization_utils import init_student, copy_layers
from utils import SummarizationDataset, pickle_load
from utils import (
use_task_specific_params,
SummarizationDataset,
pickle_load,
freeze_params,
assert_all_frozen,
any_requires_grad,
)
class TheseusDistiller(SummarizationTrainer):
@@ -49,8 +52,10 @@ class SummarizationDistiller(SummarizationTrainer):
super().__init__(hparams, model=student, config=student_cfg)
self.teacher = teacher
use_task_specific_params(self.teacher, "summarization")
freeze_params(self.teacher)
self.freeze_stuff(d_layers_to_copy)
assert len(self.model.model.decoder.layers) == len(d_layers_to_copy)
self.sanity_check_gradients()
self.ce_loss_fct = nn.KLDivLoss(reduction="batchmean")
self.temperature = 2.0
self.alpha_mlm = hparams.alpha_mlm
@@ -61,8 +66,7 @@ class SummarizationDistiller(SummarizationTrainer):
gc.collect()
torch.cuda.empty_cache()
def freeze_stuff(self, d_layers_to_copy):
assert len(self.model.model.decoder.layers) == len(d_layers_to_copy)
def sanity_check_gradients(self):
assert_all_frozen(self.teacher)
assert_all_frozen(self.model.model.decoder.embed_tokens)
assert_all_frozen(self.model.model.encoder.embed_tokens)
@@ -81,7 +85,7 @@ class SummarizationDistiller(SummarizationTrainer):
}
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)
hparams.layer_to_copy = d_layers_to_copy
hparams.d_layer_to_copy = d_layers_to_copy
hparams.e_layer_to_copy = e_layers_to_copy
kw = teacher.config.to_diff_dict()
kw.update(student_updates)
@@ -252,7 +256,7 @@ class SummarizationDistiller(SummarizationTrainer):
dec_mask = decoder_input_ids.eq(self.tokenizer.pad_token_id)
loss_ce, s_logits_slct, t_logits_slct = self.calc_ce_loss(dec_mask, slogits, tlogits)
if not self.hparams.freeze_decoder and self.alpha_hid > 0:
hid_loss_dec = self.calc_hidden_loss(dec_mask, dec_hidden, tdec_hidden, self.hparams.layer_to_copy)
hid_loss_dec = self.calc_hidden_loss(dec_mask, dec_hidden, tdec_hidden, self.hparams.d_layer_to_copy)
blended_loss = (
self.alpha_ce * loss_ce
@@ -287,7 +291,7 @@ class T5SummarizationDistiller(SummarizationDistiller):
d_layers_to_copy = get_layers_to_copy(n_layer, len(teacher.decoder.block))
e_layers_to_copy: List = get_layers_to_copy(n_layer, len(teacher.encoder.block))
student_updates = {"num_layers": n_layer}
hparams.layer_to_copy = d_layers_to_copy
hparams.d_layer_to_copy = d_layers_to_copy
hparams.e_layer_to_copy = e_layers_to_copy
kw = teacher.config.to_diff_dict()
@@ -308,7 +312,7 @@ class T5SummarizationDistiller(SummarizationDistiller):
for d in [self.model.encoder, self.model.decoder]:
freeze_params(d.embed_tokens)
def freeze_stuff(self, d_layers_to_copy):
def sanity_check_gradients(self, d_layers_to_copy):
"""T5"""
assert len(self.model.decoder.block) == len(d_layers_to_copy)
assert_all_frozen(self.teacher)
@@ -325,7 +329,6 @@ class T5SummarizationDistiller(SummarizationDistiller):
freeze_params(self.model.decoder) # TODO(SS): very suspicious
def _step(self, batch):
# assert is_frozen(self.teacher)
pad_token_id = self.tokenizer.pad_token_id
source_ids, source_mask, y = batch["input_ids"], batch["attention_mask"], batch["decoder_input_ids"]
decoder_input_ids = y[:, :-1].contiguous()
@@ -376,7 +379,7 @@ class T5SummarizationDistiller(SummarizationDistiller):
loss_ce, s_logits_slct, t_logits_slct = self.calc_ce_loss(dec_mask, slogits, tlogits)
if not self.hparams.freeze_decoder and self.alpha_hid > 0:
hid_loss_dec = self.calc_hidden_loss(dec_mask, dec_hidden, tdec_hidden, self.hparams.layer_to_copy)
hid_loss_dec = self.calc_hidden_loss(dec_mask, dec_hidden, tdec_hidden, self.hparams.d_layer_to_copy)
blended_loss = (
self.alpha_ce * loss_ce
@@ -387,15 +390,6 @@ class T5SummarizationDistiller(SummarizationDistiller):
return blended_loss, loss_ce, sloss, loss_encoder, hid_loss_enc, hid_loss_dec
def distill_main(args):
Path(args.output_dir).mkdir(exist_ok=True)
if len(os.listdir(args.output_dir)) > 3 and args.do_train:
raise ValueError("Output directory ({}) already exists and is not empty.".format(args.output_dir))
model = create_module(args)
ft_main(args, model=model)
def create_module(args):
t5 = "t5" in args.model_name_or_path
if args.no_teacher:
@@ -452,6 +446,15 @@ def get_layers_to_copy(n_to_get, tot):
return all_layers[:n_to_get]
def distill_main(args):
Path(args.output_dir).mkdir(exist_ok=True)
if len(os.listdir(args.output_dir)) > 3 and args.do_train:
raise ValueError("Output directory ({}) already exists and is not empty.".format(args.output_dir))
model = create_module(args)
ft_main(args, model=model)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser = SummarizationDistiller.add_model_specific_args(parser, os.getcwd())
+41 -92
View File
@@ -1,118 +1,61 @@
import argparse
import glob
import json
import logging
import os
import time
from pathlib import Path
from typing import Dict, Iterable, List, Tuple
from typing import Dict, List, Tuple
import git
import numpy as np
import pytorch_lightning as pl
import torch
from rouge_score import rouge_scorer, scoring
from torch import nn
from pytorch_lightning.loggers import WandbLogger
from torch.utils.data import DataLoader
from durbango import pickle_save
from lightning_base import BaseTransformer, add_generic_args, generic_train
from transformers import AutoModelWithLMHead, get_linear_schedule_with_warmup
from transformers import get_linear_schedule_with_warmup
try:
from .utils import SummarizationDataset, lmap, flatten_list
from .utils import (
use_task_specific_params,
SummarizationDataset,
lmap,
flatten_list,
pickle_save,
save_git_info,
freeze_params,
calculate_rouge,
get_git_info,
ROUGE_KEYS,
)
from .callbacks import Seq2SeqLoggingCallback, get_rouge2_checkpoint_callback
except ImportError:
from utils import SummarizationDataset, lmap, flatten_list
from .utils import (
use_task_specific_params,
SummarizationDataset,
lmap,
flatten_list,
pickle_save,
save_git_info,
freeze_params,
calculate_rouge,
get_git_info,
ROUGE_KEYS,
)
from callbacks import Seq2SeqLoggingCallback, get_rouge2_checkpoint_callback
logger = logging.getLogger(__name__)
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL"]
def save_git_info(folder_path: str):
"""
Log commit info.
"""
repo_infos = get_git_info()
with open(os.path.join(folder_path, "git_log.json"), "w") as f:
json.dump(repo_infos, f, indent=4)
def get_git_info():
repo = git.Repo(search_parent_directories=True)
repo_infos = {
"repo_id": str(repo),
"repo_sha": str(repo.head.object.hexsha),
"repo_branch": str(repo.active_branch),
}
return repo_infos
def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
scorer = rouge_scorer.RougeScorer(ROUGE_KEYS, use_stemmer=True)
aggregator = scoring.BootstrapAggregator()
for reference_ln, output_ln in zip(reference_lns, output_lns):
scores = scorer.score(reference_ln, output_ln)
aggregator.add_scores(scores)
result = aggregator.aggregate()
return {k: v.mid.fmeasure for k, v in result.items()}
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 freeze_params(model: nn.Module):
for par in model.parameters():
par.requires_grad = False
def grad_status(model: nn.Module) -> Iterable:
return (par.requires_grad for par in model.parameters())
def any_requires_grad(model: nn.Module) -> bool:
return any(grad_status(model))
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"
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"
class SummarizationTrainer(BaseTransformer):
mode = "language-modeling"
mode = "summarization"
loss_names = ["loss"]
def __init__(self, hparams, **kwargs):
super().__init__(hparams, num_labels=None, mode=self.mode, **kwargs)
use_task_specific_params(self.model, "summarization")
save_git_info(self.hparams.output_dir)
self.model: AutoModelWithLMHead
self.metrics_save_path = Path(self.output_dir) / "metrics.pkl"
self.hparams_save_path = Path(self.output_dir) / "hparams.pkl"
self.step_count = 0
@@ -294,9 +237,8 @@ class SummarizationTrainer(BaseTransformer):
@staticmethod
def add_model_specific_args(parser, root_dir):
add_generic_args(parser, root_dir)
BaseTransformer.add_model_specific_args(parser, root_dir)
add_generic_args(parser, root_dir)
parser.add_argument(
"--max_source_length",
default=1024,
@@ -325,7 +267,6 @@ class SummarizationTrainer(BaseTransformer):
help="The maximum total input sequence length after tokenization. Sequences longer "
"than this will be truncated, sequences shorter will be padded.",
)
parser.add_argument(
"--data_dir",
default=None,
@@ -339,10 +280,12 @@ class SummarizationTrainer(BaseTransformer):
parser.add_argument(
"--freeze_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("--n_train", type=int, default=-1, required=False)
parser.add_argument("--n_val", type=int, default=500, required=False)
parser.add_argument("--n_test", type=int, default=-1, required=False)
parser.add_argument("--sortish_sampler", action="store_true", default=False)
return parser
@@ -352,12 +295,19 @@ def main(args, model=None):
raise ValueError("Output directory ({}) already exists and is not empty.".format(args.output_dir))
if model is None:
model: BaseTransformer = SummarizationTrainer(args)
if args.logger == "default" or args.fast_dev_run:
logger = True
elif args.logger == "wandb":
logger = WandbLogger()
elif args.logger == "wandb_shared":
logger = WandbLogger(name=args.output_dir, project="hf_summarization")
trainer: pl.Trainer = generic_train(
model,
args,
early_stopping_callback=True,
logging_callback=Seq2SeqLoggingCallback(),
checkpoint_callback=get_rouge2_checkpoint_callback(args.output_dir),
logger=logger,
)
if not args.do_predict:
return model
@@ -368,13 +318,12 @@ def main(args, model=None):
model.hparams.test_checkpoint = checkpoints[-1]
trainer.resume_from_checkpoint = checkpoints[-1]
trainer.logger.log_hyperparams(model.hparams)
trainer.test(model)
trainer.test(model) # NOTE(SS): this will break in DDP, known lightning issue. See evaluate_checkpoint
return model
if __name__ == "__main__":
parser = argparse.ArgumentParser()
add_generic_args(parser, os.getcwd())
parser = SummarizationTrainer.add_model_specific_args(parser, os.getcwd())
args = parser.parse_args()
+23
View File
@@ -0,0 +1,23 @@
export OUTPUT_DIR=bart_cnn_finetune
# Make output directory if it doesn't exist
mkdir -p $OUTPUT_DIR
# Add parent directory to python path to access lightning_base.py
export PYTHONPATH="../":"${PYTHONPATH}"
# --model_name_or_path=t5-base for t5
python finetune.py \
--model_name_or_path=facebook/bart-large \
--learning_rate=3e-5 \
--fp16 \
--gpus 1 \
--do_train \
--do_predict \
--n_val 1000 \
--val_check_interval 0.1 \
--sortish_sampler \
--max_target_length=56 \
$@
-18
View File
@@ -1,18 +0,0 @@
export OUTPUT_DIR_NAME=bart_sum
export CURRENT_DIR=${PWD}
export OUTPUT_DIR=${CURRENT_DIR}/${OUTPUT_DIR_NAME}
# Make output directory if it doesn't exist
mkdir -p $OUTPUT_DIR
# Add parent directory to python path to access lightning_base.py
export PYTHONPATH="../":"${PYTHONPATH}"
python finetune.py \
--data_dir=./cnn-dailymail/cnn_dm \
--model_name_or_path=bart-large \
--learning_rate=3e-5 \
--train_batch_size=4 \
--eval_batch_size=4 \
--output_dir=$OUTPUT_DIR \
--do_train $@
+1 -1
View File
@@ -3,7 +3,7 @@
# Add parent directory to python path to access lightning_base.py
export PYTHONPATH="../":"${PYTHONPATH}"
python finetune.py \
python distillation.py \
--learning_rate=3e-4 \
--do_train \
--do_predict \
@@ -1,13 +1,18 @@
import argparse
import json
from pathlib import Path
import torch
from rouge_score import rouge_scorer, scoring
from tqdm import tqdm
from transformers import AutoModelWithLMHead, AutoTokenizer
try:
from .finetune import calculate_rouge, use_task_specific_params
except ImportError:
from finetune import calculate_rouge, use_task_specific_params
DEFAULT_DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
@@ -26,9 +31,7 @@ def generate_summaries(
tokenizer = AutoTokenizer.from_pretrained(model_name)
# update config with summarization specific params
task_specific_params = model.config.task_specific_params
if task_specific_params is not None:
model.config.update(task_specific_params.get("summarization", {}))
use_task_specific_params(model, "summarization")
for batch in tqdm(list(chunks(examples, batch_size))):
if "t5" in model_name:
@@ -44,27 +47,10 @@ def generate_summaries(
fout.flush()
def calculate_rouge(output_lns, reference_lns, score_path):
score_file = Path(score_path).open("w")
scorer = rouge_scorer.RougeScorer(["rouge1", "rouge2", "rougeL"], use_stemmer=True)
aggregator = scoring.BootstrapAggregator()
for reference_ln, output_ln in zip(reference_lns, output_lns):
scores = scorer.score(reference_ln, output_ln)
aggregator.add_scores(scores)
result = aggregator.aggregate()
score_file.write(
"ROUGE_1: \n{} \n\n ROUGE_2: \n{} \n\n ROUGE_L: \n{} \n\n".format(
result["rouge1"], result["rouge2"], result["rougeL"]
)
)
def run_generate():
parser = argparse.ArgumentParser()
parser.add_argument(
"input_path", type=str, help="like cnn_dm/test.source or cnn_dm/test_articles_input.txt",
"input_path", type=str, help="like cnn_dm/test.source",
)
parser.add_argument(
"output_path", type=str, help="where to save summaries",
@@ -73,11 +59,11 @@ def run_generate():
"model_name",
type=str,
default="facebook/bart-large-cnn",
help="like bart-large-cnn,'t5-small', 't5-base', 't5-large', 't5-3b', 't5-11b",
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",
"--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.",
@@ -93,7 +79,9 @@ def run_generate():
output_lns = [x.rstrip() for x in open(args.output_path).readlines()]
reference_lns = [x.rstrip() for x in open(args.reference_path).readlines()]
calculate_rouge(output_lns, reference_lns, args.score_path)
rouge: dict = calculate_rouge(output_lns, reference_lns)
json.dump(rouge, open("score_path", "w+"))
if __name__ == "__main__":
@@ -7,15 +7,14 @@ import unittest
from pathlib import Path
from unittest.mock import patch
import pandas as pd
import torch
from torch.utils.data import DataLoader
from transformers import BartTokenizer
from .distillation import distill_main, evaluate_checkpoint
from .evaluate_cnn import generate_summaries, run_generate
from .finetune import main
from .run_eval import generate_summaries, run_generate
from .utils import SummarizationDataset, lmap, pickle_load
@@ -24,6 +23,7 @@ logging.basicConfig(level=logging.DEBUG)
logger = logging.getLogger()
FP16_EVER = False
CHEAP_ARGS = {
"logger": "default",
"alpha_hid": 0,
"freeze_embeds": True,
"enc_only": False,
@@ -244,6 +244,8 @@ class TestSummarizationDistiller(unittest.TestCase):
# self.assertEqual(len(contents), 15)
metrics = pickle_load(Path(output_dir) / "metrics.pkl")
import pandas as pd
val_df = pd.DataFrame(metrics["val"])
train_df = pd.DataFrame(metrics["train"])
test_df = pd.DataFrame(metrics["test"])
@@ -267,7 +269,7 @@ class TestBartExamples(unittest.TestCase):
output_file_name = Path(tempfile.gettempdir()) / "utest_output_bart_sum.hypo"
articles = [" New York (CNN)When Liana Barrientos was 23 years old, she got married in Westchester County."]
_dump_articles(tmp, articles)
testargs = ["evaluate_cnn.py", str(tmp), str(output_file_name), "sshleifer/bart-tiny-random"]
testargs = ["run_eval.py", str(tmp), str(output_file_name), "sshleifer/bart-tiny-random"]
with patch.object(sys, "argv", testargs):
run_generate()
self.assertTrue(Path(output_file_name).exists())
+80 -2
View File
@@ -1,11 +1,15 @@
import itertools
import json
import os
import pickle
from pathlib import Path
from typing import List
from typing import Dict, Iterable, List
import git
import numpy as np
import torch
from rouge_score import rouge_scorer, scoring
from torch import nn
from torch.utils.data import Dataset, Sampler
from tqdm import tqdm
@@ -55,7 +59,7 @@ def lmap(f, x):
return list(map(f, x))
T5_PREFIX = "summarize: "
T5_PREFIX = "summarize: " # HACK, fixme
class SummarizationDataset(Dataset):
@@ -158,11 +162,85 @@ class SortishSampler(Sampler):
return iter(sort_idx)
def use_task_specific_params(model, task):
# update config with summarization specific params
task_specific_params = model.config.task_specific_params
if task_specific_params is not None:
model.config.update(task_specific_params.get(task, {}))
def pickle_load(path):
"""pickle.load(path)"""
with open(path, "rb") as f:
return pickle.load(f)
def pickle_save(obj, path):
"""pickle.dump(obj, path)"""
with open(path, "wb") as f:
return pickle.dump(obj, f)
def flatten_list(summary_ids: List[List]):
return [x for x in itertools.chain.from_iterable(summary_ids)]
def save_git_info(folder_path: str):
"""
Log commit info.
"""
repo_infos = get_git_info()
with open(os.path.join(folder_path, "git_log.json"), "w") as f:
json.dump(repo_infos, f, indent=4)
def get_git_info():
repo = git.Repo(search_parent_directories=True)
repo_infos = {
"repo_id": str(repo),
"repo_sha": str(repo.head.object.hexsha),
"repo_branch": str(repo.active_branch),
}
return repo_infos
ROUGE_KEYS = ["rouge1", "rouge2", "rougeL"]
def calculate_rouge(output_lns: List[str], reference_lns: List[str]) -> Dict:
scorer = rouge_scorer.RougeScorer(ROUGE_KEYS, use_stemmer=True)
aggregator = scoring.BootstrapAggregator()
for reference_ln, output_ln in zip(reference_lns, output_lns):
scores = scorer.score(reference_ln, output_ln)
aggregator.add_scores(scores)
result = aggregator.aggregate()
return {k: v.mid.fmeasure for k, v in result.items()}
def freeze_params(model: nn.Module):
for par in model.parameters():
par.requires_grad = False
def grad_status(model: nn.Module) -> Iterable:
return (par.requires_grad for par in model.parameters())
def any_requires_grad(model: nn.Module) -> bool:
return any(grad_status(model))
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"
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"
+1 -1
View File
@@ -59,7 +59,7 @@ BART_GENERATION_EXAMPLE = r"""
Examples::
from transformers import BartTokenizer, BartForConditionalGeneration, BartConfig
# see ``examples/summarization/bart/evaluate_cnn.py`` for a longer example
# see ``examples/summarization/bart/run_eval.py`` for a longer example
model = BartForConditionalGeneration.from_pretrained('bart-large-cnn')
tokenizer = BartTokenizer.from_pretrained('bart-large-cnn')
ARTICLE_TO_SUMMARIZE = "My friends are cool but they eat too many carbs."