Args: query: [batch_size, len_q, d_model] key: [batch_size, len_k, d_model] value: [batch_size, len_v(=len_k), d_model] attn_mask: [batch_size, seq_len, seq_len] Returns:
(self, query, key, value, attn_mask)
| 396 | self.layer_norm = LayerNorm(d_model) |
| 397 | |
| 398 | def forward(self, query, key, value, attn_mask): |
| 399 | """ |
| 400 | Args: |
| 401 | query: [batch_size, len_q, d_model] |
| 402 | key: [batch_size, len_k, d_model] |
| 403 | value: [batch_size, len_v(=len_k), d_model] |
| 404 | attn_mask: [batch_size, seq_len, seq_len] |
| 405 | Returns: |
| 406 | """ |
| 407 | residual = query |
| 408 | batch_size = query.shape[0] |
| 409 | |
| 410 | # (B, S, D) -proj-> (B, S, D_new) -split-> (B, S, H, W) -trans-> (B, H, S, W) |
| 411 | Q = self.W_Q(query) |
| 412 | Q = autograd.reshape(Q, [batch_size, -1, self.n_head, self.d_k]) |
| 413 | Q = autograd.transpose(Q, [0, 2, 1, 3]) |
| 414 | |
| 415 | K = self.W_K(key) |
| 416 | K = autograd.reshape(K, [batch_size, -1, self.n_head, self.d_k]) |
| 417 | K = autograd.transpose(K, [0, 2, 1, 3]) |
| 418 | |
| 419 | V = self.W_V(value) |
| 420 | V = autograd.reshape(V, [batch_size, -1, self.n_head, self.d_v]) |
| 421 | V = autograd.transpose(V, [0, 2, 1, 3]) |
| 422 | |
| 423 | # Q: [batch_size, n_heads, len_q, d_k] |
| 424 | # K: [batch_size, n_heads, len_k, d_k] |
| 425 | # V: [batch_size, n_heads, len_v(=len_k), d_v] |
| 426 | |
| 427 | # attn_mask : [batch_size, n_heads, seq_len, seq_len] |
| 428 | attn_mask = MultiHeadAttention._get_attn_mask(attn_mask, self.n_head) |
| 429 | |
| 430 | # context: [batch_size, n_heads, len_q, d_v] |
| 431 | # attn: [batch_size, n_heads, seq_len, seq_len] |
| 432 | context, attn = self.scaled_dot_product_attention(Q, K, V, attn_mask) |
| 433 | context = autograd.transpose(context, [0, 2, 1, 3]) |
| 434 | # context: [batch_size, len_q, n_heads * d_v] |
| 435 | context = autograd.reshape(context, [batch_size, -1, self.n_head * self.d_v]) |
| 436 | |
| 437 | output = self.linear(context) |
| 438 | output = self.add(output, residual) |
| 439 | # [batch_size, len_q, d_model] |
| 440 | output = self.layer_norm(output) |
| 441 | return output, attn |
| 442 | |
| 443 | @staticmethod |
| 444 | def _get_attn_mask(attn_mask, n_head): |
nothing calls this directly
no test coverage detected