Perform beam search decoding. Args: input_batch: A dict containing: prefix: The prefix to use for prompting of shape [batch, max_prefix_length]. The prefix for each example in the batch should begin with a prompt token (e.g.
(
self,
*,
input_batch: Nested[Tensor],
max_sequence_length: int,
num_decodes: int,
cross_attention_data: Optional[Tensor] = None,
cross_attention_logit_biases: Optional[Tensor] = None,
brevity_penalty: Optional[BrevityPenaltyFn] = None,
)
| 216 | |
| 217 | def beam_search_decode( |
| 218 | self, |
| 219 | *, |
| 220 | input_batch: Nested[Tensor], |
| 221 | max_sequence_length: int, |
| 222 | num_decodes: int, |
| 223 | cross_attention_data: Optional[Tensor] = None, |
| 224 | cross_attention_logit_biases: Optional[Tensor] = None, |
| 225 | brevity_penalty: Optional[BrevityPenaltyFn] = None, |
| 226 | ) -> BeamSearchOutputs: |
| 227 | """Perform beam search decoding. |
| 228 | |
| 229 | Args: |
| 230 | input_batch: A dict containing: |
| 231 | prefix: The prefix to use for prompting of shape [batch, max_prefix_length]. |
| 232 | The prefix for each example in the batch should begin with a prompt token (e.g. |
| 233 | BOS). |
| 234 | The prefix will be padded with `cfg.pad_token_id` to `max_sequence_length`, thus |
| 235 | it is expected that `max_prefix_length <= max_sequence_length`. |
| 236 | max_sequence_length: The maximum sequence length of tokens to generate. |
| 237 | num_decodes: The number of decoded sequences to return. These are the number of |
| 238 | hypotheses per batch example. |
| 239 | cross_attention_data: A float Tensor of shape [batch_size, source_len, hidden_dim]. |
| 240 | cross_attention_logit_biases: A Tensor of shape [batch_size, target_len, source_len]. |
| 241 | A -inf represents a disconnected position pair. |
| 242 | `target_len` should be broadcastable to `max_sequence_length`. |
| 243 | brevity_penalty: Brevity penalty function for length normalization during beam search. |
| 244 | |
| 245 | Returns: |
| 246 | The beam search outputs. |
| 247 | |
| 248 | Raises: |
| 249 | ValueError: If pad_token_id is non-zero. |
| 250 | """ |
| 251 | validate_contains_paths(input_batch, paths=["prefix"]) |
| 252 | prefix = input_batch["prefix"] |
| 253 | |
| 254 | cfg = self.config |
| 255 | tokens_to_scores_fn = self._tokens_to_scores( |
| 256 | num_decodes=num_decodes, |
| 257 | cross_attention_data=cross_attention_data, |
| 258 | cross_attention_logit_biases=cross_attention_logit_biases, |
| 259 | logits_modifier=None, |
| 260 | ) |
| 261 | input_ids = self._pad( |
| 262 | prefix, max_sequence_length=max_sequence_length, pad_id=cfg.pad_token_id |
| 263 | ) |
| 264 | time_step = infer_initial_time_step(prefix, pad_id=cfg.pad_token_id) |
| 265 | prefill_batch = {**input_batch} |
| 266 | prefill_batch["input_ids"] = input_ids |
| 267 | # Note: it prefills `k-1` tokens and used the last prefix token in first generation loop. |
| 268 | # TODO(axlearn-dev): prefill all prefix tokens in one shot. |
| 269 | init_states, _ = self._decoder.prefill_states( |
| 270 | time_step=time_step, |
| 271 | input_batch=prefill_batch, |
| 272 | cross_attention_data=cross_attention_data, |
| 273 | cross_attention_logit_biases=cross_attention_logit_biases, |
| 274 | ) |
| 275 | return beam_search_decode( |
nothing calls this directly
no test coverage detected