Keras plot_model如何创建自定义模块并为其设置名称
Keras自定义模块及plot_model可视化命名实现方案
Keras原生支持创建可被plot_model接口识别的自定义模块,无需修改框架源码即可实现自定义模块独立成块、显示指定名称的效果,和参考示例表现一致。
实现规则
- 自定义模块需继承
tf.keras.layers.Layer或tf.keras.Model基类,所有内部计算逻辑封装在类的call方法中 - 自定义模块的名称通过类初始化方法的
name参数传入,调用父类初始化方法时透传该参数即可实现命名自定义 - 未开启嵌套展开参数时,
plot_model会自动将整个自定义模块渲染为独立的功能块,不会暴露内部的算子细节
代码示例
import tensorflow as tf from tensorflow.keras.utils import plot_model # 定义通用自定义模块类 class CustomBlock(tf.keras.layers.Layer): def __init__(self, filter_num, name=None): # 透传name参数至父类,支持自定义模块名 super().__init__(name=name) self.conv_1 = tf.keras.layers.Conv2D(filter_num, kernel_size=3, padding='same', activation='relu') self.bn_1 = tf.keras.layers.BatchNormalization() self.conv_2 = tf.keras.layers.Conv2D(filter_num, kernel_size=3, padding='same', activation='relu') self.bn_2 = tf.keras.layers.BatchNormalization() def call(self, inputs, training=False): x = self.conv_1(inputs) x = self.bn_1(x, training=training) x = self.conv_2(x) x = self.bn_2(x, training=training) return x # 构建测试模型 input_tensor = tf.keras.Input(shape=(224, 224, 3)) # 初始化自定义模块时指定名称,对应示例中的New_Block x = CustomBlock(64, name="New_Block")(input_tensor) x = tf.keras.layers.GlobalAveragePooling2D()(x) output_tensor = tf.keras.layers.Dense(10, activation='softmax')(x) model = tf.keras.Model(inputs=input_tensor, outputs=output_tensor) # 生成模型结构图 plot_model( model, to_file="custom_block_vis.png", show_shapes=True, # 如需查看自定义模块内部结构,可将下方参数设为True expand_nested=False )
效果参考

补充说明
- 若自定义模块继承自
tf.keras.Model,上述命名、可视化逻辑完全通用,无需额外适配 - 模块名称支持传入任意合法字符串,最终会直接显示在可视化生成的结构块上
- 若出现名称自动追加后缀的情况,是因为同层级存在同名模块,Keras为了避免名称冲突自动添加了唯一标识,确保同层级下模块名唯一即可解决
内容的提问来源于stack exchange,提问作者RACHID BEN ABDELMALEK
相关产品推荐
相关产品推荐

