Builds a dummy attention mask where non-padding tokens are followed by padding tokens. Example: batch_size: 3 max_seq_len: 3 output: [[1, 1, 0], [1, 0, 0], [1, 1, 1]] Args: batch_size: Batch size. max_seq_len: Sequence length. Returns: A
(*, batch_size: int, max_seq_len: int)
| 614 | |
| 615 | # TODO(markblee): Update to take prng_key explicitly. |
| 616 | def dummy_padding_mask(*, batch_size: int, max_seq_len: int) -> Tensor: |
| 617 | """Builds a dummy attention mask where non-padding tokens are followed by padding tokens. |
| 618 | |
| 619 | Example: |
| 620 | batch_size: 3 |
| 621 | max_seq_len: 3 |
| 622 | output: [[1, 1, 0], [1, 0, 0], [1, 1, 1]] |
| 623 | |
| 624 | Args: |
| 625 | batch_size: Batch size. |
| 626 | max_seq_len: Sequence length. |
| 627 | |
| 628 | Returns: |
| 629 | A bool attention mask of shape [batch, max_seq_len]. A value of 0 indicates a padding |
| 630 | token, whereas 1 indicates non-padding. Each example has at least one non-padding, |
| 631 | followed by padding up to seq_len. |
| 632 | """ |
| 633 | lower_diag = jnp.arange(max_seq_len) |
| 634 | lower_diag = lower_diag[None, :] <= lower_diag[:, None] |
| 635 | input_len = jax.random.randint( |
| 636 | jax.random.PRNGKey(123), shape=(batch_size,), minval=0, maxval=max_seq_len |
| 637 | ) |
| 638 | return lower_diag[input_len].astype(jnp.bool) |
| 639 | |
| 640 | |
| 641 | # TODO(markblee): Update to take prng_key explicitly. |
no test coverage detected