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
)
| 106 | def 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. |
| 129 | class BaseDecoder(Protocol): |
no test coverage detected