-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtest_resunet.py
More file actions
83 lines (73 loc) · 3.35 KB
/
Copy pathtest_resunet.py
File metadata and controls
83 lines (73 loc) · 3.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
from re import X
import tensorflow as tf
import numpy as np
input_shape = (224,224,3)
inputs = tf.keras.layers.Input(input_shape)
#
x = tf.keras.layers.Conv2D(64, (3,3),strides=1,padding='same') (inputs)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(64, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(64, (1,1),strides=1,padding='same') (inputs)
s1 = x + s
x = tf.keras.layers.BatchNormalization() (s1)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(128, (3,3),strides=2,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(128, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(128, (1,1),strides=2,padding='same') (s1)
s2 = x + s
x = tf.keras.layers.BatchNormalization() (s2)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(256, (3,3),strides=2,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(256, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(256, (1,1),strides=2,padding='same') (s2)
s3 = x + s
x = tf.keras.layers.BatchNormalization() (s3)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(512, (3,3),strides=2,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(512, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(512, (1,1),strides=2,padding='same') (s3)
b = x + s
x = tf.keras.layers.UpSampling2D((2,2)) (b)
d3 = tf.keras.layers.Concatenate() ([x,s3])
x = tf.keras.layers.BatchNormalization() (d3)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(256, (3,3),strides=1,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(256, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(256, (1,1),strides=1,padding='same') (d3)
x = x + s
x = tf.keras.layers.UpSampling2D((2,2)) (x)
d2 = tf.keras.layers.Concatenate() ([x,s2])
x = tf.keras.layers.BatchNormalization() (d2)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(128, (3,3),strides=1,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(128, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(128, (1,1),strides=1,padding='same') (d2)
x = x + s
x = tf.keras.layers.UpSampling2D((2,2)) (x)
d1 = tf.keras.layers.Concatenate() ([x,s1])
x = tf.keras.layers.BatchNormalization() (d1)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(64, (3,3),strides=1,padding='same') (x)
x = tf.keras.layers.BatchNormalization() (x)
x = tf.keras.layers.Activation('relu') (x)
x = tf.keras.layers.Conv2D(64, (3,3),strides=1,padding='same') (x)
s = tf.keras.layers.Conv2D(64, (1,1),strides=1,padding='same') (d1)
x = x + s
x = tf.keras.layers.Conv2D(1, (1,1),strides=1,padding='same') (x)
outputs = tf.keras.layers.Activation('sigmoid') (x)
model = tf.keras.Model(inputs=inputs, outputs=outputs,name='ResUNet')
print('plotting model')
tf.keras.utils.plot_model(model,to_file='/home/heijkoop/Desktop/ResUNet/TF_ResUNet.png',show_shapes=True)
print('model plotted')
model.summary()