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

Function _map_label_sequences

axlearn/audio/decoder_asr.py:650–698  ·  view source on GitHub ↗

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
)

Source from the content-addressed store, hash-verified

648
649
650def _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
701class RNNPredictionNetwork(BaseLayer):

Callers 5

test_greedy_decodeMethod · 0.90
init_stateMethod · 0.85
greedy_decodeMethod · 0.85
_postprocess_outputsMethod · 0.85
beam_search_decodeMethod · 0.85

Calls

no outgoing calls

Tested by 1

test_greedy_decodeMethod · 0.72