【问题标题】:How to get summary graph of custom (subclass) Keras layer?如何获取自定义(子类)Keras 层的摘要图?
【发布时间】:2021-10-29 06:06:06
【问题描述】:

如何打印自定义层的层的 summary()?

model.summary() 打印了整个模型的漂亮摘要图,但是这里称为“magic_layer”的子类层,其中有很多层,是聚合的......

Model: "transformer"
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
positioning (PositionalEncod (None, 504, 6)            0         
_________________________________________________________________
magic_layer (CustomLayer)    (None, 6, 504)            3088040   
_________________________________________________________________
g_pooling (GlobalAveragePool (None, 504)               0         
_________________________________________________________________
dropout_2 (Dropout)          (None, 504)               0         
_________________________________________________________________
dense_2 (Dense)              (None, 32)                16160     
_________________________________________________________________
dropout_3 (Dropout)          (None, 32)                0         
_________________________________________________________________
dense_3 (Dense)              (None, 5)                 165       
=================================================================
Total params: 3,104,365
Trainable params: 3,104,365
Non-trainable params: 0
_________________________________________________________________

如果您有自定义的 Tensorflow/Keras 层(在此处了解更多信息:Making new layers and models via subclassing - Francis Chollet),那么摘要调用不会分解该子层中的所有层。本例中的“magic_layer”是我感兴趣的子类层。

在本例中,您如何为名为“magic_layer”的层获得相同的子层打印输出?

model.layers[1].summary() 不幸的是不起作用...也许我需要在自定义层类中包含摘要定义,但我希望有一种方法可以从模型类继承此功能。

【问题讨论】:

    标签: tensorflow keras


    【解决方案1】:

    由于模型是 layer 的子类,只需从 tf.keras.Model 而不是 tf.keras.layers.Layer 制作您的自定义 layer 子类。现在您可以通过 summary() 打印“层”的摘要。

    model.summary 不是递归的 - 它不会打印嵌入式模型的摘要。如果你想这样,你必须自己写,或者只是根据原始来源创建自己的摘要函数。

    https://github.com/keras-team/keras/blob/07a22914c8114a74238fd86741749cab5af299ce/keras/utils/layer_utils.py#L116

    【讨论】:

    • 我不知道模型是图层的子类。本来以为是反过来的。但是,我最初确实从 keras.layers.layer class EncoderLayer(keras.layers.Layer): 制作了我的自定义层子类,并按照您的建议从 tf.keras.Model 制作了模型。
    • 抱歉,我的阅读障碍症引起了我的注意,我颠倒了图层和模型。修改!来自 tf.keras.Model 而不是 Layer 的子类。
    猜你喜欢
    • 1970-01-01
    • 2023-01-25
    • 1970-01-01
    • 1970-01-01
    • 2019-05-09
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多