{"record":{"id":"791b8cd09ad65f95","repo":"jax-ml/jax","slug":"dynamic-grid-bounds-not-yet-supported-on-gpu","errorCode":null,"errorMessage":"Dynamic grid bounds not (yet) supported on GPU","messagePattern":"Dynamic grid bounds not \\(yet\\) supported on GPU","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/interpret/interpret_pallas_call.py","lineNumber":418,"sourceCode":"      gpu_callbacks.TOKEN_SHAPE_DTYPE,\n      token,\n      ordered=True,\n  )\n\n  token = gpu_callbacks.call_initialize_shared_memory(\n      token=token,\n      num_gpus=jnp.int32(device_info.num_devices),\n      num_threads_per_block=jnp.int32(num_threads_per_block),\n      num_blocks_per_cluster=jnp.int32(num_blocks_per_cluster),\n      interpret_params=interpret_params,\n  )\n\n  dynamic_grid_args, scalars, inputs = split_list(\n      args,\n      [grid_mapping.num_dynamic_grid_bounds, grid_mapping.num_index_operands],\n  )\n  if dynamic_grid_args:\n    raise NotImplementedError(\"Dynamic grid bounds not (yet) supported on GPU\")\n  if scalars:\n    raise NotImplementedError(\"Scalar arguments not (yet) supported on GPU\")\n\n  assert grid_mapping.num_index_operands == 0\n\n  token, input_buffer_keys = _allocate_buffers_for_inputs(\n      token,\n      device,\n      jaxpr.invars[: grid_mapping.num_inputs],\n      inputs,\n  )\n\n  token, output_buffers = _allocate_buffers_for_outputs(\n      token,\n      device,\n      num_threads_per_block,\n      input_output_aliases,\n      grid_mapping,","sourceCodeStart":400,"sourceCodeEnd":436,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/interpret/interpret_pallas_call.py#L400-L436","documentation":"interpret_pallas_call rejects kernels whose GridMapping declares dynamic grid bounds that arrive as runtime arguments (dynamic_grid_args non-empty). GPU interpret mode only supports fully static grids.","triggerScenarios":"Invoking a pallas kernel with dynamic grid bounds (grid sizes passed as traced arguments at call time) while GPU interpret mode is active.","commonSituations":"Debugging with interpret mode a kernel that uses call-time dynamic grids that works on device; shapes flowing through jit making grid args dynamic.","solutions":["Use a static grid for interpret-mode runs","Construct the grid from concrete Python ints outside jit","Skip interpret mode for kernels requiring dynamic grids (test on device)","Update jax in case support was added"],"exampleFix":"# before\nkernel(x, dynamic_grid_args=(n_blocks,))  # traced n_blocks\n# after\nkernel(x)  # with grid=(int(n_blocks),) fixed at pallas_call creation","handlingStrategy":"validation","validationCode":"assert not dynamic_grid_args, 'GPU interpret mode requires a static grid'","typeGuard":"def is_static_grid(grid) -> bool:\n    return isinstance(grid, tuple) and all(isinstance(g, int) for g in grid)","tryCatchPattern":null,"preventionTips":["Avoid callable/dynamic grids when interpret mode is active"],"tags":["pallas","mosaic-gpu","interpret-mode","dynamic-grid","not-implemented"],"backgroundTag":"traced-value-where-static-required","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}