jax.numpy.trunc#

jax.numpy.trunc(x)[source]#

將輸入四捨五入到最接近零的整數。

numpy.trunc() 的 JAX 實作。

參數:

x (ArrayLike) – 輸入陣列或純量。

回傳:

一個與 x 具有相同形狀和 dtype 的陣列,其中包含四捨五入的值。

回傳類型:

陣列

另請參閱

範例

>>> 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)
[[-0.23  3.6   2.33]
 [ 1.22 -0.99  1.72]
 [-8.5   5.5   3.98]]
>>> jnp.trunc(x)
Array([[-0.,  3.,  2.],
       [ 1., -0.,  1.],
       [-8.,  5.,  3.]], dtype=float32)