CoolFace
Apppublic

frozencherry/Forgery-Localization-App

sourceHugging Faceupdated 1y agoView on Hugging Face
1likes
custom_layer.py58 linesDownload Raw Back to root
1import tensorflow as tf
2from tensorflow.keras import layers
3
4
5class ChannelAttention(layers.Layer):
6    def __init__(self, ratio=8, **kwargs):
7        super(ChannelAttention, self).__init__(**kwargs)
8        self.ratio = ratio
9
10    def build(self, input_shape):
11        channels = input_shape[-1]
12        self.shared_dense_one = layers.Dense(channels//self.ratio, activation='relu', kernel_initializer='he_normal', use_bias=True)
13        self.shared_dense_two = layers.Dense(channels, kernel_initializer='he_normal', use_bias=True)
14
15    def call(self, inputs):
16        avg_pool = layers.GlobalAveragePooling2D()(inputs)    
17        avg_pool = layers.Reshape((1, 1, avg_pool.shape[1]))(avg_pool)
18        avg_pool = self.shared_dense_one(avg_pool)
19        avg_pool = self.shared_dense_two(avg_pool)
20        max_pool = layers.GlobalMaxPooling2D()(inputs)
21        max_pool = layers.Reshape((1, 1, max_pool.shape[1]))(max_pool)
22        max_pool = self.shared_dense_one(max_pool)
23        max_pool = self.shared_dense_two(max_pool)
24        cbam_feature = tf.nn.sigmoid(avg_pool + max_pool)
25        return layers.Multiply()([inputs, cbam_feature])
26
27class SpatialAttention(layers.Layer):
28    def __init__(self, kernel_size=7, **kwargs):
29        super(SpatialAttention, self).__init__(**kwargs)
30        self.kernel_size = kernel_size
31
32    def build(self, input_shape):
33        self.conv2d = layers.Conv2D(1, (self.kernel_size, self.kernel_size), padding='same', kernel_initializer='he_normal', use_bias=False)
34
35    def call(self, inputs):
36        avg_pool = layers.Lambda(lambda x: tf.reduce_mean(x, axis=3, keepdims=True))(inputs)
37        max_pool = layers.Lambda(lambda x: tf.reduce_max(x, axis=3, keepdims=True))(inputs)
38        concat = layers.Concatenate(axis=3)([avg_pool, max_pool])
39        cbam_feature = self.conv2d(concat)
40        cbam_feature = layers.Activation('sigmoid')(cbam_feature)
41        return layers.Multiply()([inputs, cbam_feature])
42
43class CBAM(layers.Layer):
44    def __init__(self, **kwargs):
45        super(CBAM, self).__init__(**kwargs)
46        ratio=8
47        kernel_size=7
48        self.channel_attention = ChannelAttention(ratio)
49        self.spatial_attention = SpatialAttention(kernel_size)
50
51    def call(self, inputs):
52        cbam_feature = self.channel_attention(inputs)
53        cbam_feature = self.spatial_attention(cbam_feature)
54        return cbam_feature
55
56    def build(self, input_shape):
57        ...
58