Holds beam search state data.
| 338 | |
| 339 | |
| 340 | class _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 | |
| 357 | class BeamSearchOutputs(flax_struct.PyTreeNode): |
no outgoing calls
no test coverage detected