frozencherry/Forgery-Localization-App
1
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 