jax.nn.logsumexpΒΆ

jax.nn.
logsumexp
(a, axis=None, b=None, keepdims=False, return_sign=False)[source]ΒΆ Compute the log of the sum of exponentials of input elements.
LAXbackend implementation of
logsumexp()
.Original docstring below.
 Parameters
a (array_like) β Input array.
axis (None or int or tuple of ints, optional) β Axis or axes over which the sum is taken. By default axis is None, and all elements are summed.
keepdims (bool, optional) β If this is set to True, the axes which are reduced are left in the result as dimensions with size one. With this option, the result will broadcast correctly against the original array.
b (arraylike, optional) β Scaling factor for exp(a) must be of the same shape as a or broadcastable to a. These values may be negative in order to implement subtraction.
return_sign (bool, optional) β If this is set to True, the result will be a pair containing sign information; if False, results that are negative will be returned as NaN. Default is False (no sign information).
 Returns
res (ndarray) β The result,
np.log(np.sum(np.exp(a)))
calculated in a numerically more stable way. If b is given thennp.log(np.sum(b*np.exp(a)))
is returned.sgn (ndarray) β If return_sign is True, this will be an array of floatingpoint numbers matching res and +1, 0, or 1 depending on the sign of the result. If False, only one result is returned.