MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / T5DecoderFinalMixin

Class T5DecoderFinalMixin

SwissArmyTransformer/sat/model/official/t5_model.py:131–147  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

129
130
131class 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
150def t5_gelu(x):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected