TensorFlow Model Subclassing用vars()无法显示参数与层的问题求解
问题描述
我编写了用于实现VGG Block的代码,想要查看模块的summary:
import tensorflow as tf from keras.layers import Conv2D, MaxPool2D, Input class VggBlock(tf.keras.Model): def __init__(self, filters, repetitions): super(VggBlock, self).__init__() self.repetitions = repetitions for i in range(repetitions): vars(self)[f'conv2D_{i}'] = Conv2D(filters=filters, kernel_size=(3, 3), padding='same', activation='relu') self.max_pool = MaxPool2D(pool_size=(2, 2)) def call(self, inputs): x = vars(self)['conv2D_0'](inputs) for i in range(1, self.repetitions): x = vars(self)[f'conv2D_{i}'](x) return self.max_pool(x) test_block = VggBlock(64, 2) temp_inputs = Input(shape=(224, 224, 3)) test_block(temp_inputs) test_block.summary()
执行代码后得到的输出:
Model: "vgg_block" _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= max_pooling2d (MaxPooling2D multiple 0 ) ================================================================= Total params: 0 Trainable params: 0 Non-trainable params: 0 _________________________________________________________________
我尝试显式检查层:
for layer in test_block.layers: print(layer)
输出仅显示一个层:
<keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7f6c18377f50>
但卷积层确实以字典形式存在(通过print(vars(test_block))输出可见):
{'_self_setattr_tracking': True, '_is_model_for_instrumentation': True, '_instrumented_keras_api': True, '_instrumented_keras_layer_class': False, '_instrumented_keras_model_class': True, '_trainable': True, '_stateful': False, 'built': True, '_input_spec': None, '_build_input_shape': None, '_saved_model_inputs_spec': TensorSpec(shape=(None, 224, 224, 3), dtype=tf.float32, name='input_10'), '_saved_model_arg_spec': ([TensorSpec(shape=(None, 224, 224, 3), dtype=tf.float32, name='input_10')], {}), '_supports_masking': False, '_name': 'vgg_block_46', '_activity_regularizer': None, '_trainable_weights': [], '_non_trainable_weights': [], '_updates': [], '_thread_local': <_thread._local object at 0x7fb9084d9ef0>, '_callable_losses': [], '_losses': [], '_metrics': [], '_metrics_lock': <unlocked _thread.lock object at 0x7fb90d88abd0>, '_dtype_policy': <Policy "float32">, '_compute_dtype_object': tf.float32, '_autocast': True, '_self_tracked_trackables': [<keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7fb9084e2510>], '_inbound_nodes_value': [<keras.engine.node.Node object at 0x7fb9087146d0>], '_outbound_nodes_value': [], '_expects_training_arg': False, '_default_training_arg': None, '_expects_mask_arg': False, '_dynamic': False, '_initial_weights': None, '_auto_track_sub_layers': True, '_preserve_input_structure_in_config': False, '_name_scope_on_declaration': '', '_captured_weight_regularizer': [], '_is_graph_network': False, 'inputs': None, 'outputs': None, 'input_names': None, 'output_names': None, 'stop_training': False, 'history': None, 'compiled_loss': None, 'compiled_metrics': None, '_compute_output_and_mask_jointly': False, '_is_compiled': False, 'optimizer': None, '_distribution_strategy': None, '_cluster_coordinator': None, '_run_eagerly': None, 'train_function': None, 'test_function': None, 'predict_function': None, 'train_tf_function': None, '_compiled_trainable_state': <WeakKeyDictionary at 0x7fb9084b0790>, '_training_state': None, '_self_unconditional_checkpoint_dependencies': [TrackableReference(name=max_pool, ref=<keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7fb9084e2510>)], '_self_unconditional_dependency_names': {'max_pool': <keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7fb9084e2510>}, '_self_unconditional_deferred_dependencies': {}, '_self_update_uid': -1, '_self_name_based_restores': set(), '_self_saveable_object_factories': {}, '_checkpoint': <tensorflow.python.training.tracking.util.Checkpoint object at 0x7fb9084b0910>, '_steps_per_execution': None, '_train_counter': <tf.Variable 'Variable:0' shape=() dtype=int64, numpy=0>, '_test_counter': <tf.Variable 'Variable:0' shape=() dtype=int64, numpy=0>, '_predict_counter': <tf.Variable 'Variable:0' shape=() dtype=int64, numpy=0>, '_base_model_initialized': True, '_jit_compile': None, '_layout_map': None, '_obj_reference_counts_dict': ObjectIdentityDictionary({<_ObjectIdentityWrapper wrapping 3>: 1, <_ObjectIdentityWrapper wrapping <keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7fb9084e2510>>: 1}), 'repetitions': 3, 'conv2D_0': <keras.layers.convolutional.conv2d.Conv2D object at 0x7fb90852e390>, 'conv2D_1': <keras.layers.convolutional.conv2d.Conv2D object at 0x7fb90852ed90>, 'conv2D_2': <keras.layers.convolutional.conv2d.Conv2D object at 0x7fb9084dac90>, 'max_pool': <keras.layers.pooling.max_pooling2d.MaxPooling2D object at 0x7fb9084e2510>}
疑问:vars()是否导致了异常?如何正确显示模型的层和参数?
问题原因与解决方法
原因分析
使用vars(self)[f'conv2D_{i}']动态添加层的方式,绕过了Keras模型的子层自动跟踪机制。Keras只会将通过直接赋值(self.layer_name = Layer(...))创建的子层加入到layers列表中,通过vars()修改实例字典的操作不会触发跟踪逻辑,因此卷积层无法被模型识别为正式子层,也就不会出现在summary和test_block.layers的输出里。
解决方法
有两种可靠的修改方式:
方式一:使用setattr显式赋值子层
在__init__中通过setattr给实例动态添加卷积层,确保Keras能跟踪到这些子层:
class VggBlock(tf.keras.Model): def __init__(self, filters, repetitions): super(VggBlock, self).__init__() self.repetitions = repetitions # 用setattr动态赋值,触发Keras的子层跟踪 for i in range(repetitions): setattr(self, f'conv2D_{i}', Conv2D(filters=filters, kernel_size=(3, 3), padding='same', activation='relu')) self.max_pool = MaxPool2D(pool_size=(2, 2)) def call(self, inputs): x = getattr(self, 'conv2D_0')(inputs) for i in range(1, self.repetitions): x = getattr(self, f'conv2D_{i}')(x) return self.max_pool(x)
方式二:用列表统一管理卷积层
创建列表存储所有卷积层,再将列表赋值给实例属性,Keras会自动跟踪列表中的所有层:
class VggBlock(tf.keras.Model): def __init__(self, filters, repetitions): super(VggBlock, self).__init__() self.repetitions = repetitions # 用列表存储所有卷积层,Keras会自动跟踪列表内的层 self.conv_layers = [ Conv2D(filters=filters, kernel_size=(3, 3), padding='same', activation='relu') for _ in range(repetitions) ] self.max_pool = MaxPool2D(pool_size=(2, 2)) def call(self, inputs): x = inputs for conv_layer in self.conv_layers: x = conv_layer(x) return self.max_pool(x)
修改后重新运行代码,test_block.summary()就能正确显示所有卷积层和参数统计,test_block.layers也会包含所有子层。
内容的提问来源于stack exchange,提问作者user20620000
相关产品推荐
相关产品推荐

