如何从数据集目录构建锚-正-负样本对以训练人脸识别单样本学习模型?
Got it, let's solve this anchor-positive-negative sampling problem for your one-shot face recognition task. Since your dataset is organized into per-person directories, we can leverage that structure to build a custom data pipeline that generates valid triplets on the fly. Here's a step-by-step solution tailored to your existing model and Triplet Loss setup:
First, Fix the Model Structure
Your current Sequential model is designed for single-input tasks, but triplet training requires three inputs sharing the same encoder weights. Let's refactor it using Keras' functional API to support this:
from tensorflow.keras import Input, Model def build_encoder(): # Reuse your existing ConvNet architecture model = models.Sequential() model.add(layers.Conv2D(16, (3,3), (3,3), activation='relu', input_shape=(384, 384, 1))) model.add(layers.MaxPooling2D((2,2))) model.add(layers.BatchNormalization()) for t in range(2): model.add(layers.Conv2D(32, (1,1), (1,1), activation='relu')) model.add(layers.Conv2D(32, (3,3), (1,1), padding='same', activation='relu')) model.add(layers.Conv2D(64, (1,1), (1,1), activation='relu')) model.add(layers.BatchNormalization()) model.add(layers.MaxPooling2D((2,2))) for t in range(3): model.add(layers.Conv2D(64, (1,1), (1,1), activation='relu')) model.add(layers.Conv2D(64, (3,3), (1,1), padding='same', activation='relu')) model.add(layers.Conv2D(128, (1,1), (1,1), activation='relu')) model.add(layers.BatchNormalization()) model.add(layers.MaxPooling2D((2,2))) for t in range(4): model.add(layers.Conv2D(128, (1,1), (1,1), activation='relu')) model.add(layers.Conv2D(128, (3,3), (1,1), padding='same', activation='relu')) model.add(layers.Conv2D(256, (1,1), (1,1), activation='relu')) model.add(layers.BatchNormalization()) model.add(layers.MaxPooling2D((2,2))) for t in range(3): model.add(layers.Conv2D(256, (1,1), (1,1), activation='relu')) model.add(layers.Conv2D(256, (3,3), (1,1), padding='same', activation='relu')) model.add(layers.Conv2D(512, (1,1), (1,1), activation='relu')) model.add(layers.BatchNormalization()) model.add(layers.AveragePooling2D((4,4))) model.add(layers.Flatten()) model.add(layers.Dense(128)) model.add(layers.Lambda(lambda x: backend.l2_normalize(x,axis=1))) return model # Build shared encoder encoder = build_encoder() # Define triplet inputs anchor_input = Input(shape=(384, 384, 1), name='anchor') positive_input = Input(shape=(384, 384, 1), name='positive') negative_input = Input(shape=(384, 384, 1), name='negative') # Get encodings for all three inputs (shared weights!) anchor_enc = encoder(anchor_input) positive_enc = encoder(positive_input) negative_enc = encoder(negative_input) # Create triplet model triplet_model = Model( inputs=[anchor_input, positive_input, negative_input], outputs=[anchor_enc, positive_enc, negative_enc] ) # Compile with your existing Triplet Loss triplet_model.compile(optimizer='adam', loss=triplet_loss)
Step 2: Build a Triplet Data Generator
We'll create a generator that pulls images from your directory structure, creates valid triplets, and feeds them to the model. The logic is:
- Pick a random person as the anchor subject
- Grab two images from this person (anchor + positive)
- Grab one image from a different person (negative)
Here's the code:
import os import random import cv2 import numpy as np def load_preprocess_image(img_path): # Load grayscale image, resize to match model input img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, (384, 384)) # Add channel dimension and normalize img = np.expand_dims(img, axis=-1) img = img / 255.0 # Adjust normalization if needed return img def get_valid_person_map(dataset_dir): # Map each person to their image paths, filter out people with <2 images person_to_imgs = {} for person_name in os.listdir(dataset_dir): person_dir = os.path.join(dataset_dir, person_name) if not os.path.isdir(person_dir): continue img_paths = [ os.path.join(person_dir, f) for f in os.listdir(person_dir) if f.lower().endswith(('.jpg', '.png', '.jpeg')) ] if len(img_paths) >= 2: person_to_imgs[person_name] = img_paths return person_to_imgs def triplet_generator(dataset_dir, batch_size=8): person_map = get_valid_person_map(dataset_dir) person_list = list(person_map.keys()) if not person_list: raise ValueError("No person in the dataset has at least 2 images!") while True: anchors, positives, negatives = [], [], [] for _ in range(batch_size): # Pick anchor person and two of their images anchor_person = random.choice(person_list) anchor_img, positive_img = random.sample(person_map[anchor_person], 2) # Pick negative person (different from anchor) and one image negative_person = random.choice([p for p in person_list if p != anchor_person]) negative_img = random.choice(person_map[negative_person]) # Preprocess images anchors.append(load_preprocess_image(anchor_img)) positives.append(load_preprocess_image(positive_img)) negatives.append(load_preprocess_image(negative_img)) # Convert to numpy arrays anchors = np.array(anchors) positives = np.array(positives) negatives = np.array(negatives) # Triplet Loss doesn't use y_true, so we pass dummy labels yield [anchors, positives, negatives], np.zeros((batch_size,))
Step 3: Train the Model
Now you can use the generator to train your triplet model:
# Replace with your dataset path dataset_dir = "/path/to/your/dataset" batch_size = 8 # Adjust based on your GPU memory (384x384 images are large!) steps_per_epoch = 100 # Number of batches per epoch train_gen = triplet_generator(dataset_dir, batch_size=batch_size) triplet_model.fit(train_gen, steps_per_epoch=steps_per_epoch, epochs=50)
Key Notes & Improvements
- Online Hard Negative Mining: For better training, you can modify the generator to select "hard" negatives (negatives whose encodings are closest to the anchor). This speeds up convergence significantly.
- Data Augmentation: Add random flips, rotations, or brightness adjustments to the
load_preprocess_imagefunction to improve generalization. - Validation: Create a separate validation generator using a held-out subset of your dataset to monitor overfitting.
内容的提问来源于stack exchange,提问作者kapil sarwat

