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

如何编写适配所有Keras模型的通用Xpredict函数以获取预测类名?

Fixing Your xpredict Function to Work with Any Keras Model

Got it, let's sort out this issue so your wrapper function can reliably return prediction class names for any Keras model—whether it was trained with an ImageDataGenerator or raw numpy arrays.

The Root of the Problem

Your current code relies on Sequential.generator.class_indices, but this only exists for generators created with ImageDataGenerator.flow_from_directory() (or similar methods that track class mappings). When you train a model directly on numpy arrays (using model.fit(x_train, y_train)), there's no attached generator with a class_indices attribute—hence the AttributeError when using NumpyArrayIterator.

Updated Universal xpredict Function

Here's a revised version that handles both scenarios, with flexible options to get class names:

def xpredict(self, img_path, batch_size=None, verbose=0, steps=None, class_indices=None):
    # Load and preprocess the image (fixed target_size to use the model's actual input shape)
    target_size = self.input_shape[1:3]  # Grab height/width from model's input shape
    x = image.load_img(img_path, target_size=target_size)
    x = image.img_to_array(x)
    x = np.expand_dims(x, axis=0)
    
    # Get model predictions
    result = self.predict(x, batch_size=batch_size, verbose=verbose, steps=steps)
    
    # Determine the predicted class index (handles multi-class and binary cases)
    if result.shape[1] > 1:
        # Multi-class classification: take index of highest probability
        pred_idx = np.argmax(result, axis=1)[0]
    else:
        # Binary classification: use 0.5 as default threshold
        pred_idx = 1 if result[0][0] > 0.5 else 0
    
    # Resolve class name mapping
    if class_indices is not None:
        # Prioritize manually passed class indices (works for all model types)
        class_map = {v: k for k, v in class_indices.items()}
    else:
        # Try to pull class indices from the model's attached generator (if it exists)
        try:
            class_map = {v: k for k, v in self.generator.class_indices.items()}
        except AttributeError:
            # Fallback: alert user to provide class indices manually
            raise ValueError(
                "Could not automatically find class indices. "
                "Please pass the `class_indices` parameter (e.g., {'cat': 0, 'dog': 1}) "
                "since this model was not trained with a generator that tracks class mappings."
            )
    
    # Return the matching class name (or a fallback if index isn't found)
    return class_map.get(pred_idx, f"Unknown class index: {pred_idx}")

Key Improvements Explained

  1. Model-Agnostic Input Shape: Uses self.input_shape[1:3] instead of hardcoding Sequential.input_shape, so it works for both Sequential and Functional API models.
  2. Handles All Classification Types: Automatically detects multi-class vs binary classification to get the correct predicted index.
  3. Flexible Class Mapping:
    • Lets you manually pass class_indices (e.g., {'cat':0, 'dog':1}) for models trained on numpy arrays.
    • Falls back to the generator's class_indices if available (for models trained with ImageDataGenerator).
  4. Clear Error Handling: Tells you exactly what to do if no class mapping can be found automatically.

Extra Tips for Smooth Usage

  • Save Class Indices When Training: If you use a generator, save its class_indices to a file so you can reuse it later:
    import json
    # Save during training
    with open("class_indices.json", "w") as f:
        json.dump(train_generator.class_indices, f)
    
    # Load when using the model
    with open("class_indices.json", "r") as f:
        class_indices = json.load(f)
    
    # Use in xpredict
    model.xpredict("test_img.jpg", class_indices=class_indices)
    
  • For Numpy Array Training: You’ll need to define class_indices based on your training data’s label order. For example, if your class order is ["cat", "dog", "bird"], create class_indices = {"cat":0, "dog":1, "bird":2}.

内容的提问来源于stack exchange,提问作者rella sirisha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:46:02