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

如何将DenseNet从Keras Functional API转换为Model子类API(报错解决)

DenseNet从Keras Functional API转Model子类API的报错解决

问题背景

尝试将DenseNet从Functional API转换为Model子类API时,出现ValueError: tf.function only supports singleton tf.Variables created on the first call错误,核心原因是call方法内动态创建包含可训练变量的层,违反了tf.function的变量创建规则。

原Functional API代码

inputs = tf.keras.layers.Input( shape=input_shape )
x = tf.keras.layers.Conv2D( num_filters , kernel_size=(3,3) , use_bias=False, kernel_initializer='he_normal' , kernel_regularizer=tf.keras.regularizers.l2( 1e-4 ) )( inputs )

for i in range( num_blocks ):
    x, num_filters = dense_block( x, num_layers_per_block , num_filters, growth_rate , dropout_rate )
    x = transition(x, num_filters , compress_factor , dropout_rate )

x = tf.keras.layers.GlobalAveragePooling2D()( x ) 
x = tf.keras.layers.Dense( 37 )( x ) # Num Classes for CIFAR-10
outputs = tf.keras.layers.Activation( 'softmax' )( x )

转换后的Model子类API代码(报错版本)

class CNN(keras.Model):
    def __init__(self,nfilters,sfilters):
        super(CNN,self).__init__()
        self.conv1 = keras.layers.Conv2D(nfilters[0] , kernel_size=(sfilters[0],sfilters[0]) , use_bias=False, kernel_initializer='he_normal', kernel_regularizer=tf.keras.regularizers.l2( 1e-4 ) )

        self.globalaverage = keras.layers.GlobalAveragePooling2D()
        self.dense = keras.layers.Dense(37) # Num Classes for CIFAR-10
        self.activation = keras.layers.Activation( 'softmax' )

    def call(self, input_tensor, training=False): 
        global num_filters
        x = self.conv1(input_tensor)

        for i in range( num_blocks ):
            x, num_filters = dense_block( x, num_layers_per_block , num_filters, growth_rate , dropout_rate )
            x = transition(x, num_filters , compress_factor , dropout_rate )

        x = self.globalaverage( x ) 
        x = self.dense( x ) # Num Classes for CIFAR-10
        return self.activation( x )

错误信息

ValueError: tf.function only supports singleton tf.Variables created on the first call. Make sure the tf.Variable is only created once or created outside tf.function.

Call arguments received:
  • input_tensor=tf.Tensor(shape=(16, 32, 32, 3), dtype=float32)
  • training=True

解决方案

错误根源

原代码中dense_block和transition函数在call方法内被调用时,每次都会动态创建新的卷积、BN等包含可训练变量的层,而Keras会将call方法包装成tf.function,tf.function不允许在每次调用时生成新的变量,所有可训练层必须在模型初始化阶段(__init__)创建完成。

修正步骤

  1. 将dense_block和transition封装为Keras Layer子类,确保内部层的变量在初始化时创建。
  2. 在CNN模型的__init__中预先实例化所有dense block和transition层,避免call时动态创建。
  3. 移除全局变量num_filters,改用实例内变量追踪通道数。

完整修正代码

1. DenseBlock层实现

