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