如何将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__)创建完成。
修正步骤
- 将
dense_block和transition封装为Keras Layer子类,确保内部层的变量在初始化时创建。 - 在CNN模型的
__init__中预先实例化所有dense block和transition层,避免call时动态创建。 - 移除全局变量
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
相关产品推荐
相关产品推荐

