Numerically stable softmax along last axis.
(logits: np.ndarray)
| 85 | |
| 86 | |
| 87 | def softmax(logits: np.ndarray) -> np.ndarray: |
| 88 | """Numerically stable softmax along last axis.""" |
| 89 | shifted = logits - logits.max(axis=-1, keepdims=True) |
| 90 | exp = np.exp(shifted) |
| 91 | return exp / exp.sum(axis=-1, keepdims=True) |
| 92 | |
| 93 | |
| 94 | def multinomial(probs: np.ndarray) -> np.ndarray: |