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.