| 129 | |
| 130 | |
| 131 | class T5DecoderFinalMixin(BaseMixin): |
| 132 | def __init__(self, vocab_size, hidden_size, tie_word_embeddings=True): |
| 133 | super().__init__() |
| 134 | self.hidden_size = hidden_size |
| 135 | self.tie_word_embeddings = tie_word_embeddings |
| 136 | if not tie_word_embeddings: |
| 137 | self.lm_head = VocabParallelEmbedding( |
| 138 | vocab_size, hidden_size, init_method=unscaled_init_method(0.02)) |
| 139 | |
| 140 | def final_forward(self, logits, **kwargs): |
| 141 | logits_parallel = copy_to_model_parallel_region(logits) |
| 142 | if self.tie_word_embeddings: |
| 143 | logits_parallel = logits_parallel * (self.hidden_size ** -0.5) |
| 144 | logits_parallel = F.linear(logits_parallel, self.transformer.word_embeddings.weight) |
| 145 | else: |
| 146 | logits_parallel = F.linear(logits_parallel, self.lm_head.weight) |
| 147 | return logits_parallel |
| 148 | |
| 149 | |
| 150 | def t5_gelu(x): |