{"record":{"id":"284e32ace1a7bb2f","repo":"jax-ml/jax","slug":"ragged-all-to-all-recv-sizes-must-be-integer-type","errorCode":null,"errorMessage":"ragged_all_to_all recv_sizes must be integer type.","messagePattern":"ragged_all_to_all recv_sizes must be integer type\\.","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/parallel.py","lineNumber":1651,"sourceCode":"              recv_sizes],\n      call_target_name=ir.StringAttr.get('ragged_all_to_all'),\n      backend_config=ir.DictAttr.get(ragged_all_to_all_attrs),\n      api_version=ir.IntegerAttr.get(ir.IntegerType.get_signless(32), 4),\n  ).results\n\ndef _ragged_all_to_all_effectful_abstract_eval(\n    operand, output, input_offsets, send_sizes, output_offsets, recv_sizes,\n    axis_name, axis_index_groups\n):\n  del operand, axis_index_groups\n  if not dtypes.issubdtype(input_offsets.dtype, np.integer):\n    raise ValueError(\"ragged_all_to_all input_offsets must be integer type.\")\n  if not dtypes.issubdtype(send_sizes.dtype, np.integer):\n    raise ValueError(\"ragged_all_to_all send_sizes must be integer type.\")\n  if not dtypes.issubdtype(output_offsets.dtype, np.integer):\n    raise ValueError(\"ragged_all_to_all output_offsets must be integer type.\")\n  if not dtypes.issubdtype(recv_sizes.dtype, np.integer):\n    raise ValueError(\"ragged_all_to_all recv_sizes must be integer type.\")\n  if len(input_offsets.shape) != 1 or input_offsets.shape[0] < 1:\n    raise ValueError(\n        \"ragged_all_to_all input_offsets must be rank 1 with positive dimension\"\n        \" size, but got shape {}\".format(input_offsets.shape)\n    )\n  if len(send_sizes.shape) != 1 or send_sizes.shape[0] < 1:\n    raise ValueError(\n        \"ragged_all_to_all send_sizes must be rank 1 with positive dimension\"\n        \" size, but got shape {}\".format(send_sizes.shape)\n    )\n  if len(output_offsets.shape) != 1 or output_offsets.shape[0] < 1:\n    raise ValueError(\n        \"ragged_all_to_all output_offsets must be rank 1 with positive\"\n        \" dimension size, but got shape {}\".format(output_offsets.shape)\n    )\n  if len(recv_sizes.shape) != 1 or recv_sizes.shape[0] < 1:\n    raise ValueError(\n        \"ragged_all_to_all recv_sizes must be rank 1 with positive dimension\"","sourceCodeStart":1633,"sourceCodeEnd":1669,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/parallel.py#L1633-L1669","documentation":"recv_sizes for ragged_all_to_all must be an integer-dtype array; the abstract eval validates this and raises ValueError for float or other dtypes.","triggerScenarios":"Passing float-typed recv_sizes to lax.ragged_all_to_all.","commonSituations":"Symmetric with send/output offsets: sizes computed as floats from division or averages.","solutions":["Cast recv_sizes to np.int64","Ensure all four arrays (input/output offsets, send/recv sizes) are integer dtype via a pre-call check"],"exampleFix":"// before\nlax.ragged_all_to_all(x, out, offs, send, 'i', recv_sizes=recv.astype(jnp.float32))\n// after\nlax.ragged_all_to_all(x, out, offs, send, 'i', recv_sizes=recv.astype(jnp.int32))","handlingStrategy":"validation","validationCode":"assert np.issubdtype(np.asarray(recv_sizes).dtype, np.integer), 'recv_sizes must be int'","typeGuard":"def is_int_array(a): return np.issubdtype(np.asarray(a).dtype, np.integer)","tryCatchPattern":null,"preventionTips":["Add a single validator covering all four offset/size arrays before calling ragged_all_to_all"],"tags":["jax","ragged-all-to-all","dtype"],"backgroundTag":"dtype-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}