如何从tf.keras.Model中查看自定义Encoder/Decoder层的详细结构?
解决自定义Keras层无法在summary中显示内部结构的方法
方法一:单独查看Encoder/Decoder的详细结构
创建输入张量,将自定义的Encoder或Decoder包装成独立的Model,直接查看其内部结构:
# 替换为你的实际输入形状,示例为(28,28,1) input_sample = K.Input(shape=(28,28,1)) # 查看Encoder详细结构 encoder_model = K.Model(inputs=input_sample, outputs=model.encoder(input_sample)) encoder_model.summary() # 查看Decoder详细结构:先获取Encoder输出形状作为Decoder输入 encoder_output = model.encoder(input_sample) decoder_model = K.Model(inputs=encoder_output, outputs=model.decoder(encoder_output)) decoder_model.summary()
方法二:修改自定义Layer实现,用Sequential封装内部层
将Encoder和Decoder的内部层用tf.keras.Sequential封装,Keras会自动在summary中展开内部结构:
class Encoder(K.layers.Layer): def __init__(self, filters): super(Encoder, self).__init__() self.encoder_layers = K.Sequential([ Conv2D(filters=filters[0], kernel_size=3, strides=1, activation='relu', padding='same'), MaxPooling2D((2, 2), padding='same'), Conv2D(filters=filters[1], kernel_size=3, strides=1, activation='relu', padding='same'), MaxPooling2D((2, 2), padding='same'), Conv2D(filters=filters[2], kernel_size=3, strides=1, activation='relu', padding='same'), MaxPooling2D((2, 2), padding='same') ]) def call(self, input_features): return self.encoder_layers(input_features) class Decoder(K.layers.Layer): def __init__(self, filters): super(Decoder, self).__init__() self.decoder_layers = K.Sequential([ Conv2D(filters=filters[2], kernel_size=3, strides=1, activation='relu', padding='same'), UpSampling2D((2, 2)), Conv2D(filters=filters[1], kernel_size=3, strides=1, activation='relu', padding='same'), UpSampling2D((2, 2)), Conv2D(filters=filters[0], kernel_size=3, strides=1, activation='relu', padding='valid'), UpSampling2D((2, 2)), Conv2D(1, 3, 1, activation='sigmoid', padding='same') ]) def call(self, encoded): return self.decoder_layers(encoded)
修改后调用model.summary(),即可看到Encoder和Decoder内部的每一层细节。
方法三:先执行一次前向传播,再查看summary
Keras需要输入形状来推断各层输出形状,先传入一个样本完成前向计算,再调用summary:
# 取训练集中第一个样本触发前向传播,替换为你的实际输入数据 model(x_train_noisy[:1]) model.summary()
此方法会让summary显示各层具体输出形状,但不会展开自定义Layer的内部结构,适合仅需输出形状信息的场景。
内容的提问来源于stack exchange,提问作者feelfree
相关产品推荐
相关产品推荐

