You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义Keras损失函数中K.argmax与K.gather配合触发索引越界错误求助

Fixing Your Custom Keras Loss Function: Index Out of Bounds Error

Alright, let's break down what's causing that InvalidArgumentError and fix your loss function properly.

Why the Error Happens

Your issue boils down to a mismatch between the axis used for K.argmax and the axis K.gather operates on:

  • When you call K.argmax(y_true) without specifying an axis, it defaults to axis=-1. For a typical classification task where y_true has shape (batch_size, num_classes), this returns a tensor of shape (batch_size,), where each value is the index of the true class for that sample (e.g., 51 for a class in a 100-class dataset).
  • But K.gather by default indexes into the 0th axis (the batch axis) of your y_pred tensor. If your batch size is 32, trying to use an index like 51 here will obviously be out of the [0, 32) range—hence the error.

Correct Implementations

Here are two straightforward ways to fix this, depending on which approach you prefer:

K.batch_gather is designed specifically for batch-wise indexing—it lets you pick the correct value from each sample's prediction vector using the true class index.

def lossFunction(self, y_true, y_pred):
    # Get true class indices for each sample (shape: (batch_size,))
    true_class_indices = K.argmax(y_true, axis=-1)
    # Expand dimensions to match batch_gather's requirements (shape: (batch_size, 1))
    true_class_indices = K.expand_dims(true_class_indices, axis=-1)
    # Extract the predicted value for the true class for each sample
    true_class_preds = K.batch_gather(y_pred, true_class_indices)
    # If y_true is one-hot encoded, K.max(y_true) is always 1—you can replace this with 1.0 directly
    max_true = K.max(y_true, axis=-1)
    # Calculate mean squared error between true max value and corresponding prediction
    return K.mean(K.square(max_true - K.squeeze(true_class_preds, axis=-1)))

Method 2: Transpose Tensors to Match K.gather

If you'd rather stick with K.gather, you can transpose your y_pred tensor so the class axis becomes the 0th axis, gather the values, then transpose back:

def lossFunction(self, y_true, y_pred):
    true_class_indices = K.argmax(y_true, axis=-1)
    # Transpose y_pred to (num_classes, batch_size) so we can index into classes
    y_pred_transposed = K.transpose(y_pred)
    # Gather the predicted value for each sample's true class
    true_class_preds = K.gather(y_pred_transposed, true_class_indices)
    max_true = K.max(y_true, axis=-1)
    return K.mean(K.square(max_true - true_class_preds))

Quick Optimization

If your y_true is one-hot encoded (which it likely is, given you're using K.argmax), K.max(y_true) will always equal 1.0. You can simplify the loss function by replacing max_true with a constant to save a computation step:

return K.mean(K.square(1.0 - K.squeeze(true_class_preds, axis=-1)))

内容的提问来源于stack exchange,提问作者Vinayak Mp

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 06:27:28