如何编写适配所有Keras模型的通用Xpredict函数以获取预测类名?
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
- Model-Agnostic Input Shape: Uses
self.input_shape[1:3]instead of hardcodingSequential.input_shape, so it works for both Sequential and Functional API models. - Handles All Classification Types: Automatically detects multi-class vs binary classification to get the correct predicted index.
- 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_indicesif available (for models trained withImageDataGenerator).
- Lets you manually pass
- 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_indicesto 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_indicesbased on your training data’s label order. For example, if your class order is["cat", "dog", "bird"], createclass_indices = {"cat":0, "dog":1, "bird":2}.
内容的提问来源于stack exchange,提问作者rella sirisha

