{"record":{"id":"740564700c3d1bcf","repo":"jax-ml/jax","slug":"requires-8-16-32-or-64-bit-field-width-740564","errorCode":null,"errorMessage":"requires 8-, 16-, 32- or 64-bit field width.","messagePattern":"requires 8-, 16-, 32- or 64-bit field width\\.","errorType":"exception","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/random/rbg.py","lineNumber":63,"sourceCode":"  if config.threefry_partitionable.value:\n    _threefry_split = threefry2x32._threefry_split_foldlike\n  else:\n    _threefry_split = threefry2x32._threefry_split_original\n  halfkeys = key.reshape(2, 2)\n  return api.vmap(\n      _threefry_split, (0, None), len(shape))(halfkeys, shape).reshape(\n          *shape, 4)\n\ndef _rbg_fold_in(key: typing.Array, data: typing.Array) -> typing.Array:\n  assert not data.shape\n  return api.vmap(threefry2x32._threefry_fold_in, (0, None), 0)(key.reshape(2, 2), data).reshape(4)\n\ndef _rbg_random_bits(key: typing.Array, bit_width: int, shape: Sequence[int]\n                     ) -> typing.Array:\n  if not key.shape == (4,) and key.dtype == np.dtype('uint32'):\n    raise TypeError(\"_rbg_random_bits got invalid prng key.\")\n  if bit_width not in (8, 16, 32, 64):\n    raise TypeError(\"requires 8-, 16-, 32- or 64-bit field width.\")\n  _, bits = lax.rng_bit_generator(key, shape, dtype=prng.UINT_DTYPES[bit_width])\n  return bits\n\nrbg_prng_impl = prng.PRNGImpl(\n    key_shape=(4,),\n    seed=_rbg_seed,\n    split=_rbg_split,\n    random_bits=_rbg_random_bits,\n    fold_in=_rbg_fold_in,\n    name='rbg',\n    tag='rbg')\n\nprng.register_prng(rbg_prng_impl)\n\n\ndef _unsafe_rbg_split(key: typing.Array, shape: prng.Shape) -> typing.Array:\n  # treat 10 iterations of random bits as a 'hash function'\n  num = math.prod(shape)","sourceCodeStart":45,"sourceCodeEnd":81,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/random/rbg.py#L45-L81","documentation":"Random-bit generation in JAX is only defined for unsigned integer widths 8, 16, 32, and 64, because the underlying kernels emit chunks of 32 or 64 random bits and split them evenly. _rbg_random_bits raises this TypeError when bit_width is anything else (e.g. bool, float widths, or arbitrary integers).","triggerScenarios":"Calling jax.random.bits with dtype=bool (bit_width interpreted oddly) or with a float dtype; calling lax.rng_bit_generator via rbg with an unsupported dtype; passing a computed width like 4 or 128.","commonSituations":"Trying to generate random booleans with jax.random.bits instead of comparisons (e.g. jax.random.uniform(...) < 0.5); dynamic dtype selection that yields non-uint kinds; assuming arbitrary bit widths are supported.","solutions":["Use one of uint8/uint16/uint32/uint64 with jax.random.bits","For booleans, generate uint bits and compare: jax.random.bits(key, shape, dtype=jnp.uint8, impl=...) % 2 == 0 or use jax.random.bernoulli","Validate dtype before calling in dynamic pipelines"],"exampleFix":"// before\nmask = jax.random.bits(key, shape=(4,), dtype=jnp.bool_, impl='rbg')\n\n// after\nmask = jax.random.bits(key, shape=(4,), dtype=jnp.uint8, impl='rbg') < 128","handlingStrategy":"validation","validationCode":"VALID = {8, 16, 32, 64}\nassert bit_width in VALID, f'bit_width must be one of {VALID}'","typeGuard":"def is_valid_bit_width(w) -> bool:\n    return w in (8, 16, 32, 64)","tryCatchPattern":null,"preventionTips":["Use only uint8/16/32/64 with jax.random.bits","Derive booleans from uint bits via comparison, not a bool dtype"],"tags":["jax","prng","bit-width","dtype"],"backgroundTag":"jax-random-bits-invalid-width","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}