| 375 | |
| 376 | |
| 377 | class T5Model(nn.Module): |
| 378 | |
| 379 | def __init__(self, |
| 380 | vocab_size, |
| 381 | dim, |
| 382 | dim_attn, |
| 383 | dim_ffn, |
| 384 | num_heads, |
| 385 | encoder_layers, |
| 386 | decoder_layers, |
| 387 | num_buckets, |
| 388 | shared_pos=True, |
| 389 | dropout=0.1): |
| 390 | super(T5Model, self).__init__() |
| 391 | self.vocab_size = vocab_size |
| 392 | self.dim = dim |
| 393 | self.dim_attn = dim_attn |
| 394 | self.dim_ffn = dim_ffn |
| 395 | self.num_heads = num_heads |
| 396 | self.encoder_layers = encoder_layers |
| 397 | self.decoder_layers = decoder_layers |
| 398 | self.num_buckets = num_buckets |
| 399 | |
| 400 | # layers |
| 401 | self.token_embedding = nn.Embedding(vocab_size, dim) |
| 402 | self.encoder = T5Encoder(self.token_embedding, dim, dim_attn, dim_ffn, |
| 403 | num_heads, encoder_layers, num_buckets, |
| 404 | shared_pos, dropout) |
| 405 | self.decoder = T5Decoder(self.token_embedding, dim, dim_attn, dim_ffn, |
| 406 | num_heads, decoder_layers, num_buckets, |
| 407 | shared_pos, dropout) |
| 408 | self.head = nn.Linear(dim, vocab_size, bias=False) |
| 409 | |
| 410 | # initialize weights |
| 411 | self.apply(init_weights) |
| 412 | |
| 413 | def forward(self, encoder_ids, encoder_mask, decoder_ids, decoder_mask): |
| 414 | x = self.encoder(encoder_ids, encoder_mask) |
| 415 | x = self.decoder(decoder_ids, decoder_mask, x, encoder_mask) |
| 416 | x = self.head(x) |
| 417 | return x |
| 418 | |
| 419 | |
| 420 | def _t5(name, |
nothing calls this directly
no outgoing calls
no test coverage detected