MCPcopy Create free account
hub / github.com/ModalityDance/Omni-R1 / stable_softmax

Function stable_softmax

src/transformers/src/transformers/tf_utils.py:51–72  ·  view source on GitHub ↗

Stable wrapper that returns the same output as `tf.nn.softmax`, but that works reliably with XLA on CPU. It is meant as a workaround for the [following issue](https://github.com/tensorflow/tensorflow/issues/55682), and will be removed after it gets fixed. The arguments and outputs are t

(logits: tf.Tensor, axis: Optional[int] = None, name: Optional[str] = None)

Source from the content-addressed store, hash-verified

49
50
51def stable_softmax(logits: tf.Tensor, axis: Optional[int] = None, name: Optional[str] = None) -> tf.Tensor:
52 """
53 Stable wrapper that returns the same output as `tf.nn.softmax`, but that works reliably with XLA on CPU. It is
54 meant as a workaround for the [following issue](https://github.com/tensorflow/tensorflow/issues/55682), and will be
55 removed after it gets fixed. The arguments and outputs are the same as `tf.nn.softmax`, and relies on the fact that
56 `softmax(x) = softmax(x + c)` (see https://ogunlao.github.io/2020/04/26/you_dont_really_know_softmax.html).
57
58 Args:
59 logits (`tf.Tensor`):
60 Must be one of the following types: half, float32, float64.
61 axis (`int`, *optional*):
62 The dimension softmax would be performed on. The default is -1 which indicates the last dimension.
63 name (`str`, *optional*):
64 A name for the operation.
65
66 Returns:
67 `tf.Tensor`:
68 A Tensor. Has the same type and shape as logits.
69 """
70 # TODO: When the issue linked above gets sorted, add a check on TF version here and use the original function if
71 # it has the fix. After we drop the support for unfixed versions, remove this function.
72 return tf.nn.softmax(logits=logits + 1e-9, axis=axis, name=name)
73
74
75def functional_layernorm(inputs, weight, bias, epsilon=1e-5, axis=-1):

Callers 15

masked_softmaxMethod · 0.90
__call__Method · 0.85
postprocessMethod · 0.85
postprocessMethod · 0.85
hard_softmaxFunction · 0.85
gumbel_softmaxFunction · 0.85
get_attnMethod · 0.85
callMethod · 0.85
callMethod · 0.85

Calls

no outgoing calls

Tested by 2

masked_softmaxMethod · 0.72