jax-ml/jax · error · ValueError

No 4+ dimensional dimension_number defaults.

Error message

No 4+ dimensional dimension_number defaults.

What it means

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.'

Source

Thrown at jax/_src/lax/convolution.py:367

  Returns:
    Transposed N-d convolution, with output padding following the conventions of
    keras.layers.Conv2DTranspose.
  """
  assert len(lhs.shape) == len(rhs.shape) and len(lhs.shape) >= 2
  ndims = len(lhs.shape)
  one = (1,) * (ndims - 2)
  # Set dimensional layout defaults if not specified.
  if dimension_numbers is None:
    if ndims == 2:
      dimension_numbers = ('NC', 'IO', 'NC')
    elif ndims == 3:
      dimension_numbers = ('NHC', 'HIO', 'NHC')
    elif ndims == 4:
      dimension_numbers = ('NHWC', 'HWIO', 'NHWC')
    elif ndims == 5:
      dimension_numbers = ('NHWDC', 'HWDIO', 'NHWDC')
    else:
      raise ValueError('No 4+ dimensional dimension_number defaults.')
  dn = conv_dimension_numbers(lhs.shape, rhs.shape, dimension_numbers)
  k_shape = np.take(rhs.shape, dn.rhs_spec)
  k_sdims = k_shape[2:]
  # Calculate correct output shape given padding and strides.
  if rhs_dilation is None:
    rhs_dilation = (1,) * (rhs.ndim - 2)
  pads: str | Sequence[tuple[int, int]]
  if use_consistent_padding or (isinstance(padding, str) and padding in {'SAME', 'VALID'}):
    effective_k_size = map(lambda k, r: core.dilate_dim(k, r), k_sdims, rhs_dilation)
    replicated_padding = [padding] * len(strides) if isinstance(padding, str) else padding
    pads = tuple(_conv_transpose_padding(k, s, p)
      for k,s,p in zip(effective_k_size, strides, replicated_padding))
  else:
      pads = padding
  if transpose_kernel:
    # flip spatial dims and swap input / output channel axes
    rhs = _flip_axes(rhs, np.array(dn.rhs_spec)[2:])
    rhs = rhs.swapaxes(dn.rhs_spec[0], dn.rhs_spec[1])

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Supply explicit dimension_numbers (a tuple of three layout strings or lax.ConvDimensionNumbers)
  2. Squeeze unnecessary singleton dims to reduce rank to <= 5
  3. Use conv_general_dilated which also needs explicit dnums but gives full control

Example fix

// before
lax.conv_transpose(x6d, k6d, strides=(2,2,2,2))  # ndim == 6
// after
dn = 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))
lax.conv_transpose(x6d, k6d, strides=(2,2,2,2), dimension_numbers=dn)
Defensive patterns

Strategy: validation

Validate before calling

if lhs.ndim > 5 and dimension_numbers is None:
    raise ValueError('provide dimension_numbers for rank > 5')

Prevention

When it happens

Trigger: 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.

Common situations: Volumetric/time-series models with extra axes; accidentally left-in extra singleton dimensions pushing rank past 5; 4-D spatial climate/medical data.

Related errors


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/7621d2732b11f98c. Report an issue: GitHub.