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

Function log_probs_from_logits

axlearn/common/decoder.py:108–126  ·  view source on GitHub ↗

Computes log probabilities from logits, with an optional modifier. Args: logits: A float Tensor of shape [..., vocab_size]. logits_modifier: An optional function to modify the log probabilities (e.g. top-k, top-p filtering). Applied after log_softmax. Returns:

(
    logits: Tensor, logits_modifier: Optional[LogitsToLogitsFn] = None
)

Source from the content-addressed store, hash-verified

106def log_probs_from_logits(
107 logits: Tensor, logits_modifier: Optional[LogitsToLogitsFn] = None
108) -> Tensor:
109 """Computes log probabilities from logits, with an optional modifier.
110
111 Args:
112 logits: A float Tensor of shape [..., vocab_size].
113 logits_modifier: An optional function to modify the log probabilities
114 (e.g. top-k, top-p filtering). Applied after log_softmax.
115
116 Returns:
117 Log probabilities of same shape as logits.
118 """
119 if logits.dtype in (jnp.bfloat16, jnp.float16):
120 logits = logits.astype(jnp.float32)
121 log_probs = jax.nn.log_softmax(logits)
122 if logits_modifier is not None:
123 log_probs = logits_modifier(log_probs)
124 return log_probs
125
126
127# NOTE: We use a Protocol instead of defining a base layer so that decoder implementations can
128# inherit from other base classes without resorting to multiple inheritance.
129class BaseDecoder(Protocol):

Callers 2

sample_decodeMethod · 0.85
tokens_to_scoresMethod · 0.85

Calls 1

astypeMethod · 0.80

Tested by

no test coverage detected