Compare commits

...
Author SHA1 Message Date
Sam Shleifer c2788f19aa boom boom 2020-08-21 19:55:50 -04:00
Sam Shleifer 6aaac683c5 merge batch parity 2020-08-21 19:54:03 -04:00
Sam Shleifer d68c671132 Merge branch 'master' into dropper-celoss 2020-08-21 19:50:36 -04:00
Sam Shleifer b9ca11e3ae Dropc 0.5 2020-08-18 22:03:35 -04:00
Sam Shleifer 6bdf998dff dropc03 2020-08-18 10:17:15 -04:00
Sam Shleifer fcb96a2c1e test dropc=0, ce loss 2020-08-17 23:48:58 -04:00
Sam Shleifer e5a3f09cd4 dropc_zero 2020-08-17 23:44:48 -04:00
Sam Shleifer 7d5bdb82eb Merge branch 'master' into dropper 2020-08-17 23:42:36 -04:00
Sam Shleifer 556fdea687 boom boom 2020-08-16 22:50:05 -04:00
Sam Shleifer 9ba04a609e Merge branch 'master' into dropper 2020-08-16 22:50:02 -04:00
Sam Shleifer d63f811af5 Doesnt break 2020-08-16 22:43:36 -04:00
Sam Shleifer faac0718cb boom boom 2020-08-16 21:36:09 -04:00
Sam Shleifer 7618c08e55 merge master 2020-08-16 21:06:19 -04:00
Sam Shleifer 1734ba169b asked for help 2020-07-22 10:54:11 -04:00
6 changed files with 80 additions and 12 deletions

No files matched your search

+3 -3
View File
@@ -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
+22 -6
View File
@@ -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"
+53
View File
@@ -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
+1 -1
View File
@@ -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 \
+1 -1
View File
@@ -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 \
"$@"