Removes blanks, paddings, and repeats from the input sequences, as used in CTC or RNN-T. Note that unless pad_id is the same as blank_id, pad_id should not be in the range of the vocab. We use pad_id to infer padding positions in the inputs. Args: inputs: An int Tensor of shape
(
inputs: Tensor, *, remove_repeats: bool, blank_id: int = 0, pad_id: int = 0
)
| 648 | |
| 649 | |
| 650 | def _map_label_sequences( |
| 651 | inputs: Tensor, *, remove_repeats: bool, blank_id: int = 0, pad_id: int = 0 |
| 652 | ) -> Nested[Tensor]: |
| 653 | """Removes blanks, paddings, and repeats from the input sequences, as used in CTC or RNN-T. |
| 654 | |
| 655 | Note that unless pad_id is the same as blank_id, pad_id should not be in the range of the vocab. |
| 656 | We use pad_id to infer padding positions in the inputs. |
| 657 | |
| 658 | Args: |
| 659 | inputs: An int Tensor of shape [..., max_decode_len] containing decoded sequences. |
| 660 | remove_repeats: A boolean indicating whether we remove repeats or not. It is True for CTC, |
| 661 | False for RNN-T. |
| 662 | blank_id: Token ID corresponding to blanks. |
| 663 | pad_id: Token ID corresponding to paddings. |
| 664 | |
| 665 | Returns: |
| 666 | A dict containing: |
| 667 | sequences: A Tensor of shape [..., max_decode_len] containing label sequences. |
| 668 | paddings: A 0/1 Tensor of shape [..., max_decode_len]. 1's represent paddings. |
| 669 | lengths: A Tensor of shape [..., 1] containing the length of each sequence. |
| 670 | """ |
| 671 | max_decode_len = inputs.shape[-1] |
| 672 | # Identify points at which the token is a valid label token to keep. |
| 673 | # `indicators` has shape [batch_size, num_decodes, max_decode_len], and has a value of 1 |
| 674 | # in positions corresponding to inputs we intend to keep, |
| 675 | # i.e., the token is not blank or padding. |
| 676 | indicators = (inputs != blank_id) & (inputs != pad_id) |
| 677 | if remove_repeats: |
| 678 | # Identify points at which curr != prev. |
| 679 | # I.e., the token is not blank or padding, and is different from the previous token. |
| 680 | y = jnp.concatenate([jnp.full(inputs.shape[:-1] + (1,), pad_id), inputs], axis=-1) |
| 681 | indicators = (y[..., 1:] != y[..., :-1]) & indicators |
| 682 | |
| 683 | # Compute lengths of final sequences. [..., 1]. |
| 684 | lens = jnp.sum(indicators, axis=-1, keepdims=True, dtype=inputs.dtype) |
| 685 | |
| 686 | # Compute sequences by left-justifying the tokens-to-keep. Under jit, we use a dispatch matrix |
| 687 | # of shape [batch_size, num_decodes, max_decode_len, max_decode_len]. |
| 688 | # dispatch[..., from, to] == 1 means inputs[:, :, from] is put at sequences[:, :, to]. |
| 689 | # dispatch[..., i, :] == 0 means we drop token i in the inputs. |
| 690 | # [batch_size, num_decodes, max_decode_len, max_decode_len]. |
| 691 | dispatch = jax.nn.one_hot( |
| 692 | jnp.cumsum(indicators, axis=-1) * indicators - 1, max_decode_len, dtype=inputs.dtype |
| 693 | ) |
| 694 | sequences = jnp.einsum("...nm,...n->...m", dispatch, inputs) |
| 695 | paddings = jnp.arange(max_decode_len) >= lens |
| 696 | if pad_id != 0: |
| 697 | sequences = jnp.where(paddings, pad_id, sequences) |
| 698 | return dict(sequences=sequences, paddings=paddings, lengths=lens) |
| 699 | |
| 700 | |
| 701 | class RNNPredictionNetwork(BaseLayer): |
no outgoing calls