{"record":{"id":"6a8ebf15369fd2ea","repo":"jax-ml/jax","slug":"requires-8-16-32-or-64-bit-field-width-6a8ebf","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/threefry2x32.py","lineNumber":326,"sourceCode":"  bits1, bits2 = threefry2x32_p.bind(k1, k2, counts1, counts2)\n  return jnp.stack([bits1, bits2], axis=bits1.ndim)\n\n\ndef threefry_fold_in(key: typing.Array, data: typing.Array) -> typing.Array:\n  assert not data.shape\n  return _threefry_fold_in(key, jnp.asarray(data, dtype='uint32'))\n\n@api.jit\ndef _threefry_fold_in(key, data):\n  return threefry_2x32(key, threefry_seed(data))\n\n\ndef threefry_random_bits(key: typing.Array, bit_width, shape):\n  \"\"\"Sample uniform random bits of given width and shape using PRNG key.\"\"\"\n  if not _is_threefry_prng_key(key):\n    raise TypeError(\"threefry_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\n  if config.threefry_partitionable.value:\n    return _threefry_random_bits_partitionable(key, bit_width, shape)\n  else:\n    return _threefry_random_bits_original(key, bit_width, shape)\n\ndef _threefry_random_bits_partitionable(key: typing.Array, bit_width, shape):\n  if all(core.is_constant_dim(d) for d in shape) and math.prod(shape) > 2 ** 64:\n    raise NotImplementedError('random bits array of size exceeding 2 ** 64')\n\n  k1, k2 = key\n  counts1, counts2 = prng.iota_2x32_shape(shape)\n  bits1, bits2 = threefry2x32_p.bind(k1, k2, counts1, counts2)\n\n  dtype = prng.UINT_DTYPES[bit_width]\n  if bit_width == 64:\n    bits_hi = lax.convert_element_type(bits1, dtype)\n    bits_lo = lax.convert_element_type(bits2, dtype)","sourceCodeStart":308,"sourceCodeEnd":344,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/random/threefry2x32.py#L308-L344","documentation":"threefry_random_bits only produces unsigned integer output of width 8, 16, 32, or 64 bits, matching the UINT_DTYPES the kernels can emit. This TypeError fires when bit_width is outside that set — e.g. requesting float widths or other integer sizes through the internal API (the public equivalent is jax.random.bits with a bad dtype).","triggerScenarios":"Calling threefry_random_bits with bit_width in {4, 128} or derived from a float dtype; jax.random.bits(..., dtype=jnp.float32); dynamic width computation landing on an unsupported value.","commonSituations":"Users expecting float uniform values from a bits API (should use jax.random.uniform); dtype-driven generic code passing arbitrary dtypes; porting NumPy randint-style code with unusual widths.","solutions":["Use uint8/uint16/uint32/uint64 only, then convert: bits.astype(jnp.float32)/2**32 for floats","Use jax.random.uniform/normal for real-valued randomness","Validate bit_width against (8,16,32,64) in dynamic pipelines"],"exampleFix":"// before\nu = jax.random.bits(key, shape=(4,), dtype=jnp.float32)\n\n// after\nu = jax.random.bits(key, shape=(4,), dtype=jnp.uint32).astype(jnp.float32) / 2**32","handlingStrategy":"validation","validationCode":"VALID = (8, 16, 32, 64)\nassert bit_width in VALID, f'bit_width must be in {VALID}'","typeGuard":"def is_valid_bit_width(w) -> bool:\n    return w in (8, 16, 32, 64)","tryCatchPattern":null,"preventionTips":["Use only unsigned integer dtypes with bit-generating APIs","Convert bits to floats explicitly with astype after generation"],"tags":["jax","prng","threefry","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"}