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

基于硬盘图像的TensorFlow自编码器数据集构建问题排查

TensorFlow自编码器训练报错排查

问题核心

报错提示Layer "sequential_2" expects 1 input(s), but it received 2 input tensors,根源有两点:

  • image_dataset_from_directory返回的数据集每个元素是**(图像数组, 标签数组)**的元组,你直接传入batch时,会把标签也当成输入传给模型,但自编码器是无监督任务,只需要图像作为输入(输入=输出)
  • 你的模型输入形状未考虑图像的3通道(默认加载彩色图,shape为(171,256,3)),和实际输入数据维度不匹配

修复步骤

  1. 处理数据集:丢弃标签+归一化
    用map函数转换数据集,只保留图像部分,同时将像素值归一化到0-1区间(适配解码器的sigmoid激活输出):

    data = data.map(lambda x, y: x / 255.0)
    
  2. 修正模型的通道维度
    编码器输入要包含3通道信息,解码器最后输出要对应3通道的图像:

    encoder = Sequential()
    # 修正输入形状:加入3通道
    encoder.add(Flatten(input_shape=[171,256,3]))
    encoder.add(Dense(400,activation='relu'))
    encoder.add(Dense(200,activation='relu'))
    encoder.add(Dense(100,activation='relu'))
    encoder.add(Dense(50,activation='relu'))
    encoder.add(Dense(25,activation='relu'))
    
    decoder = Sequential()
    decoder.add(Dense(50,input_shape=[25],activation='relu'))
    decoder.add(Dense(100,activation='relu'))
    decoder.add(Dense(200,activation='relu'))
    decoder.add(Dense(400,activation='relu'))
    # 修正输出维度:对应3通道图像的总像素数
    decoder.add(Dense(171*256*3,activation='sigmoid'))
    # 修正Reshape:恢复3通道图像形状
    decoder.add(Reshape([171,256,3]))
    
  3. 正确传入训练数据

    • 直接用处理后的数据集训练:
      autoencoder.fit(data, epochs=5)
      
    • 如果用迭代器测试单batch:
      batch = data_iterator.next()
      # 取batch中的图像部分作为输入和目标
      autoencoder.fit(batch, batch, epochs=5)
      

完整修正后代码片段

import os
import pandas as pd
import numpy as np
import seaborn as sns
import matplotlib.pyplot as plt
from matplotlib.image import imread
import matplotlib.image as mpimg
import cv2
%matplotlib inline

from google.colab import drive
drive.mount('/content/gdrive')

my_data_dir = '/content/gdrive/MyDrive/Skyrmion Vision/testFiles/train/'
images = os.listdir(my_data_dir)

# 加载数据集
data = tf.keras.utils.image_dataset_from_directory(
    '/content/gdrive/MyDrive/Skyrmion Vision/testFiles/train/',
    batch_size=1,
    image_size=(171,256)
)
# 处理数据集:丢弃标签+归一化
data = data.map(lambda x, y: x / 255.0)

data_iterator = data.as_numpy_iterator()
batch = data_iterator.next()

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense,Flatten,Reshape
from tensorflow.keras.optimizers import SGD

# 修正后的编码器
encoder = Sequential()
encoder.add(Flatten(input_shape=[171,256,3]))
encoder.add(Dense(400,activation='relu'))
encoder.add(Dense(200,activation='relu'))
encoder.add(Dense(100,activation='relu'))
encoder.add(Dense(50,activation='relu'))
encoder.add(Dense(25,activation='relu'))

# 修正后的解码器
decoder = Sequential()
decoder.add(Dense(50,input_shape=[25],activation='relu'))
decoder.add(Dense(100,activation='relu'))
decoder.add(Dense(200,activation='relu'))
decoder.add(Dense(400,activation='relu'))
decoder.add(Dense(171*256*3,activation='sigmoid'))
decoder.add(Reshape([171,256,3]))

autoencoder = Sequential([encoder,decoder])
autoencoder.compile(loss='binary_crossentropy',optimizer=SGD(learning_rate=1.5),metrics=['accuracy'])
# 用处理后的数据集训练
autoencoder.fit(data, epochs=5)

内容的提问来源于stack exchange,提问作者Joseph

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 18:05:27