在DenseNet201中添加DeformConv2D后出现维度不匹配错误求解决
问题
希望构建修改版DenseNet201模型,在conv1/conv层后添加DeformConv2D(可变形卷积)层,其余层保持不变。编写的代码如下:
import tensorflow as tf from tensorflow.keras.applications import DenseNet201 from tensorflow.keras.layers import Layer, Input from tensorflow.keras.models import Model class DeformConv2D(Layer): # Placeholder for the actual implementation def __init__(self, out_channels, kernel_size, stride=1, padding='same', **kwargs): super(DeformConv2D, self).__init__(**kwargs) # Initialize your DeformConv2D layer here def call(self, inputs): # Implement the forward pass of your DeformConv2D layer here return inputs # Placeholder return class ModifiedDenseNet201(Model): def __init__(self, **kwargs): super(ModifiedDenseNet201, self).__init__(**kwargs) self.base_model = DenseNet201(include_top=False, input_shape=(224, 224, 3)) self.deform_conv2d = DeformConv2D(out_channels=64, kernel_size=(3, 3), stride=1, padding='same') # Freeze the layers of the base model for layer in self.base_model.layers: layer.trainable = False def call(self, inputs): x = inputs for layer in self.base_model.layers: if layer.name == 'conv1/conv': x = layer(x) x = self.deform_conv2d(x) # Insert the DeformConv2D layer after 'conv1/conv' else: x = layer(x) return x # Create an instance of the modified model input_tensor = Input(shape=(224, 224, 3)) modified_densenet201 = ModifiedDenseNet201() output = modified_densenet201(input_tensor) # Create the modified model model = Model(inputs=input_tensor, outputs=output) # Summary to verify the model structure model.summary()
运行时出现错误:
ValueError: Exception encountered when calling layer 'conv2_block1_concat' (type Concatenate). A merge layer should be called on a list of inputs. Received: inputs=Tensor("modified_dense_net201_1/conv2_block1_2_conv/Conv2D:0", shape=(None, 56, 56, 32), dtype=float32) (not a list of tensors) Call arguments received by layer 'conv2_block1_concat' (type Concatenate): • inputs=tf.Tensor(shape=(None, 56, 56, 32), dtype=float32) Call arguments received by layer "modified_dense_net201_1" (type ModifiedDenseNet201): • inputs=tf.Tensor(shape=(None, 224, 224, 3), dtype=float32)
错误显示concat层输入格式异常,需修改代码解决问题,实现预期模型架构。
解决方案
错误根源是直接遍历DenseNet的层并逐个调用时,破坏了原模型中Concatenate层依赖的多输入分支结构——原DenseNet的concat层接收来自多个分支的张量列表,但逐个调用层时无法保留这种分支输入关系,导致concat层只收到单个张量而非列表。
需修改以下核心部分:
- 放弃遍历所有层的方式,改为截取原模型到
conv1/conv的输出,插入DeformConv2D后,再接上原模型剩余部分 - 确保DeformConv2D的输出通道数与
conv1/conv一致,避免后续层维度不匹配 - 特殊处理Concatenate层,手动维护其输入张量列表,替换为经过可变形卷积后的张量
修改后的完整代码:
import tensorflow as tf from tensorflow.keras.applications import DenseNet201 from tensorflow.keras.layers import Layer, Input from tensorflow.keras.models import Model class DeformConv2D(Layer): # 可变形卷积占位实现,后续替换为实际逻辑 def __init__(self, out_channels, kernel_size, stride=1, padding='same', **kwargs): super(DeformConv2D, self).__init__(**kwargs) # 这里可添加实际可变形卷积的参数初始化,比如偏移量预测卷积 self.conv = tf.keras.layers.Conv2D(out_channels, kernel_size, stride=stride, padding=padding) def call(self, inputs): # 示例实现,替换为实际可变形卷积前向逻辑 return self.conv(inputs) def build_modified_densenet201(): # 加载基础DenseNet201,不包含顶层 base_model = DenseNet201(include_top=False, input_shape=(224, 224, 3)) # 冻结基础模型参数 for layer in base_model.layers: layer.trainable = False # 获取conv1/conv层的输出张量 conv1_output = base_model.get_layer('conv1/conv').output # 插入DeformConv2D层,输出通道数与conv1/conv保持一致(DenseNet201的conv1/conv输出为64通道) deform_output = DeformConv2D(out_channels=64, kernel_size=(3, 3), stride=1, padding='same')(conv1_output) # 从conv1/conv的下一层(conv1/bn)开始,重构后续模型结构 x = deform_output start_process = False for layer in base_model.layers: if layer.name == 'conv1/bn': start_process = True x = layer(x) continue if not start_process: continue # 特殊处理Concatenate层:替换原输入中来自conv1/conv的张量为deform_output if isinstance(layer, tf.keras.layers.Concatenate): concat_inputs = [] for node in layer._inbound_nodes: for input_tensor in node.input_tensors: if input_tensor == conv1_output: concat_inputs.append(deform_output) else: concat_inputs.append(input_tensor) x = layer(concat_inputs) else: x = layer(x) # 构建完整模型 model = Model(inputs=base_model.input, outputs=x) return model # 创建修改后的模型并打印结构 modified_model = build_modified_densenet201() modified_model.summary()
关键修改说明
- 模型拼接逻辑优化:通过截取原模型中间层输出、插入自定义层后再接后续层的方式,保留原模型的分支结构
- Concatenate层适配:手动维护concat层的输入列表,将原本来自
conv1/conv的输入替换为经过可变形卷积后的张量,确保concat层接收正确的输入格式 - 通道数一致性:DeformConv2D输出通道数设为64,与原
conv1/conv输出通道数匹配,避免后续层维度不兼容
内容的提问来源于stack exchange,提问作者shel coop
相关产品推荐
相关产品推荐

