MCPcopy Create free account
hub / github.com/apple/axlearn / dummy_padding_mask

Function dummy_padding_mask

axlearn/common/test_utils.py:616–638  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

614
615# TODO(markblee): Update to take prng_key explicitly.
616def 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.

Callers 2

test_decodeMethod · 0.90

Calls 1

astypeMethod · 0.80

Tested by

no test coverage detected