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

如何适配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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 16:15:55