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

TensorFlow多输入模型数据预处理问题求助

TensorFlow双输入数据预处理问题修正方案

一、数据读取函数的错误修正

你的read_image_data函数存在拼写错误、路径不全、未转numpy数组等问题,修正后如下:

import os, cv2
import numpy as np

train_path = r"..\Dataset\Train"
val_path = r"..\Dataset\Val"

def read_image_data(root):
    all_files = os.listdir(root)  # 修复变量名空格问题
    data1 = []
    data2 = []
    for folder_name in all_files:
        folder_full_path = os.path.join(root, folder_name)
        image_files = os.listdir(folder_full_path)
        # 假设每个文件夹下有两张图片,按顺序取第一张和第二张
        img1_path = os.path.join(folder_full_path, image_files[0])
        img2_path = os.path.join(folder_full_path, image_files[1])
        
        # 读取图片并转RGB(cv2默认BGR),resize到目标尺寸
        img1 = cv2.resize(cv2.cvtColor(cv2.imread(img1_path), cv2.COLOR_BGR2RGB), (320, 320))
        img2 = cv2.resize(cv2.cvtColor(cv2.imread(img2_path), cv2.COLOR_BGR2RGB), (320, 320))
        
        data1.append(img1)
        data2.append(img2)
    # 转换为numpy数组,ImageDataGenerator需要数组格式
    return np.array(data1), np.array(data2)

train_data1, train_data2 = read_image_data(train_path)
val_data1, val_data2 = read_image_data(val_path)

# 标签创建保持不变
train_labels = np.ones(len(train_data1), dtype=int)
val_labels = np.ones(len(val_data1), dtype=int)

二、双输入生成器的错误修正

你的生成器函数存在语法错误、拼写错误、输出格式错误,修正后如下:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 修正rescale参数(1/255.0而非1/.255),可添加其他增强参数
datagen = ImageDataGenerator(
    rescale=1.0/255.0,
    # 可按需添加数据增强,比如:
    # rotation_range=15,
    # width_shift_range=0.1,
    # height_shift_range=0.1,
    # horizontal_flip=True
)

def dual_input_generator(X1, X2, y, batch_size=32):
    # 两个生成器使用相同seed,保证增强同步
    gen_X1 = datagen.flow(X1, y, batch_size=batch_size, seed=7)
    gen_X2 = datagen.flow(X2, batch_size=batch_size, seed=7)
    
    while True:
        X1_batch, y_batch = gen_X1.next()
        X2_batch = gen_X2.next()
        # 返回双输入列表和标签
        yield [X1_batch, X2_batch], y_batch

# 测试生成器输出
temp = next(dual_input_generator(train_data1, train_data2, train_labels))
print("输入1形状:", temp[0][0].shape)
print("输入2形状:", temp[0][1].shape)
print("标签形状:", temp[1].shape)

三、关键注意事项

  • 数据格式一致性:确保两个输入的图片尺寸、通道数完全一致,否则模型无法接收。
  • 同步增强:两个生成器必须使用相同的seed,否则输入1和输入2的增强会错位,破坏样本配对关系。
  • 通道转换:cv2读取的图片是BGR格式,必须转换为RGB,否则和TensorFlow默认的颜色通道顺序不符,影响模型训练效果。
  • 批量大小:根据GPU显存调整batch_size,避免显存溢出。

四、模型训练时的使用示例

from tensorflow.keras.models import Model
from tensorflow.keras.layers import Input, Conv2D, Flatten, Dense, concatenate

# 定义双输入模型
input1 = Input(shape=(320, 320, 3))
x1 = Conv2D(32, (3,3), activation='relu')(input1)
x1 = Flatten()(x1)

input2 = Input(shape=(320, 320, 3))
x2 = Conv2D(32, (3,3), activation='relu')(input2)
x2 = Flatten()(x2)

concat = concatenate([x1, x2])
output = Dense(1, activation='sigmoid')(concat)

model = Model(inputs=[input1, input2], outputs=output)
model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])

# 使用生成器训练
train_generator = dual_input_generator(train_data1, train_data2, train_labels, batch_size=32)
val_generator = dual_input_generator(val_data1, val_data2, val_labels, batch_size=32)

model.fit(
    train_generator,
    steps_per_epoch=len(train_data1)//32,
    validation_data=val_generator,
    validation_steps=len(val_data1)//32,
    epochs=10
)

内容的提问来源于stack exchange,提问作者Tờ Ân

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 23:14:58