{"record":{"id":"2d9eb891acd43bb9","repo":"jax-ml/jax","slug":"cannot-partition-over-multiple-dynamic-parallel-di","errorCode":null,"errorMessage":"Cannot partition over multiple dynamic parallel dimensions: {grid=}","messagePattern":"Cannot partition over multiple dynamic parallel dimensions: (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/pipeline.py","lineNumber":1612,"sourceCode":"        first_divisible_dimension,\n        partitioned_dim_offset,\n    )\n    return new_grid, offsets\n\n  # Separate the remaining dimensions into dynamic and static.\n  dynamic_dims = [\n      i\n      for i in range(len(grid))\n      if i in parallel_dimensions and not isinstance(grid[i], int)\n  ]\n  static_dims = [\n      i\n      for i in range(len(grid))\n      if i in parallel_dimensions and isinstance(grid[i], int)\n  ]\n\n  if len(dynamic_dims) > 1:\n    raise NotImplementedError(\n        f\"Cannot partition over multiple dynamic parallel dimensions: {grid=}\"\n    )\n\n  if dynamic_dims and not static_dims:\n    # Exactly one dynamic dimension and no static non-divisible dimensions\n    partition_dimension = dynamic_dims[0]\n  else:\n    # No divisible static dimensions, so we can't evenly partition the grid.\n    # Let's pick the largest dimension and try to divide it as evenly as\n    # possible.\n    # TODO(sharadmv): take the product of many nondivisible dimensions to\n    # potentially divide it more evenly\n    largest_parallel_dimension = max(grid[i] for i in static_dims)\n    partition_dimension, *_ = (\n        i for i in static_dims if grid[i] == largest_parallel_dimension\n    )\n\n  base_num_iters, rem = divmod(grid[partition_dimension], num_cores)","sourceCodeStart":1594,"sourceCodeEnd":1630,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/pipeline.py#L1594-L1630","documentation":"When partitioning a grid with a dynamic (non-Python-int) dimension, the partitioner can handle at most one dynamic parallel dimension. If two or more parallel axes are dynamic JAX arrays, it cannot statically distribute work, so NotImplementedError is raised.","triggerScenarios":"Passing a grid like (jax_arr_size, jax_arr_size2, 8) where two of the parallel-marked dimensions are JAX Arrays rather than Python ints, together with num_cores/core_id.","commonSituations":"Fully dynamic grids derived from runtime shapes in multi-core TPU kernels; migrating a static-grid kernel to dynamic sizing while keeping core partitioning on.","solutions":["Make all but at most one parallel grid dimension a static Python int","Convert dynamic sizes with int(...) where the value is known at trace time","Disable multi-core partitioning (omit num_cores/core_id) for fully dynamic grids"],"exampleFix":"# before\ngrid=(seq_len_jax, batch_jax, 8)  # two dynamic parallel dims\n# after\ngrid=(int(seq_len), batch_jax, 8)  # one dynamic dim max","handlingStrategy":"validation","validationCode":"sem = dimension_semantics or ()\ndyn = [i for i, (g, d) in enumerate(zip(grid, sem)) if d == 'parallel' and not isinstance(g, int)]\nassert len(dyn) <= 1, f'multiple dynamic parallel dims: {dyn}'","typeGuard":"def at_most_one_dynamic_parallel(grid, sems) -> bool:\n    return sum(1 for g, d in zip(grid, sems) if d == 'parallel' and not isinstance(g, int)) <= 1","tryCatchPattern":null,"preventionTips":["Keep parallel grid dims static Python ints when partitioning across cores","Only one dynamic dim can be parallel"],"tags":["jax","pallas","grid-partition","dynamic-shape"],"backgroundTag":"dynamic-shape-unsupported","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}