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

Method tokens_to_scores

axlearn/common/decoder.py:417–456  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 3

log_probs_from_logitsFunction · 0.85
extend_stepMethod · 0.45

Tested by

no test coverage detected