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

Class _BeamState

axlearn/common/decoding.py:340–354  ·  view source on GitHub ↗

Holds beam search state data.

Source from the content-addressed store, hash-verified

338
339
340class _BeamState(NamedTuple):
341 """Holds beam search state data."""
342
343 # The position of the decoding loop in the length dimension.
344 cur_index: Tensor # scalar int32: current decoded length index.
345 # The active sequence log probabilities and finished sequence scores.
346 live_scores: Tensor # float32: [batch_size, beam_size].
347 finished_scores: Tensor # float32: [batch_size, beam_size].
348 # The current active-beam-searching and finished sequences.
349 live_seqs: Tensor # int32: [batch_size, beam_size, max_decode_len].
350 finished_seqs: Tensor # int32: [batch_size, beam_size, max_decode_len].
351 # The current state of the autoregressive decoding caches.
352 cache: NestedTensor
353 # The prefix merger state.
354 prefix_merger: NestedTensor
355
356
357class BeamSearchOutputs(flax_struct.PyTreeNode):

Callers 2

_beam_initFunction · 0.85
beam_search_loop_body_fnFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected