{"record":{"id":"2f305853adb0b3dc","repo":"jax-ml/jax","slug":"batch-dimension-sizes-must-match-got-sorted-arr","errorCode":null,"errorMessage":"batch dimension sizes must match; got {sorted_arr_aval.shape[:batch_dims]} != {query_aval.shape[:batch_dims]}","messagePattern":"batch dimension sizes must match; got (.+?) != (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/hijax.py","lineNumber":92,"sourceCode":"      raise ValueError(\n          f\"dimension={dimension} must be in range [{batch_dims},\"\n          f\" {sorted_arr_aval.ndim})\"\n      )\n    if sorted_arr_aval.dtype != query_aval.dtype:\n      raise ValueError(\n          \"dtypes of sorted_arr and query must match; got \"\n          f\"{sorted_arr_aval.dtype} and {query_aval.dtype}\"\n      )\n    if side not in [\"left\", \"right\"]:\n      raise ValueError(\n          f\"invalid argument side={side!r}, expected 'left' or 'right'\"\n      )\n    if method not in self.valid_methods:\n      raise ValueError(\n          f\"invalid argument {method=}, expected one of {list(self.valid_methods)}\"\n      )\n    if sorted_arr_aval.shape[:batch_dims] != query_aval.shape[:batch_dims]:\n      raise ValueError(\n          \"batch dimension sizes must match; got\"\n          f\" {sorted_arr_aval.shape[:batch_dims]} != {query_aval.shape[:batch_dims]}\"\n      )\n    if not dtypes.issubdtype(out_dtype, np.integer):\n      raise ValueError(f\"out_dtype should be an integer type; got {out_dtype}\")\n    # Attempt this here to catch overflow errors early.\n    out_dtype.type(sorted_arr_aval.shape[dimension])\n    self.in_avals = (sorted_arr_aval, query_aval)\n    self.out_aval = core.typeof(api.eval_shape(\n      functools.partial(_searchsorted_impl,\n        dimension=dimension, batch_dims=batch_dims, side=side,\n        dtype=out_dtype, method=method),\n        sorted_arr_aval, query_aval))\n    self.params = dict(\n      side=side,\n      dimension=dimension,\n      batch_dims=batch_dims,\n      method=method,","sourceCodeStart":74,"sourceCodeEnd":110,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/hijax.py#L74-L110","documentation":"Raised by the SearchSorted HiJAX primitive when the leading batch dimensions of sorted_arr and query differ. When batch_dims > 0, both arrays must carry identical leading shape[0:batch_dims] so the primitive can vmap over them consistently.","triggerScenarios":"Calling batched searchsorted (batch_dims=k) where, e.g., sorted_arr.shape=(B1, N, M) and query.shape=(B2, M) with B1 != B2; often caused by mismatched leading axes after reshaping or vmap.","commonSituations":"Using vmap/pmap where one operand got an extra or different batch axis; a reshape that merged or split the batch dim in only one operand.","solutions":["Align batch shapes: query = jnp.broadcast_to(query, sorted_arr.shape[:batch_dims] + query.shape[batch_dims:]) or reshape the sorted array","Check which operand gained/lost a batch dim under vmap and use in_axes appropriately","Print sorted_arr.shape[:batch_dims] and query.shape[:batch_dims] right before the call"],"exampleFix":"# before\nout = searchsorted_batched(sorted_arr, query)  # (8,100) vs (16,)\n# after\nquery = jnp.broadcast_to(query[None], (8, query.shape[0]))\nout = searchsorted_batched(sorted_arr, query)","handlingStrategy":"validation","validationCode":"assert sorted_arr.shape[:batch_dims] == query.shape[:batch_dims], (sorted_arr.shape, query.shape)","typeGuard":"def batch_shapes_match(a, q, k: int) -> bool:\n    return a.shape[:k] == q.shape[:k]","tryCatchPattern":null,"preventionTips":["Print both shapes before batched search calls","Use jnp.broadcast_to on the smaller operand when batch dims legitimately differ","Be explicit with vmap in_axes when only one operand is batched"],"tags":["jax","shape-mismatch","searchsorted","batching"],"backgroundTag":"shape-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}