Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2788f19aa | ||
|
|
6aaac683c5 | ||
|
|
d68c671132 | ||
|
|
b9ca11e3ae | ||
|
|
6bdf998dff | ||
|
|
fcb96a2c1e | ||
|
|
e5a3f09cd4 | ||
|
|
7d5bdb82eb | ||
|
|
556fdea687 | ||
|
|
9ba04a609e | ||
|
|
d63f811af5 | ||
|
|
faac0718cb | ||
|
|
7618c08e55 | ||
|
|
1734ba169b |
No files matched your search
@@ -122,8 +122,8 @@ Best performing command:
|
||||
export ENRO_DIR='wmt_en_ro' # Download instructions above
|
||||
# export WANDB_PROJECT="MT" # optional
|
||||
export MAX_LEN=128
|
||||
export BS=4
|
||||
./train_mbart_cc25_enro.sh --output_dir enro_finetune_baseline --label_smoothing 0.1 --fp16_opt_level=O1 --logger_name wandb --sortish_sampler
|
||||
export BS=8
|
||||
./train_mbart_cc25_enro.sh --output_dir enro_finetune_baseline_dropper --label_smoothing 0 --fp16_opt_level=O1 --logger_name wandb --sortish_sampler
|
||||
```
|
||||
This should take < 6h/epoch on a 16GB v100 and achieve test BLEU above 26
|
||||
To get results in line with fairseq, you need to do some postprocessing. (see `romanian_postprocessing.md`)
|
||||
@@ -141,7 +141,7 @@ export BS=4
|
||||
As you train, `output_dir` will be filled with files, that look kind of like this (comments are mine).
|
||||
Some of them are metrics, some of them are checkpoints, some of them are metadata. Here is a quick tour:
|
||||
|
||||
```bash
|
||||
```
|
||||
output_dir
|
||||
├── best_tfmr # this is a huggingface checkpoint generated by save_pretrained. It is the same model as the PL .ckpt file below
|
||||
│ ├── config.json
|
||||
|
||||
@@ -37,6 +37,7 @@ try:
|
||||
)
|
||||
|
||||
from .callbacks import Seq2SeqLoggingCallback, get_checkpoint_callback, get_early_stopping_callback
|
||||
from .loss_dropper import LossDropper
|
||||
except ImportError:
|
||||
from utils import (
|
||||
Seq2SeqDataset,
|
||||
@@ -56,19 +57,21 @@ except ImportError:
|
||||
label_smoothed_nll_loss,
|
||||
)
|
||||
from callbacks import Seq2SeqLoggingCallback, get_checkpoint_callback, get_early_stopping_callback
|
||||
from loss_dropper import LossDropper
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SummarizationModule(BaseTransformer):
|
||||
mode = "summarization"
|
||||
loss_names = ["loss"]
|
||||
loss_names = ["loss", "dropper_mask_mean"]
|
||||
metric_names = ROUGE_KEYS
|
||||
val_metric = "rouge2"
|
||||
|
||||
def __init__(self, hparams, **kwargs):
|
||||
super().__init__(hparams, num_labels=None, mode=self.mode, **kwargs)
|
||||
use_task_specific_params(self.model, "summarization")
|
||||
self.dropper = LossDropper(dropc=.05)
|
||||
save_git_info(self.hparams.output_dir)
|
||||
self.metrics_save_path = Path(self.output_dir) / "metrics.json"
|
||||
self.hparams_save_path = Path(self.output_dir) / "hparams.pkl"
|
||||
@@ -140,7 +143,7 @@ class SummarizationModule(BaseTransformer):
|
||||
|
||||
if "labels" in batch:
|
||||
lm_labels = batch["labels"]
|
||||
decoder_input_ids = shift_tokens_right(lm_labels)
|
||||
decoder_input_ids = shift_tokens_right(lm_labels, pad_token_id)
|
||||
elif isinstance(self.model, T5ForConditionalGeneration):
|
||||
decoder_input_ids = self.model._shift_right(target_ids)
|
||||
lm_labels = target_ids
|
||||
@@ -149,19 +152,32 @@ class SummarizationModule(BaseTransformer):
|
||||
lm_labels = target_ids[:, 1:].clone() # why clone?
|
||||
|
||||
outputs = self(source_ids, attention_mask=source_mask, decoder_input_ids=decoder_input_ids, use_cache=False)
|
||||
bs = source_ids.shape[0]
|
||||
|
||||
if self.hparams.label_smoothing == 0:
|
||||
# Same behavior as modeling_bart.py, besides pad_token_id
|
||||
loss_fct = torch.nn.CrossEntropyLoss(ignore_index=pad_token_id)
|
||||
|
||||
# Same behavior as modeling_bart.py
|
||||
loss_fct = torch.nn.CrossEntropyLoss(reduction='none', ignore_index=pad_token_id)
|
||||
lm_logits = outputs[0]
|
||||
assert lm_logits.shape[-1] == self.model.config.vocab_size
|
||||
|
||||
#loss_fct = torch.nn.NLLLoss(reduction='none', ignore_index=pad_token_id)
|
||||
#logit_shape =
|
||||
#weights = torch.ones(logit_shape
|
||||
loss = loss_fct(lm_logits.view(-1, lm_logits.shape[-1]), lm_labels.view(-1))
|
||||
loss = loss.view(-1, bs)
|
||||
loss = loss.mean(dim=0)
|
||||
mask = self.dropper(loss)
|
||||
loss *= mask
|
||||
loss = loss.mean()
|
||||
return (loss, 1-mask.mean())
|
||||
#loss = loss.view(-1, bs)
|
||||
else:
|
||||
lprobs = torch.nn.functional.log_softmax(outputs[0], dim=-1)
|
||||
loss, nll_loss = label_smoothed_nll_loss(
|
||||
lprobs, lm_labels, self.hparams.label_smoothing, ignore_index=pad_token_id
|
||||
)
|
||||
return (loss,)
|
||||
return (loss,torch.tensor(1.))
|
||||
|
||||
@property
|
||||
def pad(self) -> int:
|
||||
@@ -306,6 +322,7 @@ class SummarizationModule(BaseTransformer):
|
||||
"--task", type=str, default="summarization", required=False, help="# examples. -1 means use all."
|
||||
)
|
||||
parser.add_argument("--label_smoothing", type=float, default=0.0, required=False)
|
||||
parser.add_argument("--loss_dropper", type=float, default=0.0, required=False)
|
||||
parser.add_argument("--src_lang", type=str, default="", required=False)
|
||||
parser.add_argument("--tgt_lang", type=str, default="", required=False)
|
||||
parser.add_argument(
|
||||
@@ -320,7 +337,6 @@ class SummarizationModule(BaseTransformer):
|
||||
|
||||
class TranslationModule(SummarizationModule):
|
||||
mode = "translation"
|
||||
loss_names = ["loss"]
|
||||
metric_names = ["bleu"]
|
||||
val_metric = "bleu"
|
||||
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class LossDropper(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dropc=0.4,
|
||||
min_count=10000,
|
||||
recompute=10000,
|
||||
verbose=True
|
||||
):
|
||||
super().__init__()
|
||||
self.keepc = 1. - dropc
|
||||
self.count = 0
|
||||
self.min_count = min_count
|
||||
|
||||
self.recompute = recompute
|
||||
self.last_computed = 0
|
||||
self.percentile_val = 100000000.
|
||||
self.cur_idx = 0
|
||||
|
||||
self.verbose = verbose
|
||||
|
||||
self.vals = np.zeros(self.recompute, dtype=np.float32)
|
||||
|
||||
def forward(self, loss):
|
||||
if loss is None:
|
||||
return loss
|
||||
|
||||
self.last_computed += loss.numel()
|
||||
self.count += loss.numel()
|
||||
if self.count < len(self.vals):
|
||||
self.vals[self.count - loss.numel():self.count] = loss.detach().cpu().numpy().flatten()
|
||||
self.cur_idx += loss.numel()
|
||||
return (loss < np.inf).type(loss.dtype)
|
||||
else:
|
||||
for idx, item in enumerate(loss):
|
||||
self.vals[self.cur_idx] = item
|
||||
self.cur_idx += 1
|
||||
if self.cur_idx >= len(self.vals):
|
||||
self.cur_idx = 0
|
||||
if self.count < self.min_count:
|
||||
return (loss < np.inf).type(loss.dtype)
|
||||
|
||||
if self.last_computed > self.recompute:
|
||||
self.percentile_val = np.percentile(self.vals, self.keepc * 100)
|
||||
if self.verbose:
|
||||
print('Using cutoff', self.percentile_val)
|
||||
self.last_computed = 0
|
||||
|
||||
mask = (loss < self.percentile_val).type(loss.dtype)
|
||||
return mask
|
||||
@@ -6,7 +6,7 @@ export GAS=1
|
||||
|
||||
python finetune.py \
|
||||
--learning_rate=3e-5 \
|
||||
--fp16 \
|
||||
--fp16 --fp16_opt_level=O1 \
|
||||
--gpus 1 \
|
||||
--do_train \
|
||||
--do_predict \
|
||||
|
||||
@@ -6,7 +6,7 @@ python distillation.py \
|
||||
--learning_rate=3e-4 \
|
||||
--do_train \
|
||||
--do_predict \
|
||||
--fp16 \
|
||||
--fp16 --fp16_opt_level=O1 \
|
||||
--val_check_interval 0.1 --n_val 1000 \
|
||||
--teacher facebook/bart-large-xsum --data_dir $XSUM_DIR \
|
||||
--max_target_length=60 --val_max_target_length=60 --test_max_target_length=100 \
|
||||
|
||||
@@ -14,5 +14,4 @@ python finetune.py \
|
||||
--task translation \
|
||||
--warmup_steps 500 \
|
||||
--freeze_embeds \
|
||||
--model_name_or_path=facebook/mbart-large-cc25 \
|
||||
"$@"
|
||||
Reference in new issue
Block a user