如何在Keras函数式API构建神经网络过程中查看模型摘要
在Keras函数式API构建过程中实时查看各阶段模型摘要
当然可以!你完全不用等到整个模型定义完成才能查看层结构和参数信息。在Keras函数式API的构建流程中,每完成一个阶段的层堆叠后,都可以通过临时创建小型Model的方式,快速查看当前阶段的模型摘要,就像调用完整模型的summary()一样。
具体实现方法
思路很简单:针对每个阶段的输入和输出张量,临时生成一个Model实例,调用它的summary()方法查看信息。这些临时模型不会干扰后续的层构建,因为Keras函数式API是基于张量的连接关系来工作的,临时模型只是对现有张量链路的封装,不会修改原有的层或张量。
修改后的示例代码
from tensorflow import keras from tensorflow.keras.layers import Input, Conv2D, MaxPooling2D # 输入张量 input_img = Input(shape=(256, 256, 3)) # Stage 1: 完成tower_1构建后查看 tower_1 = Conv2D(64, (1, 1), padding='same', activation='relu')(input_img) tower_1 = Conv2D(64, (3, 3), padding='same', activation='relu')(tower_1) # stage1 # 临时创建模型查看Stage1的结构与参数 temp_stage1 = keras.Model(inputs=input_img, outputs=tower_1) print("=== Stage 1 模型摘要 ===") temp_stage1.summary() # Stage 2: 完成tower_2构建后查看 tower_2 = Conv2D(64, (1, 1), padding='same', activation='relu')(input_img) tower_2 = Conv2D(64, (5, 5), padding='same', activation='relu')(tower_2) # stage2 temp_stage2 = keras.Model(inputs=input_img, outputs=tower_2) print("\n=== Stage 2 模型摘要 ===") temp_stage2.summary() # Stage 3: 完成tower_3构建后查看 tower_3 = MaxPooling2D((3, 3), strides=(1, 1), padding='same')(input_img) tower_3 = Conv2D(64, (1, 1), padding='same', activation='relu')(tower_3) # stage3 temp_stage3 = keras.Model(inputs=input_img, outputs=tower_3) print("\n=== Stage 3 模型摘要 ===") temp_stage3.summary() # 最终构建完整模型 output = keras.layers.concatenate([tower_1, tower_2, tower_3], axis=1) final_model = keras.Model(inputs=input_img, outputs=output) print("\n=== 最终完整模型摘要 ===") final_model.summary()
额外小技巧
如果你只想查看某一层的参数数量,也可以直接访问层的count_params()方法,比如查看Stage1最后一层的参数:
print(f"Stage1最后一层参数数量: {tower_1._keras_history.layer.count_params()}")
这种方式能让你在构建大型模型的过程中,随时验证每一部分的层结构、输出形状和参数是否符合预期,非常方便调试。
内容的提问来源于stack exchange,提问作者Eka
相关产品推荐
相关产品推荐

