| 334 | |
| 335 | |
| 336 | class TransformerEncoderLayer(layer.Layer): |
| 337 | def __init__(self, d_model=512, n_head=8, dim_feedforward=2048): |
| 338 | super(TransformerEncoderLayer, self).__init__() |
| 339 | self.d_model = d_model |
| 340 | self.n_head = n_head |
| 341 | self.dim_feedforward = dim_feedforward |
| 342 | self.enc_self_attn = MultiHeadAttention(d_model, n_head) |
| 343 | self.pos_ffn = PoswiseFeedForwardNet(d_model=d_model, dim_feedforward=dim_feedforward, bias=False) |
| 344 | |
| 345 | def forward(self, enc_inputs, enc_self_attn_mask): |
| 346 | """ |
| 347 | Args: |
| 348 | enc_inputs: [batch_size, src_len, d_model] |
| 349 | enc_self_attn_mask: [batch_size, src_len, src_len] |
| 350 | |
| 351 | Returns: |
| 352 | enc_outputs: [batch_size, src_len, d_model] |
| 353 | attn: [batch_size, n_heads, src_len, src_len] |
| 354 | """ |
| 355 | # enc_inputs to same Q,K,V |
| 356 | enc_outputs, attn = self.enc_self_attn(enc_inputs, enc_inputs, enc_inputs, enc_self_attn_mask) |
| 357 | enc_outputs = self.pos_ffn(enc_outputs) |
| 358 | return enc_outputs, attn |
| 359 | |
| 360 | |
| 361 | def matmul4d(x1, x2): |