You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用TensorFlow层构建模型,如何实现类Keras风格的模型摘要展示?

如何用TensorFlow层构建模型并输出Keras风格的模型摘要?

当然可以!其实TensorFlow和Keras现在深度整合,只要你的模型能和Keras的Model类兼容,就能直接输出和keras.Model.summary()完全一致风格的摘要。下面分几种常见情况给你说明:

情况1:直接用Keras层构建模型(最常用)

如果你是用tf.keras.layers通过Sequential或者Functional API搭建模型,那调用summary()就是原生支持的,完全不用额外操作:

import tensorflow as tf

# 用Functional API示例
inputs = tf.keras.Input(shape=(28, 28))
x = tf.keras.layers.Flatten()(inputs)
x = tf.keras.layers.Dense(64, activation='relu')(x)
outputs = tf.keras.layers.Dense(10, activation='softmax')(x)

model = tf.keras.Model(inputs=inputs, outputs=outputs)
# 直接输出Keras风格摘要
model.summary()

运行后你会看到熟悉的层级结构、参数数量、输出形状等信息,和传统Keras的summary完全一致。

情况2:自定义tf.Module包装成Keras模型

如果你的模型是用低级API tf.Module构建的,也可以轻松把它包装成Keras Model来调用summary():

class MyCustomModule(tf.Module):
    def __init__(self):
        super().__init__()
        self.flatten = tf.keras.layers.Flatten()
        self.dense1 = tf.keras.layers.Dense(64, activation='relu')
        self.dense2 = tf.keras.layers.Dense(10, activation='softmax')
    
    @tf.function
    def __call__(self, x):
        x = self.flatten(x)
        x = self.dense1(x)
        return self.dense2(x)

# 实例化自定义Module
module = MyCustomModule()
# 先传入一个dummy输入让模型初始化权重和形状
dummy_input = tf.random.normal((1, 28, 28))
_ = module(dummy_input)

# 包装成Keras Model
keras_model = tf.keras.Model(inputs=dummy_input, outputs=module(dummy_input))
keras_model.summary()

情况3:子类化Keras Model的自定义模型

如果你用子类化方式写模型,需要注意先让模型接收一次输入(或者显式指定input_shape),才能调用summary(),因为子类化模型的输入形状是动态确定的:

class MyModel(tf.keras.Model):
    def __init__(self):
        super().__init__()
        self.flatten = tf.keras.layers.Flatten()
        self.dense1 = tf.keras.layers.Dense(64, activation='relu')
        self.dense2 = tf.keras.layers.Dense(10, activation='softmax')
    
    def call(self, x):
        x = self.flatten(x)
        x = self.dense1(x)
        return self.dense2(x)

model = MyModel()
# 先传入一个输入让模型确定形状
model(tf.random.normal((1, 28, 28)))
# 现在可以正常输出summary了
model.summary()

核心思路就是:只要你的模型能转化为tf.keras.Model的实例,就能直接调用summary()得到你想要的Keras风格模型摘要~

内容的提问来源于stack exchange,提问作者Ram Sure

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 04:15:37