{"record":{"id":"7621d2732b11f98c","repo":"jax-ml/jax","slug":"no-4-dimensional-dimension-number-defaults","errorCode":null,"errorMessage":"No 4+ dimensional dimension_number defaults.","messagePattern":"No 4\\+ dimensional dimension_number defaults\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/convolution.py","lineNumber":367,"sourceCode":"  Returns:\n    Transposed N-d convolution, with output padding following the conventions of\n    keras.layers.Conv2DTranspose.\n  \"\"\"\n  assert len(lhs.shape) == len(rhs.shape) and len(lhs.shape) >= 2\n  ndims = len(lhs.shape)\n  one = (1,) * (ndims - 2)\n  # Set dimensional layout defaults if not specified.\n  if dimension_numbers is None:\n    if ndims == 2:\n      dimension_numbers = ('NC', 'IO', 'NC')\n    elif ndims == 3:\n      dimension_numbers = ('NHC', 'HIO', 'NHC')\n    elif ndims == 4:\n      dimension_numbers = ('NHWC', 'HWIO', 'NHWC')\n    elif ndims == 5:\n      dimension_numbers = ('NHWDC', 'HWDIO', 'NHWDC')\n    else:\n      raise ValueError('No 4+ dimensional dimension_number defaults.')\n  dn = conv_dimension_numbers(lhs.shape, rhs.shape, dimension_numbers)\n  k_shape = np.take(rhs.shape, dn.rhs_spec)\n  k_sdims = k_shape[2:]\n  # Calculate correct output shape given padding and strides.\n  if rhs_dilation is None:\n    rhs_dilation = (1,) * (rhs.ndim - 2)\n  pads: str | Sequence[tuple[int, int]]\n  if use_consistent_padding or (isinstance(padding, str) and padding in {'SAME', 'VALID'}):\n    effective_k_size = map(lambda k, r: core.dilate_dim(k, r), k_sdims, rhs_dilation)\n    replicated_padding = [padding] * len(strides) if isinstance(padding, str) else padding\n    pads = tuple(_conv_transpose_padding(k, s, p)\n      for k,s,p in zip(effective_k_size, strides, replicated_padding))\n  else:\n      pads = padding\n  if transpose_kernel:\n    # flip spatial dims and swap input / output channel axes\n    rhs = _flip_axes(rhs, np.array(dn.rhs_spec)[2:])\n    rhs = rhs.swapaxes(dn.rhs_spec[0], dn.rhs_spec[1])","sourceCodeStart":349,"sourceCodeEnd":385,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/convolution.py#L349-L385","documentation":"conv_transpose can infer default dimension_numbers only for 1-D, 2-D (NHC), 3-D spatial (NHWC/HWIO style) and 4-D spatial (NHWDC) inputs. For operands with more dimensions there is no letter-scheme default, so the else-branch raises ValueError 'No 4+ dimensional dimension_number defaults.'","triggerScenarios":"Calling lax.conv_transpose on arrays with ndim > 5 (more than 4 spatial dims), e.g. 6-D (video-plus or 4D spatial + extra axis) without passing dimension_numbers explicitly.","commonSituations":"Volumetric/time-series models with extra axes; accidentally left-in extra singleton dimensions pushing rank past 5; 4-D spatial climate/medical data.","solutions":["Supply explicit dimension_numbers (a tuple of three layout strings or lax.ConvDimensionNumbers)","Squeeze unnecessary singleton dims to reduce rank to <= 5","Use conv_general_dilated which also needs explicit dnums but gives full control"],"exampleFix":"// before\nlax.conv_transpose(x6d, k6d, strides=(2,2,2,2))  # ndim == 6\n// after\ndn = lax.ConvDimensionNumbers(lhs_spec=(0,5,1,2,3,4), rhs_spec=(5,4,0,1,2,3), out_spec=(0,5,1,2,3,4))\nlax.conv_transpose(x6d, k6d, strides=(2,2,2,2), dimension_numbers=dn)","handlingStrategy":"validation","validationCode":"if lhs.ndim > 5 and dimension_numbers is None:\n    raise ValueError('provide dimension_numbers for rank > 5')","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Squeeze singleton axes before convs","Prepare explicit ConvDimensionNumbers templates for high-rank data"],"tags":["jax","lax","conv-transpose","dimension-numbers","high-rank"],"backgroundTag":"unsupported-input-rank","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}