Update Deblurring model_modded.py
Browse files- Deblurring model_modded.py +5 -21
Deblurring model_modded.py
CHANGED
|
@@ -8,7 +8,6 @@ from tensorflow.keras.models import Model
|
|
| 8 |
from tensorflow.keras.optimizers import Adam
|
| 9 |
import cv2
|
| 10 |
|
| 11 |
-
# Residual Block for the Generator
|
| 12 |
def residual_block(x, filters):
|
| 13 |
res = Conv2D(filters, kernel_size=3, strides=1, padding='same',
|
| 14 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
|
@@ -19,7 +18,6 @@ def residual_block(x, filters):
|
|
| 19 |
res = InstanceNormalization()(res)
|
| 20 |
return Add()([x, res])
|
| 21 |
|
| 22 |
-
# Generator Architecture
|
| 23 |
def build_generator():
|
| 24 |
inputs = Input(shape=(256, 256, 3))
|
| 25 |
x = Conv2D(64, kernel_size=7, strides=1, padding='same',
|
|
@@ -27,7 +25,6 @@ def build_generator():
|
|
| 27 |
x = InstanceNormalization()(x)
|
| 28 |
x = Activation('relu')(x)
|
| 29 |
|
| 30 |
-
# Down-sampling
|
| 31 |
x = Conv2D(128, kernel_size=3, strides=2, padding='same',
|
| 32 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 33 |
x = InstanceNormalization()(x)
|
|
@@ -38,11 +35,9 @@ def build_generator():
|
|
| 38 |
x = InstanceNormalization()(x)
|
| 39 |
x = Activation('relu')(x)
|
| 40 |
|
| 41 |
-
# Residual Blocks
|
| 42 |
for _ in range(9):
|
| 43 |
x = residual_block(x, 256)
|
| 44 |
|
| 45 |
-
# Up-sampling
|
| 46 |
x = Conv2DTranspose(128, kernel_size=3, strides=2, padding='same',
|
| 47 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 48 |
x = InstanceNormalization()(x)
|
|
@@ -57,7 +52,6 @@ def build_generator():
|
|
| 57 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 58 |
return Model(inputs, outputs, name="Generator")
|
| 59 |
|
| 60 |
-
# Discriminator Architecture
|
| 61 |
def build_discriminator():
|
| 62 |
inputs = Input(shape=(256, 256, 3))
|
| 63 |
|
|
@@ -81,14 +75,12 @@ def build_discriminator():
|
|
| 81 |
x = LeakyReLU(0.2)(x)
|
| 82 |
|
| 83 |
x = GlobalAveragePooling2D()(x)
|
| 84 |
-
outputs = Dense(1)(x)
|
| 85 |
return Model(inputs, outputs, name="Discriminator")
|
| 86 |
|
| 87 |
-
# Wasserstein Loss
|
| 88 |
def wasserstein_loss(y_true, y_pred):
|
| 89 |
return -tf.reduce_mean(y_true * y_pred)
|
| 90 |
|
| 91 |
-
# Load and Preprocess Data
|
| 92 |
def load_images(folder):
|
| 93 |
images = []
|
| 94 |
for filename in os.listdir(folder):
|
|
@@ -98,7 +90,6 @@ def load_images(folder):
|
|
| 98 |
images.append(img)
|
| 99 |
return np.array(images)
|
| 100 |
|
| 101 |
-
# Training Pipeline
|
| 102 |
def train(generator, discriminator, blurred_images, clear_images, epochs, batch_size):
|
| 103 |
optimizer_g = Adam(learning_rate=0.0001, beta_1=0.5, beta_2=0.999)
|
| 104 |
optimizer_d = Adam(learning_rate=0.0002, beta_1=0.5, beta_2=0.999)
|
|
@@ -106,38 +97,31 @@ def train(generator, discriminator, blurred_images, clear_images, epochs, batch_
|
|
| 106 |
for epoch in range(epochs):
|
| 107 |
print(f"Epoch {epoch + 1}/{epochs}")
|
| 108 |
for i in range(0, len(blurred_images), batch_size):
|
| 109 |
-
# Prepare batches
|
| 110 |
blurred_batch = blurred_images[i:i + batch_size]
|
| 111 |
clear_batch = clear_images[i:i + batch_size]
|
| 112 |
|
| 113 |
-
# Generate fake images
|
| 114 |
fake_images = generator.predict(blurred_batch)
|
| 115 |
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
fake_labels = np.ones((len(fake_images), 1)) # Label fake as 1
|
| 119 |
d_loss_real = discriminator.train_on_batch(clear_batch, real_labels)
|
| 120 |
d_loss_fake = discriminator.train_on_batch(fake_images, fake_labels)
|
| 121 |
d_loss = d_loss_real + d_loss_fake
|
| 122 |
|
| 123 |
-
|
| 124 |
-
misleading_labels = -np.ones((len(blurred_batch), 1)) # Fool discriminator
|
| 125 |
g_loss = generator.train_on_batch(blurred_batch, misleading_labels)
|
| 126 |
|
| 127 |
print(f"Batch {i // batch_size + 1}: D Loss: {d_loss:.4f}, G Loss: {g_loss:.4f}")
|
| 128 |
|
| 129 |
-
# Paths and Data Loading
|
| 130 |
blurred_folder = "blurred_sketches"
|
| 131 |
clear_folder = "clear_sketches"
|
| 132 |
blurred_images = load_images(blurred_folder)
|
| 133 |
clear_images = load_images(clear_folder)
|
| 134 |
|
| 135 |
-
# Instantiate Models
|
| 136 |
generator = build_generator()
|
| 137 |
discriminator = build_discriminator()
|
| 138 |
|
| 139 |
-
# Compile Discriminator
|
| 140 |
discriminator.compile(optimizer=Adam(0.0002, 0.5, 0.999), loss=wasserstein_loss)
|
| 141 |
|
| 142 |
-
|
| 143 |
train(generator, discriminator, blurred_images, clear_images, epochs=500, batch_size=16)
|
|
|
|
| 8 |
from tensorflow.keras.optimizers import Adam
|
| 9 |
import cv2
|
| 10 |
|
|
|
|
| 11 |
def residual_block(x, filters):
|
| 12 |
res = Conv2D(filters, kernel_size=3, strides=1, padding='same',
|
| 13 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
|
|
|
| 18 |
res = InstanceNormalization()(res)
|
| 19 |
return Add()([x, res])
|
| 20 |
|
|
|
|
| 21 |
def build_generator():
|
| 22 |
inputs = Input(shape=(256, 256, 3))
|
| 23 |
x = Conv2D(64, kernel_size=7, strides=1, padding='same',
|
|
|
|
| 25 |
x = InstanceNormalization()(x)
|
| 26 |
x = Activation('relu')(x)
|
| 27 |
|
|
|
|
| 28 |
x = Conv2D(128, kernel_size=3, strides=2, padding='same',
|
| 29 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 30 |
x = InstanceNormalization()(x)
|
|
|
|
| 35 |
x = InstanceNormalization()(x)
|
| 36 |
x = Activation('relu')(x)
|
| 37 |
|
|
|
|
| 38 |
for _ in range(9):
|
| 39 |
x = residual_block(x, 256)
|
| 40 |
|
|
|
|
| 41 |
x = Conv2DTranspose(128, kernel_size=3, strides=2, padding='same',
|
| 42 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 43 |
x = InstanceNormalization()(x)
|
|
|
|
| 52 |
kernel_initializer=tf.keras.initializers.RandomNormal(mean=0.0, stddev=0.02))(x)
|
| 53 |
return Model(inputs, outputs, name="Generator")
|
| 54 |
|
|
|
|
| 55 |
def build_discriminator():
|
| 56 |
inputs = Input(shape=(256, 256, 3))
|
| 57 |
|
|
|
|
| 75 |
x = LeakyReLU(0.2)(x)
|
| 76 |
|
| 77 |
x = GlobalAveragePooling2D()(x)
|
| 78 |
+
outputs = Dense(1)(x)
|
| 79 |
return Model(inputs, outputs, name="Discriminator")
|
| 80 |
|
|
|
|
| 81 |
def wasserstein_loss(y_true, y_pred):
|
| 82 |
return -tf.reduce_mean(y_true * y_pred)
|
| 83 |
|
|
|
|
| 84 |
def load_images(folder):
|
| 85 |
images = []
|
| 86 |
for filename in os.listdir(folder):
|
|
|
|
| 90 |
images.append(img)
|
| 91 |
return np.array(images)
|
| 92 |
|
|
|
|
| 93 |
def train(generator, discriminator, blurred_images, clear_images, epochs, batch_size):
|
| 94 |
optimizer_g = Adam(learning_rate=0.0001, beta_1=0.5, beta_2=0.999)
|
| 95 |
optimizer_d = Adam(learning_rate=0.0002, beta_1=0.5, beta_2=0.999)
|
|
|
|
| 97 |
for epoch in range(epochs):
|
| 98 |
print(f"Epoch {epoch + 1}/{epochs}")
|
| 99 |
for i in range(0, len(blurred_images), batch_size):
|
|
|
|
| 100 |
blurred_batch = blurred_images[i:i + batch_size]
|
| 101 |
clear_batch = clear_images[i:i + batch_size]
|
| 102 |
|
|
|
|
| 103 |
fake_images = generator.predict(blurred_batch)
|
| 104 |
|
| 105 |
+
real_labels = -np.ones((len(clear_batch), 1))
|
| 106 |
+
fake_labels = np.ones((len(fake_images), 1))
|
|
|
|
| 107 |
d_loss_real = discriminator.train_on_batch(clear_batch, real_labels)
|
| 108 |
d_loss_fake = discriminator.train_on_batch(fake_images, fake_labels)
|
| 109 |
d_loss = d_loss_real + d_loss_fake
|
| 110 |
|
| 111 |
+
misleading_labels = -np.ones((len(blurred_batch), 1))
|
|
|
|
| 112 |
g_loss = generator.train_on_batch(blurred_batch, misleading_labels)
|
| 113 |
|
| 114 |
print(f"Batch {i // batch_size + 1}: D Loss: {d_loss:.4f}, G Loss: {g_loss:.4f}")
|
| 115 |
|
|
|
|
| 116 |
blurred_folder = "blurred_sketches"
|
| 117 |
clear_folder = "clear_sketches"
|
| 118 |
blurred_images = load_images(blurred_folder)
|
| 119 |
clear_images = load_images(clear_folder)
|
| 120 |
|
|
|
|
| 121 |
generator = build_generator()
|
| 122 |
discriminator = build_discriminator()
|
| 123 |
|
|
|
|
| 124 |
discriminator.compile(optimizer=Adam(0.0002, 0.5, 0.999), loss=wasserstein_loss)
|
| 125 |
|
| 126 |
+
|
| 127 |
train(generator, discriminator, blurred_images, clear_images, epochs=500, batch_size=16)
|