如何在TensorFlow中动态创建指定数量的网络层?
add()是tf.keras.Sequential序列模型的专属方法,继承tf.keras.Model实现的自定义模型类默认不支持该方法,要实现根据参数动态生成层的需求,可以用以下两种方案:
方案1:使用LayerList容器存储层(推荐)
tf.keras.layers.LayerList是TensorFlow提供的专门用来存储层的容器,容器内所有层的参数会被自动注册到模型中,可正常参与训练。
示例代码如下:
import tensorflow as tf class generic_vns_function(tf.keras.Model): def __init__(self, num_layers, num_class=10): super().__init__() # 用LayerList存储动态生成的卷积层 self.conv_layers = tf.keras.layers.LayerList() for i in range(num_layers): self.conv_layers.append(tf.keras.layers.Conv2D(64, 3, activation="relu")) # 按需添加其他层,比如池化、全连接层等 self.max_pool = tf.keras.layers.MaxPool2D(2) self.flatten = tf.keras.layers.Flatten() self.fc = tf.keras.layers.Dense(num_class, activation="softmax") def call(self, inputs): x = inputs # 遍历LayerList依次执行卷积计算 for conv in self.conv_layers: x = conv(x) x = self.max_pool(x) x = self.flatten(x) x = self.fc(x) return x
方案2:封装Sequential容器
如果不想手动遍历层,也可以直接在自定义类内创建一个Sequential容器,把动态生成的层都添加到这个容器中,调用时直接执行该容器即可:
import tensorflow as tf class generic_vns_function(tf.keras.Model): def __init__(self, num_layers, num_class=10): super().__init__() # 创建Sequential容器统一管理层 self.feature_extractor = tf.keras.Sequential() for i in range(num_layers): self.feature_extractor.add(tf.keras.layers.Conv2D(64, 3, activation="relu")) self.feature_extractor.add(tf.keras.layers.MaxPool2D(2)) self.feature_extractor.add(tf.keras.layers.Flatten()) self.feature_extractor.add(tf.keras.layers.Dense(num_class, activation="softmax")) def call(self, inputs): return self.feature_extractor(inputs)
注意事项
不要直接用普通Python列表存储生成的层,普通列表内的层参数不会被TensorFlow自动识别,会导致训练时模型无可用参数的问题。
内容的提问来源于stack exchange,提问作者logankilpatrick
相关产品推荐
相关产品推荐

