如何适配Keras siamese_contrastive.py示例加载自定义图像数据集
自定义三元组数据集接入对比损失孪生网络实现方案
核心逻辑说明
你当前的目录结构是标准的三元组存储格式,同文件名的三个文件刚好组成(anchor, positive, negative)样本组,只需要把每个三元组拆成1组正样本对、1组负样本对,就能完全适配原示例对比损失的输入要求;如果需要先提取图像嵌入,直接在数据预处理阶段接入预训练特征提取器即可,不需要改动后续孪生网络的对比损失逻辑。
分步实现
1. 生成三元组路径清单
先扫描三个目录,校验文件配对关系,避免缺漏文件,所有路径按需加载不会提前占用内存:
import os import tensorflow as tf from tensorflow import keras # 按需修改配置 IMG_SHAPE = (28, 28, 3) # 对齐原MNIST输入尺寸,可根据业务调整 BATCH_SIZE = 32 EMBED_DIM = 128 # 提前提取嵌入时使用的向量维度 DATA_ROOT = "./" # 替换成你的数据集根目录 anchor_dir = os.path.join(DATA_ROOT, "anchor") positive_dir = os.path.join(DATA_ROOT, "positive") negative_dir = os.path.join(DATA_ROOT, "negative") # 过滤jpg文件并排序,保证同序号文件一一对应 img_list = sorted([f for f in os.listdir(anchor_dir) if f.lower().endswith(".jpg")]) triplet_path_list = [] for img_name in img_list: p_path = os.path.join(positive_dir, img_name) n_path = os.path.join(negative_dir, img_name) if os.path.exists(p_path) and os.path.exists(n_path): triplet_path_list.append( (os.path.join(anchor_dir, img_name), p_path, n_path) )
2. 实现预处理逻辑
分两种场景选择对应的预处理函数:
- 直接用原始图像训练:做解码、resize、归一化即可
- 先提取嵌入再训练:预处理阶段接入预训练特征提取器,直接输出固定维度的嵌入向量
# -------------------------- # 场景1:直接输入原始图像训练用这个预处理 def preprocess(file_path): img = tf.io.read_file(file_path) img = tf.io.decode_jpeg(img, channels=3) img = tf.image.resize(img, IMG_SHAPE[:2]) img = tf.cast(img, tf.float32) / 255.0 # 像素值归一化到0~1 return img # -------------------------- # 场景2:提前提取图像嵌入训练用这个预处理 # 先加载你自己的预训练嵌入提取模型,不需要微调的话设trainable=False # pretrained_embed = keras.models.load_model("your_feat_extractor.h5") # pretrained_embed.trainable = False # def preprocess(file_path): # img = tf.io.read_file(file_path) # img = tf.io.decode_jpeg(img, channels=3) # img = tf.image.resize(img, pretrained_embed.input_shape[1:3]) # 如果用ImageNet预训练模型,替换成对应模型的预处理逻辑 # img = tf.keras.applications.resnet50.preprocess_input(img) # img = tf.expand_dims(img, 0) # emb = pretrained_embed(img, training=False) # return tf.squeeze(emb, 0)
3. 构建TensorFlow数据集流水线
把每个三元组拆成正、负两个样本对,生成符合原示例输入格式的数据集,做加载性能优化:
def sample_generator(): for a_path, p_path, n_path in triplet_path_list: a_feat = preprocess(a_path) p_feat = preprocess(p_path) n_feat = preprocess(n_path) # 正样本对 标签=1 yield (a_feat, p_feat), 1 # 负样本对 标签=0 yield (a_feat, n_feat), 0 # 定义输出格式,场景2用嵌入的话把IMG_SHAPE替换成(EMBED_DIM,) ds = tf.data.Dataset.from_generator( sample_generator, output_signature=( (tf.TensorSpec(shape=IMG_SHAPE, dtype=tf.float32), tf.TensorSpec(shape=IMG_SHAPE, dtype=tf.float32)), tf.TensorSpec(shape=(), dtype=tf.int32) ) ) # 拆分训练/验证集,8:2比例可按需调整 train_size = int(len(ds) * 0.8) train_ds = ds.take(train_size).shuffle(1024).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE) val_ds = ds.skip(train_size).batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
4. 对接原示例代码
- 场景1(原始图像输入):直接删除原示例中
keras.datasets.mnist.load_data()相关的加载代码,把train_ds、val_ds传入model.fit()即可,原有的卷积嵌入backbone不需要改动,注意输入通道改成3对应RGB图像。 - 场景2(嵌入输入):把原示例中用于提特征的卷积嵌入网络替换成轻量全连接头即可,参考实现:
def get_embed_head(input_dim=EMBED_DIM): inputs = keras.Input(shape=(input_dim,)) x = keras.layers.Dense(256, activation="relu")(inputs) x = keras.layers.Dense(EMBED_DIM)(x) outputs = keras.layers.UnitNormalization()(x) return keras.Model(inputs, outputs)
注意:如果同文件名的三个文件不是严格一一对应的配对关系,不要用文件名排序生成三元组,建议先维护一份CSV文件记录每一组三元组的路径,再按CSV加载避免样本配错。
内容的提问来源于stack exchange,提问作者Exploring
相关产品推荐
相关产品推荐

