TensorFlow自定义层输出命名设置:匹配tf.data数据集字典键
问题解决方法
你遇到的问题是自定义层内部操作直接暴露在模型摘要中,根因是重写了tf.keras.layers.Layer的__call__方法,而非Keras规范要求的call方法。父类Layer的__call__内置了层封装、节点追踪等逻辑,直接覆盖会导致自定义层的封装失效。
修改后的完整代码
import tensorflow as tf class CustomLayer(tf.keras.layers.Layer): def __init__(self, name=None): super(CustomLayer, self).__init__(name=name) self.dense = tf.keras.layers.Dense(32) # 把__call__改为call即可实现层整体封装 def call(self, inputs): x = self.dense(inputs) x = tf.add(x, 42) # 如果需要让输出张量名称严格匹配层名,适配tf.data字典结构,可取消下行注释 # x = tf.identity(x, name=self.name) return x m = tf.keras.models.Sequential() m.add(tf.keras.Input(shape=(100,))) m.add(tf.keras.layers.Dense(84)) m.add(CustomLayer(name='custom_layer')) m.summary()
运行后输出的模型摘要
Model: "sequential" _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= dense (Dense) (None, 84) 8484 custom_layer (CustomLayer) (None, 32) 2720 ================================================================= Total params: 11,204 Trainable params: 11,204 Non-trainable params: 0 _________________________________________________________________
完全符合你期望的显示效果,自定义层会作为整体展示,不会暴露内部操作。
适配tf.data字典输入输出的额外说明
如果你需要让模型输出张量的名称严格匹配tf.data.Dataset的字典键,只需要取消上述代码中tf.identity那行的注释即可,最终输出张量的名称会和你传入的层名custom_layer完全一致,不需要额外调整其他配置。
内容的提问来源于stack exchange,提问作者Scholar
相关产品推荐
相关产品推荐

