{"record":{"id":"1d7cd5cf20e9506c","repo":"jax-ml/jax","slug":"cannot-operate-on-split-stateful-generator","errorCode":null,"errorMessage":"cannot operate on split stateful generator","messagePattern":"cannot operate on split stateful generator","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/random/stateful_rng.py","lineNumber":105,"sourceCode":"      shape: an optional shape if returning multiple keys.\n\n    Returns:\n      A new, independent PRNG key with the same impl/dtype as\n      ``self._base_key``.\n\n    Examples:\n      >>> from jax.experimental import random\n      >>> rng = random.stateful_rng(0)\n      >>> rng.key()\n      Array((), dtype=key<fry>) overlaying:\n      [1797259609 2579123966]\n      >>> rng.key()\n      Array((), dtype=key<fry>) overlaying:\n      [ 928981903 3453687069]\n    \"\"\"\n    if self._base_key.shape:\n      # TODO(jakevdp): better error message.\n      raise ValueError(\"cannot operate on split stateful generator\")\n\n    key = random.fold_in(self._base_key, ref_primitives.ref_get(self._counter))\n    ref_primitives.ref_addupdate(self._counter, ..., 1)\n    shape_tuple = _canonicalize_size(shape)\n    return random.split(key, shape_tuple) if shape_tuple else key\n\n  def random(\n      self,\n      size: int | Sequence[int] | None = None,\n      dtype: DTypeLike = float,\n  ):\n    \"\"\"Return random floats in the half-open interval [0.0, 1.0).\"\"\"\n    # TODO(jakevdp): write docstring\n    return random.uniform(self.key(), shape=_canonicalize_size(size), dtype=dtype)\n\n\n  def uniform(\n      self,","sourceCodeStart":87,"sourceCodeEnd":123,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/random/stateful_rng.py#L87-L123","documentation":"jax.experimental.random.stateful_rng() supports splitting the generator so children draw independent streams, but a split generator's _base_key has a non-empty shape. Operations that need to advance the scalar counter (key(), random(), uniform(), normal(), integers(), split(), spawn()) refuse to run on such a split generator via this ValueError.","triggerScenarios":"Calling rng.key() or any sampling method on an object obtained from rng.spawn() or rng.split(); reusing a child generator for direct draws instead of spawning further children.","commonSituations":"Adopting the stateful API and assuming split children behave like independent top-level generators; passing spawned generators into code that calls rng.random() directly.","solutions":["Only call draw methods on the top-level generator; use spawned children to spawn further or pass as base_key to new stateful_rng instances","Use rng.spawn() at the point where independent substreams are needed instead of pre-splitting","Restructure so each worker receives its own stateful_rng(seed=unique_seed)"],"exampleFix":"// before\nchildren = rng.split(2)\nk = children[0].key()  # ValueError: cannot operate on split stateful generator\n\n// after\nchild = rng.spawn(1)[0]\nrng2 = stateful_rng(base_key=child.key())\nk = rng2.key()","handlingStrategy":"validation","validationCode":"import jax\nif jax.random.key_data(rng._base_key).shape != (2,):\n    raise ValueError('generator is split; spawn children instead of drawing')","typeGuard":"def is_unsplit_generator(rng) -> bool:\n    return rng._base_key.shape == ()","tryCatchPattern":null,"preventionTips":["Draw only from top-level generators; use spawn() for children","Give each worker its own stateful_rng(seed)"],"tags":["jax","prng","stateful-rng","split","spawn"],"backgroundTag":"jax-stateful-rng-split-misuse","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}