Maps current token IDs and model state to next logits and updated state. Args: token_ids: An int Tensor of shape [batch*num_decodes, 1]. cache: A NestedTensor of cached states. Returns: (log_probs, updated_cache), where log_pr
(token_ids: Tensor, cache: NestedTensor)
| 415 | (log_probs, updated_cache), where log_probs has shape |
| 416 | [token_ids.shape[0], vocab_size] and represents log probabilities of the next |
| 417 | tokens; and updated_cache is the updated cache. |
| 418 | """ |
| 419 | time_step = cache["time_step"] |
| 420 | assert time_step.ndim == 1 |
| 421 | |
| 422 | # Select attention biases corresponding to the current time steps. |
| 423 | # We alias the nonlocal variable. |
| 424 | cross_attention_biases = cross_attention_logit_biases |
| 425 | if cross_attention_biases is not None: |
| 426 | # Note: the target_len dimension can be 1 during decoding. |
| 427 | # When indexing, we clip the indices instead of producing NaNs. |
| 428 | # TODO(markblee): Consider removing `take_along_axis` entirely if we restrict |
| 429 | # target_len to always be 1 during decoding. |
| 430 | # [batch*num_decodes, num_heads, 1, source_len]. |
| 431 | cross_attention_biases = jnp.take_along_axis( |
| 432 | cross_attention_biases, time_step[:, None, None, None], mode="clip", axis=2 |
| 433 | ) |
| 434 | |
| 435 | # Use a temporary output collection to avoid tracer leaks during extend_step. |
| 436 | with _temporary_output_collection(): |
| 437 | updated_state, outputs = self._decoder.extend_step( |
| 438 | cached_states=cache, |
| 439 | input_batch={"input_ids": token_ids}, |
| 440 | cross_attention_data=cross_attention_data, |
| 441 | cross_attention_logit_biases=cross_attention_biases, |
| 442 | ) |
| 443 | |
| 444 | logits = outputs["logits"] |
| 445 | log_probs = log_probs_from_logits(logits[:, -1, :], logits_modifier=logits_modifier) |
| 446 | return log_probs, updated_state |
| 447 | |
| 448 | return tokens_to_scores |
| 449 | |
| 450 | @staticmethod |
| 451 | def _pad(prefix: Tensor, *, max_sequence_length: int, pad_id: int) -> Tensor: |
| 452 | """Accept token IDs input tensor and pad if necessary to max_sequence_length.""" |
| 453 | return jnp.concatenate( |
| 454 | [ |
| 455 | prefix, |
| 456 | jnp.full( |
| 457 | (prefix.shape[0], max_sequence_length - prefix.shape[1]), |
| 458 | pad_id, |
| 459 | dtype=prefix.dtype, |
nothing calls this directly
no test coverage detected