| 2286 | |
| 2287 | @staticmethod |
| 2288 | def logits_to_logprobs( |
| 2289 | logits: Union[npt.NDArray[np.single], List], axis: int = -1 |
| 2290 | ) -> npt.NDArray[np.single]: |
| 2291 | # https://docs.scipy.org/doc/scipy/reference/generated/scipy.special.log_softmax.html |
| 2292 | logits_maxs: np.ndarray = np.amax(logits, axis=axis, keepdims=True) |
| 2293 | if logits_maxs.ndim > 0: |
| 2294 | logits_maxs[~np.isfinite(logits_maxs)] = 0 |
| 2295 | elif not np.isfinite(logits_maxs): |
| 2296 | logits_maxs = 0 |
| 2297 | subtract_maxs = np.subtract(logits, logits_maxs, dtype=np.single) |
| 2298 | exp = np.exp(subtract_maxs) |
| 2299 | # Suppress warnings about log of zero |
| 2300 | with np.errstate(divide="ignore"): |
| 2301 | summed = np.sum(exp, axis=axis, keepdims=True) |
| 2302 | out = np.log(summed) |
| 2303 | return subtract_maxs - out |
| 2304 | |
| 2305 | @staticmethod |
| 2306 | def longest_token_prefix(a: Sequence[int], b: Sequence[int]): |