class DenseBlock(tf.keras.layers.Layer):
    def __init__(self, num_layers, growth_rate, dropout_rate):
        super().__init__()
        self.layers_seq = []
        for _ in range(num_layers):
            # DenseNet瓶颈层结构:BN -> ReLU -> 1x1Conv -> Dropout -> BN -> ReLU -> 3x3Conv -> Dropout
            self.layers_seq.extend([
                tf.keras.layers.BatchNormalization(),
                tf.keras.layers.Activation('relu'),
                tf.keras.layers.Conv2D(4*growth_rate, (1,1), use_bias=False,
                                      kernel_initializer='he_normal',
                                      kernel_regularizer=tf.keras.regularizers.l2(1e-4)),
                tf.keras.layers.Dropout(dropout_rate),
                tf.keras.layers.BatchNormalization(),
                tf.keras.layers.Activation('relu'),
                tf.keras.layers.Conv2D(growth_rate, (3,3), padding='same', use_bias=False,
                                      kernel_initializer='he_normal',
                                      kernel_regularizer=tf.keras.regularizers.l2(1e-4)),
                tf.keras.layers.Dropout(dropout_rate)
            ])
        self.growth_rate = growth_rate

    def call(self, inputs, training=False):
        x = inputs
        for layer in self.layers_seq:
            if isinstance(layer, (tf.keras.layers.Dropout, tf.keras.layers.BatchNormalization)):
                x = layer(x, training=training)
            else:
                x = layer(x)
            # 完成3x3Conv后拼接原始输入与新特征
            if isinstance(layer, tf.keras.layers.Conv2D) and layer.filters == self.growth_rate:
                x = tf.concat([inputs, x], axis=-1)
                inputs = x
        return x, inputs.shape[-1]

2. TransitionLayer层实现

class TransitionLayer(tf.keras.layers.Layer):
    def __init__(self, num_filters, compress_factor, dropout_rate):
        super().__init__()
        self.target_filters = int(num_filters * compress_factor)
        self.bn = tf.keras.layers.BatchNormalization()
        self.relu = tf.keras.layers.Activation('relu')
        self.conv = tf.keras.layers.Conv2D(self.target_filters, (1,1), use_bias=False,
                                           kernel_initializer='he_normal',
                                           kernel_regularizer=tf.keras.regularizers.l2(1e-4))
        self.dropout = tf.keras.layers.Dropout(dropout_rate)
        self.pool = tf.keras.layers.AveragePooling2D((2,2), strides=(2,2))

    def call(self, inputs, training=False):
        x = self.bn(inputs, training=training)
        x = self.relu(x)
        x = self.conv(x)
        x = self.dropout(x, training=training)
        x = self.pool(x)
        return x, self.target_filters

3. 最终CNN模型实现

class CNN(tf.keras.Model):
    def __init__(self, initial_num_filters, num_blocks, num_layers_per_block, growth_rate, compress_factor, dropout_rate):
        super().__init__()
        # 初始卷积层
        self.conv1 = tf.keras.layers.Conv2D(initial_num_filters, kernel_size=(3,3), padding='same', use_bias=False,
                                            kernel_initializer='he_normal',
                                            kernel_regularizer=tf.keras.regularizers.l2(1e-4))
        # 预先创建所有DenseBlock和TransitionLayer
        self.dense_blocks = []
        self.transitions = []
        current_filters = initial_num_filters
        
        for _ in range(num_blocks):
            self.dense_blocks.append(DenseBlock(num_layers_per_block, growth_rate, dropout_rate))
            # 计算DenseBlock输出后的通道数
            current_filters += num_layers_per_block * growth_rate
            self.transitions.append(TransitionLayer(current_filters, compress_factor, dropout_rate))
            # 更新通道数为Transition层输出后的数量
            current_filters = int(current_filters * compress_factor)
        
        # 分类头部
        self.globalaverage = tf.keras.layers.GlobalAveragePooling2D()
        self.dense = tf.keras.layers.Dense(37)
        self.activation = tf.keras.layers.Activation('softmax')

    def call(self, input_tensor, training=False):
        x = self.conv1(input_tensor)
        current_filters = self.conv1.filters

        # 遍历所有DenseBlock和TransitionLayer
        for db, trans in zip(self.dense_blocks, self.transitions):
            x, current_filters = db(x, training=training)
            x, current_filters = trans(x, training=training)
        
        x = self.globalaverage(x)
        x = self.dense(x)
        return self.activation(x)

使用示例

# 模型参数配置
input_shape = (32, 32, 3)
initial_num_filters = 64
num_blocks = 3
num_layers_per_block = 4
growth_rate = 16
compress_factor = 0.5
dropout_rate = 0.2

# 初始化并构建模型
model = CNN(initial_num_filters, num_blocks, num_layers_per_block, growth_rate, compress_factor, dropout_rate)
model.build((None,) + input_shape)
model.summary()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 03:16:05