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
- 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
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
- Squeeze singleton axes before convs
- Prepare explicit ConvDimensionNumbers templates for high-rank data
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
- Invalid padding mode: {padding}
- {name} in {op_name} op must be sorted; got {dims}
- {name1} and {name2} in {op_name} op must be disjoint; got: {
- lax.platform_dependent: the '{pname}' branch must be a calla
- Use 'cuda', 'rocm', or 'oneapi' for lax.platform_dependent.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/7621d2732b11f98c.
Report an issue: GitHub.