jax.numpy.trunc

目录

jax.numpy.trunc#

jax.numpy.trunc(x)[源代码][源代码]#

将输入四舍五入到最接近的整数,趋向于零。

JAX 实现的 numpy.trunc()

参数:

x (ArrayLike) – 输入数组或标量。

返回:

一个与 x 形状和数据类型相同的数组,包含四舍五入后的值。

返回类型:

Array

参见

示例

>>> key = jax.random.key(42)
>>> x = jax.random.uniform(key, (3, 3), minval=-10, maxval=10)
>>> with jnp.printoptions(precision=2, suppress=True):
...     print(x)
[[ 2.88 -3.55 -6.13]
 [ 7.73  4.49 -6.16]
 [-3.1  -4.95  2.64]]
>>> jnp.trunc(x)
Array([[ 2., -3., -6.],
       [ 7.,  4., -6.],
       [-3., -4.,  2.]], dtype=float32)