Compare commits

...
Author SHA1 Message Date
Patrick von Platen 27fbbeab1b Add clear description of how to train T5 2020-03-29 03:43:46 +02:00
+31
View File
@@ -16,6 +16,37 @@ To facilitate future work on transfer learning for NLP, we release our dataset,
The Authors' code can be found `here <https://github.com/google-research/text-to-text-transfer-transformer>`_ .
Training
~~~~~~~~~~~~~~~~~~~~
T5 is an encoder-decoder model and casts all NLP problems as sequence to sequence tasks. It is trained using teacher forcing.
This means that for training we always need an input sequence and a target sequence.
The input sequence is fed to the model using ``input_ids``. In teacher forcing style, the target sequence shifted to the right, *i.e.* prepended by the PAD "<pad>" token, is fed to the decoder and the target sequence appended by the EOS "</s>" token represents the ``lm_labels``.
T5 can be trained / fine-tuned on two types of objectives:
- Unsupervised denoising training
In this setup spans of the input sequence are masked by so-called sentinel tokens (*a.k.a* unique mask tokens)
and the output sequence is formed as a concatenation of the same sentinel tokens and the *real* masked tokens.
*E.g.* the sentence "The cute dog walks in the park" and mask "cute dog" and "the" should be processed as follows:
::
input_ids = tokenizer.encode('The <extra_id_1> walks in <extra_id_2> park')
decoder_input_ids = tokenizer.encode('<pad> <extra_id_1> cute dog <extra_id_2> the <extra_id_3>')
lm_labels = tokenizer.encode('<extra_id_1> cute dog <extra_id_2> the <extra_id_3> </s>')
model(input_ids=input_ids, decoder_input_ids=decoder_input_ids, lm_labels=lm_labels)
- Supervised training
In this setup the input sequence and output sequence are standart sequence to sequence input output mapping.
In translation, *e.g.* the input sequence "The house is wonderful." and output sequence "Das Haus ist wunderbar." should
be processed as follows:
::
input_ids = tokenizer.encode('The house is wonderful. </s>')
decoder_input_ids = tokenizer.encode('<pad> Das Haus ist wunderbar. ')
lm_labels = tokenizer.encode('Das Haus ist wunderbar. </s>')
model(input_ids=input_ids, decoder_input_ids=decoder_input_ids, lm_labels=lm_labels)
Tips
~~~~~~~~~~~~~~~~~~~~
- T5 is an encoder-decoder model pre-trained on a multi-task mixture of unsupervised