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,
)
| 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 | ), |
nothing calls this directly
no test coverage detected