Checks that all elements in `x` are finite.
(x: Tensor, msg_fmt: str = "", **msg_kwargs)
| 357 | |
| 358 | |
| 359 | def check_numerics(x: Tensor, msg_fmt: str = "", **msg_kwargs): |
| 360 | """Checks that all elements in `x` are finite.""" |
| 361 | global _enable_numeric_checks # pylint: disable=global-statement,global-variable-not-assigned |
| 362 | if _enable_numeric_checks: |
| 363 | assert bool(jnp.isfinite(x).all()), f"Check numerics {msg_fmt.format(**msg_kwargs)}: {x}" |
| 364 | return x |
| 365 | |
| 366 | |
| 367 | def shapes(nested_tensor: NestedTensor) -> NestedTree: |
no outgoing calls
no test coverage detected