如何用Keras构建双分支CNN实现硬币双面图像分类及数据输入?
Great question—you’re spot-on that trying to train a single CNN to handle both coin sides at once isn’t the best approach. By splitting into two specialized branches (one for front, one for back) and merging their features into a shared classifier, you’ll let each branch learn the unique patterns of each side while combining their strengths for better classification. Let’s walk through exactly how to build this in Keras, including data handling.
The setup we’ll build has three key parts:
- Two parallel CNN branches (same structure, either shared or independent weights) to extract features from front and back images respectively.
- A feature fusion step to combine the outputs of the two branches.
- A shared fully-connected classifier that takes the merged features and outputs the coin category.
Let’s start coding this step-by-step.
2.1 Define the Feature Extractor Branch
First, create a reusable function to build the CNN structure for each branch—this ensures both branches have identical architecture:
from tensorflow.keras import layers, Model def build_feature_branch(input_shape): inputs = layers.Input(shape=input_shape) # Customize this CNN stack to fit your image size and data complexity x = layers.Conv2D(32, (3, 3), activation='relu', padding='same')(inputs) x = layers.MaxPooling2D((2, 2))(x) x = layers.Conv2D(64, (3, 3), activation='relu', padding='same')(x) x = layers.MaxPooling2D((2, 2))(x) x = layers.Conv2D(128, (3, 3), activation='relu', padding='same')(x) x = layers.MaxPooling2D((2, 2))(x) x = layers.Flatten()(x) return Model(inputs, x, name="feature_extractor")
2.2 Assemble the Dual-Branch & Merge Features
Now, create inputs for both image types, instantiate the branches, and merge their feature outputs:
# Set your image dimensions (adjust to match your dataset) input_shape = (128, 128, 3) num_classes = 10 # Replace with your number of coin categories # Branch 1: For front-facing coin images front_branch = build_feature_branch(input_shape) front_input = front_branch.input # Branch 2: For back-facing coin images # If you already pre-trained separate models for front/back, load their weights here: # back_branch = build_feature_branch(input_shape) # back_branch.load_weights("path/to/pretrained_back_branch_weights.h5") back_branch = build_feature_branch(input_shape) back_input = back_branch.input # Extract features from both branches front_features = front_branch.output back_features = back_branch.output # Merge features (concatenation is the most common choice; try add/multiply if needed) merged_features = layers.concatenate([front_features, back_features], axis=-1)
2.3 Add the Shared Classifier
Attach the shared fully-connected layers that will learn to classify coins using the merged features:
# Shared classifier head x = layers.Dense(256, activation='relu')(merged_features) x = layers.Dropout(0.5)(x) # Prevent overfitting outputs = layers.Dense(num_classes, activation='softmax')(x) # Build the full model model = Model(inputs=[front_input, back_input], outputs=outputs) # Compile the model model.compile( optimizer='adam', loss='sparse_categorical_crossentropy', # Use 'categorical_crossentropy' if labels are one-hot metrics=['accuracy'] ) # Optional: Print model summary to verify the structure model.summary()
The critical part here is ensuring each sample’s front and back images are paired correctly. Here’s how to handle both small and large datasets:
3.1 For Small Datasets (In-Memory Arrays)
Organize your data into three arrays:
front_images: Shape(num_samples, height, width, channels)back_images: Shape(num_samples, height, width, channels)labels: Shape(num_samples,), with integer labels for each coin category
Train the model by passing a list of the two image arrays as input:
# Example training call (replace with your train/validation splits) model.fit( x=[train_front_images, train_back_images], y=train_labels, validation_data=([val_front_images, val_back_images], val_labels), epochs=20, batch_size=32 )
3.2 For Large Datasets (Using tf.data.Dataset)
If your data is too big to fit in memory, use TensorFlow’s tf.data pipeline to load and pair images on the fly:
import tensorflow as tf # Assume you have lists of file paths: front_paths, back_paths, and corresponding labels dataset = tf.data.Dataset.from_tensor_slices((front_paths, back_paths, labels)) # Define a preprocessing function to load and normalize images def load_and_preprocess(front_path, back_path, label): # Load images front_img = tf.io.read_file(front_path) front_img = tf.image.decode_jpeg(front_img, channels=3) back_img = tf.io.read_file(back_path) back_img = tf.image.decode_jpeg(back_img, channels=3) # Resize and normalize front_img = tf.image.resize(front_img, (128, 128)) / 255.0 back_img = tf.image.resize(back_img, (128, 128)) / 255.0 return (front_img, back_img), label # Apply preprocessing, shuffle, and batch dataset = dataset.map(load_and_preprocess).shuffle(1000).batch(32) # Train the model model.fit(dataset, epochs=20)
If you already trained separate CNNs for front and back images, here’s how to integrate them:
- Load the pre-trained weights into each branch as shown earlier.
- Freeze the feature extraction layers first to avoid destroying pre-trained patterns:
# Freeze all layers in both branches for layer in front_branch.layers: layer.trainable = False for layer in back_branch.layers: layer.trainable = False - Recompile the model and train only the shared classifier.
- Optionally, unfreeze the last few layers of the branches later and fine-tune with a small learning rate to adapt the features to the combined task.
- Feature Fusion: Concatenation works for most cases, but if you suspect front/back features are complementary, try
layers.add()orlayers.multiply()to emphasize shared patterns. - Weight Sharing: If coin front/back share common features (e.g., edge texture, material), you can use a single shared branch instead of two separate ones (just reuse the same
front_branchfor both inputs). This reduces model size and training time. - Input Alignment: Double-check that your front/back image pairs are always correctly matched—misaligned pairs will ruin model performance.
内容的提问来源于stack exchange,提问作者Steve Brown

