Estoy tratando de producir una CNN usando Keras y escribí el siguiente código:
batch_size = 64 epochs = 20 num_classes = 5 cnn_model = Sequential() cnn_model.add(Conv2D(32, kernel_size=(3, 3), activation='linear', input_shape=(380, 380, 1), padding='same')) cnn_model.add(Activation('relu')) cnn_model.add(MaxPooling2D((2, 2), padding='same')) cnn_model.add(Conv2D(64, (3, 3), activation='linear', padding='same')) cnn_model.add(Activation('relu')) cnn_model.add(MaxPooling2D(pool_size=(2, 2), padding='same')) cnn_model.add(Conv2D(128, (3, 3), activation='linear', padding='same')) cnn_model.add(Activation('relu')) cnn_model.add(MaxPooling2D(pool_size=(2, 2), padding='same')) cnn_model.add(Flatten()) cnn_model.add(Dense(128, activation='linear')) cnn_model.add(Activation('relu')) cnn_model.add(Dense(num_classes, activation='softmax')) cnn_model.compile(loss=keras.losses.categorical_crossentropy, optimizer=keras.optimizers.Adam(), metrics=['accuracy']) Quiero usar la capa de activación LeakyReLU de Keras en lugar de usar Activation('relu') . Sin embargo, intenté usar LeakyReLU(alpha=0.1) , pero esta es una capa de activación en Keras y aparece un error sobre el uso de una capa de activación y no una función de activación.
¿Cómo puedo usar LeakyReLU en este ejemplo?
Todas las activaciones avanzadas en Keras, incluido LeakyReLU , están disponibles como capas y no como activaciones; por lo tanto, debe usarlo como tal:
from keras.layers import LeakyReLU # instead of cnn_model.add(Activation('relu')) # use cnn_model.add(LeakyReLU(alpha=0.1))A veces, solo desea un reemplazo directo para una capa de activación integrada y no tener que agregar capas de activación adicionales solo para este propósito.
Para eso, puede usar el hecho de que el argumento de activation puede ser un objeto invocable.
lrelu = lambda x: tf.keras.activations.relu(x, alpha=0.1) model.add(Conv2D(..., activation=lrelu, ...) Dado que una Layer también es un objeto invocable, también podría simplemente usar
model.add(Conv2D(..., activation=tf.keras.layers.LeakyReLU(alpha=0.1), ...) que ahora funciona en TF2. Esta es una mejor solución ya que evita la necesidad de usar un objeto personalizado durante la carga como mencionó custom_object .