TensorFlow 2.0中MaxPool2D操作后通道数减少原因及Xception模块通道不匹配问题排查
解决Xception Entry Flow模块的通道不匹配问题
首先咱们先揪出问题的核心:你在call方法里的主分支循环写法出错了!
看你这段代码:
Z = inputs for layer in self.main_layer: Z_1 = layer(Z)
这里每次循环都是把**初始的Z(也就是模块的输入)**直接传入每个层,而不是把前一层的输出作为后一层的输入。也就是说,三个层是并行处理输入,而非串联执行:
- 第一个
SeparableConv2D处理输入得到(?,56,56,128),但这个结果没有被传给下一层; - 第二个
SeparableConv2D还是处理原始输入,同样得到(?,56,56,128),也没往下传递; - 最后
MaxPool2D处理的是原始输入(通道数是前面Conv2D(64)的输出,也就是64),MaxPool2D不会改变通道数,所以最终主分支输出通道是64——这就是你看到主分支输出通道为64的原因!
而你的跳过分支用Conv2D(128)把输入通道从64转成了128,自然和主分支的64通道没法相加,导致报错。
另外要告诉你:你关于SeparableConv2D的filters参数的假设是完全正确的,这个参数确实决定了输出通道数,问题根本不在这个参数上。
修正后的代码
我们需要修改call方法,让主分支的层链式执行,把前一层的输出传给下一层:
import tensorflow as tf from tensorflow import keras class Xception_entry_flow(keras.layers.Layer): def __init__(self,filters_1,activation,size_1,**kwargs): super().__init__(**kwargs) self.main_layer = [ keras.layers.SeparableConv2D(filters = filters_1,kernel_size = size_1,activation = 'relu',padding='same'), keras.layers.SeparableConv2D(filters = filters_1,kernel_size = size_1,padding='same'), # 原始Xception中这一层不提前加激活 keras.layers.MaxPool2D(pool_size= 3,strides=2,padding='same') ] self.skip_layer = [ keras.layers.Conv2D(filters = filters_1,kernel_size=1,strides=2,padding='same') ] self.activation = keras.activations.get(activation) def call(self,inputs): Z = inputs # 主分支链式执行,逐层传递输出 for layer in self.main_layer: Z = layer(Z) print('Shape of Image after main layer step:') print(Z.shape) # 跳过分支处理原始输入 Z_skip = inputs for layer in self.skip_layer: Z_skip = layer(Z_skip) print('Shape after skip layer:') print(Z_skip.shape) # 残差相加后再激活 return self.activation(Z + Z_skip) model = keras.models.Sequential() model.add(keras.layers.Conv2D(32,kernel_size=3,strides = 2,activation='relu',input_shape =[229,229,3])) model.add(keras.layers.Conv2D(64,kernel_size=3,strides=1,activation='relu')) model.add(Xception_entry_flow(filters_1=128,size_1=3,activation ='relu'))
额外贴合原始Xception架构的提示
原始Xception的Entry Flow模块里,第二个SeparableConv2D是不提前加激活函数的,而是先和残差分支的输出相加,之后再一起激活。我在修正代码里已经调整了这一点,如果你严格遵循原始架构,可以保留这个修改。
现在运行修正后的代码,主分支的输出通道会变成128,和跳过分支的输出通道一致,就能正常完成残差相加了。
内容的提问来源于stack exchange,提问作者adhok
相关产品推荐
相关产品推荐

