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

Method beam_search_decode

axlearn/audio/decoder_asr.py:1016–1099  ·  view source on GitHub ↗

Transducer label-synchronous search. Each hypothesis in the beam has the same length of tokens, including both blank and label tokens. Args: input_batch: A dict containing: inputs: A Tensor of shape [batch_size, num_frames, dim] from encoder

(  # pytype: disable=signature-mismatch
        self,
        input_batch: Nested[Tensor],
        num_decodes: int,
        max_decode_len: int,
        prefix_merger: Optional[PrefixMerger] = None,
    )

Source from the content-addressed store, hash-verified

1014 return tokens_to_scores
1015
1016 def beam_search_decode( # pytype: disable=signature-mismatch
1017 self,
1018 input_batch: Nested[Tensor],
1019 num_decodes: int,
1020 max_decode_len: int,
1021 prefix_merger: Optional[PrefixMerger] = None,
1022 ) -> DecodeOutputs:
1023 """Transducer label-synchronous search.
1024
1025 Each hypothesis in the beam has the same length of tokens, including
1026 both blank and label tokens.
1027
1028 Args:
1029 input_batch: A dict containing:
1030 inputs: A Tensor of shape [batch_size, num_frames, dim] from encoder outputs.
1031 paddings: A 0/1 Tensor of shape [batch_size, num_frames]. 1's represent paddings.
1032 num_decodes: Beam size.
1033 max_decode_len: maximum number of decode steps to run beam search.
1034 Decoding terminates if an eos token is not emitted after max_decode_steps
1035 steps. This value can depend on the tokenization.
1036 prefix_merger: An optional PrefixMerger to apply during decoding.
1037
1038 Returns:
1039 DecodeOutputs, containing
1040 raw_sequences: An int Tensor of shape [batch_size, num_decodes, max_decode_len].
1041 sequences: An int Tensor of shape [batch_size, num_decodes, max_decode_len].
1042 paddings: A 0/1 Tensor of shape [batch_size, num_decodes, max_decode_len].
1043 scores: A Tensor of shape [batch_size, num_decodes].
1044
1045 Raises:
1046 ValueError: If max_decode_len <= src_max_len.
1047 """
1048 paddings: Tensor = input_batch["paddings"]
1049 batch_size, src_max_len = paddings.shape
1050 if max_decode_len <= src_max_len:
1051 raise ValueError(f"{max_decode_len=} is smaller than {src_max_len=}.")
1052
1053 cfg = self.config
1054 blank_id, eos_id, bos_id = cfg.blank_id, cfg.eos_id, cfg.bos_id
1055
1056 # Starts decoding with [BOS] token.
1057 inputs = jnp.zeros((batch_size, max_decode_len), dtype=jnp.int32)
1058 inputs = inputs.at[:, 0].set(bos_id)
1059
1060 init_states = {
1061 "am_step": jnp.zeros(batch_size),
1062 "lm_states": self.prediction_network.init_states(batch_size=batch_size),
1063 "lm_data": jnp.zeros((batch_size, 1, self.config.joint_dim)),
1064 "decode_step": jnp.array(0),
1065 }
1066
1067 beam_search_outputs = beam_search_decode(
1068 inputs=inputs,
1069 time_step=infer_initial_time_step(inputs, pad_id=0),
1070 cache=init_states,
1071 tokens_to_scores=self._tokens_to_scores(
1072 input_batch, num_decodes=num_decodes, max_decode_len=max_decode_len
1073 ),

Callers

nothing calls this directly

Calls 7

_tokens_to_scoresMethod · 0.95
beam_search_decodeFunction · 0.90
infer_initial_time_stepFunction · 0.90
_map_label_sequencesFunction · 0.85
DecodeOutputsClass · 0.85
setMethod · 0.45
init_statesMethod · 0.45

Tested by

no test coverage detected