自定义Keras损失函数中K.argmax与K.gather配合触发索引越界错误求助
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 toaxis=-1. For a typical classification task wherey_truehas 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.gatherby default indexes into the 0th axis (the batch axis) of youry_predtensor. 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:
Method 1: Use K.batch_gather (Recommended)
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

