Keras Dense层报错:输入最后维度未定义,输入形状为(None, None)
Keras孪生网络子类化Model报错问题解决
问题代码
class modelSIAMESE(keras.Model): def __init__(self): super().__init__() input_shape = target_shape + (3,) input_a = Input(shape = input_shape) input_b = Input(shape = input_shape) input_activation = 'relu' hidden_activation = 'relu' output_activation = 'sigmoid' self.conv1 = Conv2D(128, kernel_size = (3, 3), activation=input_activation, input_shape=input_shape) self.pool1 = MaxPooling2D(pool_size = (2, 2)) self.conv2 = Conv2D(128, kernel_size = (3, 3), activation=hidden_activation) self.pool2 = MaxPooling2D(pool_size = (2, 2)) self.conv3 = Conv2D(128, kernel_size = (3, 3), activation=hidden_activation) self.pool3 = MaxPooling2D(pool_size = (2, 2)) self.flatten = Flatten() self.dense1 = Dense(128, activation="relu", input_shape=(None,)) # Procesar las imágenes de entrada a través de las capas compartidas self.norm_a = layers.BatchNormalization() #(self.flatten(self.pool3(self.conv3(self.pool2(self.conv2(self.pool1(self.conv1(input_a)))))))) self.norm_b = layers.BatchNormalization() #(self.flatten(self.pool3(self.conv3(self.pool2(self.conv2(self.pool1(self.conv1(input_b)))))))) # Definir la capa de comparación self.distance = keras.layers.Subtract() self.prediction = Dense(units=1, activation=output_activation, input_shape=(None,None)) self.dropout = keras.layers.Dropout(0.5) def call(self, inputs, training=False): input_a = inputs["input_a"] input_b = inputs["input_b"] x = self.conv1(input_a) y = self.conv1(input_b) x = self.pool1(x) y = self.pool1(y) x = self.conv2(x) y = self.conv2(y) x = self.pool2(x) y = self.pool2(y) x = self.conv3(x) y = self.conv3(y) x = self.pool3(x) y = self.pool3(y) x = self.dense1(x) y = self.dense1(y) x = self.norm_a(x) y = self.norm_b(y) x = self.flatten(x) y = self.flatten(y) d = self.distance([x, y]) d = K.abs(d) d = self.dropout(d, training=training) return self.prediction(d)
错误分析与修改方案
核心错误点
- 张量维度顺序错误:
call方法中先调用Dense再Flatten,但pool3输出是4D张量((batch_size, height, width, channels)),Dense层只能处理2D张量((batch_size, features)),直接对接会导致维度不匹配。 Dense层input_shape参数冗余且错误:子类化keras.Model时无需手动指定input_shape,框架会自动推导;self.prediction的input_shape=(None,None)完全不符合张量实际维度。- 孪生网络参数未共享:
self.norm_a和self.norm_b是两个独立的BatchNorm层,违背孪生网络共享特征提取参数的设计原则。 - 冗余Input定义:
__init__中定义的input_a、input_b属于函数式API用法,在子类化Model中完全多余,会造成混淆。
修改后的完整代码
import keras from keras import layers from keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Input import keras.backend as K class modelSIAMESE(keras.Model): def __init__(self, target_shape): super().__init__() input_shape = target_shape + (3,) input_activation = 'relu' hidden_activation = 'relu' output_activation = 'sigmoid' # 共享特征提取网络 self.conv1 = Conv2D(128, kernel_size=(3, 3), activation=input_activation, input_shape=input_shape) self.pool1 = MaxPooling2D(pool_size=(2, 2)) self.conv2 = Conv2D(128, kernel_size=(3, 3), activation=hidden_activation) self.pool2 = MaxPooling2D(pool_size=(2, 2)) self.conv3 = Conv2D(128, kernel_size=(3, 3), activation=hidden_activation) self.pool3 = MaxPooling2D(pool_size=(2, 2)) self.flatten = Flatten() self.dense1 = Dense(128, activation="relu") self.norm = layers.BatchNormalization() # 共享BatchNorm # 比较层 self.distance = keras.layers.Subtract() self.dropout = keras.layers.Dropout(0.5) self.prediction = Dense(units=1, activation=output_activation) def call(self, inputs, training=False): input_a = inputs["input_a"] input_b = inputs["input_b"] # 处理输入A x = self.conv1(input_a) x = self.pool1(x) x = self.conv2(x) x = self.pool2(x) x = self.conv3(x) x = self.pool3(x) x = self.flatten(x) # 先Flatten转为2D张量 x = self.dense1(x) x = self.norm(x, training=training) # 处理输入B(共享同一套特征提取层) y = self.conv1(input_b) y = self.pool1(y) y = self.conv2(y) y = self.pool2(y) y = self.conv3(y) y = self.pool3(y) y = self.flatten(y) y = self.dense1(y) y = self.norm(y, training=training) # 计算距离并输出预测 d = self.distance([x, y]) d = K.abs(d) d = self.dropout(d, training=training) return self.prediction(d)
关键修改说明
- 调整张量处理顺序:将
Flatten移至Dense之前,确保输入到Dense的是2D张量。 - 共享BatchNorm层:删除
norm_a和norm_b,改用单个norm层,保证孪生网络两侧参数完全共享。 - 移除冗余参数与定义:删除
__init__中多余的Input定义,以及Dense层错误的input_shape参数。 - 将target_shape作为构造参数传入:避免依赖全局变量,提升代码可复用性。
内容的提问来源于stack exchange,提问作者Alexander Pascau
相关产品推荐
相关产品推荐

