keras-team/keras · error · ValueError
`cdist` inputs must have rank >= 2
Error message
`cdist` inputs must have rank >= 2
What it means
Error "`cdist` inputs must have rank >= 2" thrown in keras-team/keras.
Source
Thrown at keras/src/backend/jax/math.py:86
)
# `nan` shouldn't be considered as large probability.
preds_at_label = jnp.where(
jnp.isnan(preds_at_label), -jnp.inf, preds_at_label
)
rank = 1 + jnp.sum(jnp.greater(predictions, preds_at_label), axis=-1)
return jnp.less_equal(rank, k)
def logsumexp(x, axis=None, keepdims=False):
x = convert_to_tensor(x)
return jax.scipy.special.logsumexp(x, axis=axis, keepdims=keepdims)
def cdist(x, y):
x = jnp.asarray(x)
y = jnp.asarray(y)
if x.ndim < 2 or y.ndim < 2:
raise ValueError("`cdist` inputs must have rank >= 2")
if x.shape[-1] != y.shape[-1]:
raise ValueError("Last dimension of inputs to `cdist` must match")
diff = jnp.expand_dims(x, -2) - jnp.expand_dims(y, -3)
return jnp.sqrt(jnp.sum(diff * diff, axis=-1))
def extract_sequences(x, sequence_length, sequence_stride):
*batch_shape, signal_length = x.shape
batch_shape = list(batch_shape)
x = jnp.reshape(x, (math.prod(batch_shape), signal_length, 1))
x = jax.lax.conv_general_dilated_patches(
x,
(sequence_length,),
(sequence_stride,),
"VALID",
dimension_numbers=("NTC", "OIT", "NTC"),
)
return jnp.reshape(x, (*batch_shape, *x.shape[-2:]))View on GitHub (pinned to 7a34a03db6)
When it happens
Trigger: Thrown at keras/src/backend/jax/math.py:86 when the library encounters an invalid state.
Common situations: See trigger scenarios.
AI-assisted analysis of keras-team/keras@7a34a03db6 (2026-08-25).
Data as JSON: /api/errors/222cb19bb7da36a4.
Report an issue: GitHub.