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

Method beam_search_decode

axlearn/common/decoder.py:218–285  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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(

Callers

nothing calls this directly

Calls 6

_tokens_to_scoresMethod · 0.95
_padMethod · 0.95
validate_contains_pathsFunction · 0.90
infer_initial_time_stepFunction · 0.90
beam_search_decodeFunction · 0.90
prefill_statesMethod · 0.45

Tested by

no test coverage detected