如何在Keras中实现卷积核共享及核旋转?
Great question! Sharing convolutional kernels with rotation across layers is a clever approach for tasks where rotational invariance or weight sharing with transformed kernels is needed. Here's how you can implement this in Keras/TensorFlow:
Core Idea
Instead of just copying the kernel, we'll share a single base kernel variable and apply rotation transformations to it dynamically when using it in different layers. This way, all layers depend on the same trainable parameter, but each uses a rotated version of it during computation.
Method 1: Custom Rotated Convolution Layer
Subclass Keras' Layer to create a reusable layer that takes a shared kernel, rotates it, and performs convolution.
Step 1: Define the Shared Kernel
First, create a single trainable kernel variable that all layers will share:
import tensorflow as tf from tensorflow.keras import layers, Model # Define shared kernel (adjust shape to match your input/output channels and kernel size) shared_kernel = tf.Variable( tf.random.normal(shape=(3, 3, 3, 16)), # (kernel_height, kernel_width, in_channels, out_channels) name="shared_conv_kernel", trainable=True )
Step 2: Build the Custom Rotated Conv Layer
This layer handles rotation and convolution logic:
class RotatedConv2D(layers.Layer): def __init__(self, rotation_angle=90, padding="same", activation=None, **kwargs): super().__init__(**kwargs) self.rotation_angle = rotation_angle self.padding = padding self.activation = tf.keras.activations.get(activation) def call(self, inputs, shared_kernel): # Rotate the shared kernel based on the specified angle if self.rotation_angle == 90: rotated_kernel = tf.image.rot90(shared_kernel, k=1) # Rotate 90° clockwise elif self.rotation_angle == 180: rotated_kernel = tf.image.rot90(shared_kernel, k=2) elif self.rotation_angle == 270: rotated_kernel = tf.image.rot90(shared_kernel, k=3) else: rotated_kernel = shared_kernel # No rotation as default # Perform convolution with the rotated kernel x = tf.nn.conv2d( inputs, rotated_kernel, strides=(1, 1), padding=self.padding.upper() ) if self.activation is not None: x = self.activation(x) return x # Add get_config for model serialization compatibility def get_config(self): config = super().get_config() config.update({ "rotation_angle": self.rotation_angle, "padding": self.padding, "activation": tf.keras.activations.serialize(self.activation) }) return config
Step 3: Build the Model with Shared Rotated Kernels
Use the custom layer with the shared kernel across different layers:
input_layer = layers.Input(shape=(224, 224, 3)) # Layer 1: Use the original shared kernel x_original = tf.nn.conv2d( input_layer, shared_kernel, strides=(1, 1), padding="SAME" ) x_original = layers.Activation("relu")(x_original) # Layer 2: Use 90° rotated shared kernel x_rot90 = RotatedConv2D(rotation_angle=90, activation="relu")(input_layer, shared_kernel) # Layer 3: Use 180° rotated shared kernel x_rot180 = RotatedConv2D(rotation_angle=180, activation="relu")(input_layer, shared_kernel) # Combine outputs and add downstream layers merged = layers.concatenate([x_original, x_rot90, x_rot180]) flattened = layers.Flatten()(merged) output = layers.Dense(10, activation="softmax")(flattened) model = Model(inputs=input_layer, outputs=output) model.compile( optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"] )
Method 2: Reuse a Base Conv Layer's Weights
If you prefer using built-in Conv2D layers, you can extract the kernel from a base layer, rotate it, and pass it to another layer (with trainable=False since we're reusing the base kernel):
input_layer = layers.Input(shape=(224, 224, 3)) # Base convolution layer (holds the shared kernel) base_conv = layers.Conv2D(16, 3, padding="same", activation="relu") x_base = base_conv(input_layer) # Create a 90° rotated version of the base kernel rotated_kernel_90 = tf.image.rot90(base_conv.kernel, k=1) # Use the rotated kernel in another Conv2D layer (disable training for this layer) rot90_conv = layers.Conv2D(16, 3, padding="same", activation="relu", trainable=False) rot90_conv.build(input_layer.shape) rot90_conv.kernel = rotated_kernel_90 x_rot90 = rot90_conv(input_layer) # Continue building your model as needed
Key Notes
- Rotation Correctness:
tf.image.rot90rotates the spatial dimensions (first two axes) of the kernel, which is exactly what we need for convolutional kernels. For arbitrary angles, you'd need custom rotation logic (since built-in rotation layers are designed for inputs, not kernels). - Weight Sharing Guarantee: All rotated kernels are derived from the same trainable variable (
shared_kernelorbase_conv.kernel), so updates to this variable will propagate to all layers using rotated versions. - Model Serialization: If you need to save/load your model, ensure your custom layer implements
get_config(as shown in Method 1) so Keras can reconstruct it correctly.
内容的提问来源于stack exchange,提问作者Matrixwira

