{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction and Setup\n\nThis notebook utilizes a CycleGAN architecture to add Monet-style to photos. For this tutorial, we will be using the TFRecord dataset. Import the following packages and change the accelerator to TPU.\n\nFor more information, check out [TensorFlow](https://www.tensorflow.org/tutorials/generative/cyclegan) and [Keras](https://keras.io/examples/generative/cyclegan/) CycleGAN documentation pages.\n\nIt is based on the tutorial from [Amy Jang](https://www.kaggle.com/code/amyjang/monet-cyclegan-tutorial/notebook). The main addition is the usage of a 9 \"ResNet\" blocks generators and a decreasing learning rate, as is described in the CycleGAN paper.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nimport tensorflow_addons as tfa\n\nfrom kaggle_datasets import KaggleDatasets\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    print('Device:', tpu.master())\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nexcept:\n    strategy = tf.distribute.get_strategy()\nprint('Number of replicas:', strategy.num_replicas_in_sync)\n\nAUTOTUNE = tf.data.experimental.AUTOTUNE\n    \nprint(tf.__version__)","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.status.busy":"2022-07-06T18:16:18.579387Z","iopub.execute_input":"2022-07-06T18:16:18.579821Z","iopub.status.idle":"2022-07-06T18:16:18.589195Z","shell.execute_reply.started":"2022-07-06T18:16:18.57979Z","shell.execute_reply":"2022-07-06T18:16:18.588141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load in the data\n\nWe want to keep our photo dataset and our Monet dataset separate. First, load in the filenames of the TFRecords.","metadata":{}},{"cell_type":"code","source":"GCS_PATH = KaggleDatasets().get_gcs_path()","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:18.591126Z","iopub.execute_input":"2022-07-06T18:16:18.592074Z","iopub.status.idle":"2022-07-06T18:16:18.910955Z","shell.execute_reply.started":"2022-07-06T18:16:18.592037Z","shell.execute_reply":"2022-07-06T18:16:18.910005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MONET_FILENAMES = tf.io.gfile.glob(str(GCS_PATH + '/monet_tfrec/*.tfrec'))\nprint('Monet TFRecord Files:', len(MONET_FILENAMES))\n\nPHOTO_FILENAMES = tf.io.gfile.glob(str(GCS_PATH + '/photo_tfrec/*.tfrec'))\nprint('Photo TFRecord Files:', len(PHOTO_FILENAMES))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:18.913187Z","iopub.execute_input":"2022-07-06T18:16:18.913891Z","iopub.status.idle":"2022-07-06T18:16:19.117853Z","shell.execute_reply.started":"2022-07-06T18:16:18.91385Z","shell.execute_reply":"2022-07-06T18:16:19.116944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"All the images for the competition are already sized to 256x256. As these images are RGB images, set the channel to 3. Additionally, we need to scale the images to a [-1, 1] scale. Because we are building a generative model, we don't need the labels or the image id so we'll only return the image from the TFRecord.","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [256, 256]\n\ndef decode_image(image):\n    image = tf.image.decode_jpeg(image, channels=3)\n    image = (tf.cast(image, tf.float32) / 127.5) - 1\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image\n\ndef read_tfrecord(example):\n    tfrecord_format = {\n        \"image_name\": tf.io.FixedLenFeature([], tf.string),\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"target\": tf.io.FixedLenFeature([], tf.string)\n    }\n    example = tf.io.parse_single_example(example, tfrecord_format)\n    image = decode_image(example['image'])\n    return image","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:19.119728Z","iopub.execute_input":"2022-07-06T18:16:19.120077Z","iopub.status.idle":"2022-07-06T18:16:19.127694Z","shell.execute_reply.started":"2022-07-06T18:16:19.120042Z","shell.execute_reply":"2022-07-06T18:16:19.126687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define the function to extract the image from the files.","metadata":{}},{"cell_type":"code","source":"def load_dataset(filenames, labeled=True, ordered=False):\n    dataset = tf.data.TFRecordDataset(filenames)\n    dataset = dataset.map(read_tfrecord, num_parallel_calls=AUTOTUNE)\n    return dataset","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:19.13199Z","iopub.execute_input":"2022-07-06T18:16:19.132344Z","iopub.status.idle":"2022-07-06T18:16:19.139106Z","shell.execute_reply.started":"2022-07-06T18:16:19.132311Z","shell.execute_reply":"2022-07-06T18:16:19.137896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's load in our datasets.","metadata":{}},{"cell_type":"code","source":"monet_ds = load_dataset(MONET_FILENAMES, labeled=True).batch(1)\nphoto_ds = load_dataset(PHOTO_FILENAMES, labeled=True).batch(1)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:19.141229Z","iopub.execute_input":"2022-07-06T18:16:19.14163Z","iopub.status.idle":"2022-07-06T18:16:19.240271Z","shell.execute_reply.started":"2022-07-06T18:16:19.141596Z","shell.execute_reply":"2022-07-06T18:16:19.239294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_monet = next(iter(monet_ds))\nexample_photo = next(iter(photo_ds))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:19.241775Z","iopub.execute_input":"2022-07-06T18:16:19.242113Z","iopub.status.idle":"2022-07-06T18:16:19.86161Z","shell.execute_reply.started":"2022-07-06T18:16:19.242068Z","shell.execute_reply":"2022-07-06T18:16:19.860633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's  visualize a photo example and a Monet example.","metadata":{}},{"cell_type":"code","source":"plt.subplot(121)\nplt.title('Photo')\nplt.imshow(example_photo[0] * 0.5 + 0.5)\n\nplt.subplot(122)\nplt.title('Monet')\nplt.imshow(example_monet[0] * 0.5 + 0.5)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:19.865303Z","iopub.execute_input":"2022-07-06T18:16:19.865572Z","iopub.status.idle":"2022-07-06T18:16:20.18002Z","shell.execute_reply.started":"2022-07-06T18:16:19.865547Z","shell.execute_reply":"2022-07-06T18:16:20.179131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_CHANNELS = 3","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:20.182305Z","iopub.execute_input":"2022-07-06T18:16:20.182994Z","iopub.status.idle":"2022-07-06T18:16:20.187668Z","shell.execute_reply.started":"2022-07-06T18:16:20.182952Z","shell.execute_reply":"2022-07-06T18:16:20.186792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's build our generator!\n\nThe generator first downsamples the input image and then upsample while establishing long skip connections. Skip connections are a way to help bypass the vanishing gradient problem by concatenating the output of a layer to multiple layers instead of only one. Here we concatenate the output of the downsample layer to the upsample layer in a symmetrical fashion.","metadata":{}},{"cell_type":"code","source":"def _get_norm_layer(norm):\n    if norm == 'none':\n        return lambda: lambda x: x\n    elif norm == 'batch_norm':\n        return keras.layers.BatchNormalization\n    elif norm == 'instance_norm':\n        return tfa.layers.InstanceNormalization\n    elif norm == 'layer_norm':\n        return keras.layers.LayerNormalization\n\n\ndef ResnetGenerator(input_shape=(256, 256, 3),\n                    output_channels=3,\n                    dim=64,\n                    n_downsamplings=2,\n                    n_blocks=9,\n                    norm='instance_norm'):\n    Norm = _get_norm_layer(norm)\n\n    def _residual_block(x):\n        dim = x.shape[-1]\n        h = x\n\n        h = tf.pad(h, [[0, 0], [1, 1], [1, 1], [0, 0]], mode='REFLECT')\n        h = keras.layers.Conv2D(dim, 3, padding='valid', use_bias=False)(h)\n        h = Norm()(h)\n        h = tf.nn.relu(h)\n\n        h = tf.pad(h, [[0, 0], [1, 1], [1, 1], [0, 0]], mode='REFLECT')\n        h = keras.layers.Conv2D(dim, 3, padding='valid', use_bias=False)(h)\n        h = Norm()(h)\n\n        return keras.layers.add([x, h])\n\n    # 0\n    h = inputs = keras.Input(shape=input_shape)\n\n    # 1\n    h = tf.pad(h, [[0, 0], [3, 3], [3, 3], [0, 0]], mode='REFLECT')\n    h = keras.layers.Conv2D(dim, 7, padding='valid', use_bias=False)(h)\n    h = Norm()(h)\n    h = tf.nn.relu(h)\n\n    # 2\n    for _ in range(n_downsamplings):\n        dim *= 2\n        h = keras.layers.Conv2D(dim, 3, strides=2, padding='same', use_bias=False)(h)\n        h = Norm()(h)\n        h = tf.nn.relu(h)\n\n    # 3\n    for _ in range(n_blocks):\n        h = _residual_block(h)\n\n    # 4\n    for _ in range(n_downsamplings):\n        dim //= 2\n        h = keras.layers.Conv2DTranspose(dim, 3, strides=2, padding='same', use_bias=False)(h)\n        h = Norm()(h)\n        h = tf.nn.relu(h)\n\n    # 5\n    h = tf.pad(h, [[0, 0], [3, 3], [3, 3], [0, 0]], mode='REFLECT')\n    h = keras.layers.Conv2D(output_channels, 7, padding='valid')(h)\n    h = tf.tanh(h)\n\n    return keras.Model(inputs=inputs, outputs=h)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:20.189442Z","iopub.execute_input":"2022-07-06T18:16:20.189811Z","iopub.status.idle":"2022-07-06T18:16:20.209504Z","shell.execute_reply.started":"2022-07-06T18:16:20.189777Z","shell.execute_reply":"2022-07-06T18:16:20.208431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build the discriminator\n\nThe discriminator takes in the input image and classifies it as real or fake (generated). Instead of outputing a single node, the discriminator outputs a smaller 2D image with higher pixel values indicating a real classification and lower values indicating a fake classification.","metadata":{}},{"cell_type":"code","source":"def ConvDiscriminator(input_shape=(256, 256, 3),\n                      dim=64,\n                      n_downsamplings=3,\n                      norm='instance_norm'):\n    dim_ = dim\n    Norm = _get_norm_layer(norm)\n\n    # 0\n    h = inputs = keras.Input(shape=input_shape)\n\n    # 1\n    h = keras.layers.Conv2D(dim, 4, strides=2, padding='same')(h)\n    h = tf.nn.leaky_relu(h, alpha=0.2)\n\n    for _ in range(n_downsamplings - 1):\n        dim = min(dim * 2, dim_ * 8)\n        h = keras.layers.Conv2D(dim, 4, strides=2, padding='same', use_bias=False)(h)\n        h = Norm()(h)\n        h = tf.nn.leaky_relu(h, alpha=0.2)\n\n    # 2\n    dim = min(dim * 2, dim_ * 8)\n    h = keras.layers.Conv2D(dim, 4, strides=1, padding='same', use_bias=False)(h)\n    h = Norm()(h)\n    h = tf.nn.leaky_relu(h, alpha=0.2)\n\n    # 3\n    h = keras.layers.Conv2D(1, 4, strides=1, padding='same')(h)\n\n    return keras.Model(inputs=inputs, outputs=h)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:20.211289Z","iopub.execute_input":"2022-07-06T18:16:20.211663Z","iopub.status.idle":"2022-07-06T18:16:20.225607Z","shell.execute_reply.started":"2022-07-06T18:16:20.211627Z","shell.execute_reply":"2022-07-06T18:16:20.224724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Learning Rate Scheduler","metadata":{}},{"cell_type":"code","source":"class LinearDecay(keras.optimizers.schedules.LearningRateSchedule):\n    # if `step` < `step_decay`: use fixed learning rate\n    # else: linearly decay the learning rate to zero\n\n    def __init__(self, initial_learning_rate, total_steps, step_decay):\n        super(LinearDecay, self).__init__()\n        self._initial_learning_rate = initial_learning_rate\n        self._steps = total_steps\n        self._step_decay = step_decay\n        self.current_learning_rate = tf.Variable(initial_value=initial_learning_rate, trainable=False, dtype=tf.float32)\n\n    def __call__(self, step):\n        self.current_learning_rate.assign(tf.cond(\n            step >= self._step_decay,\n            true_fn=lambda: self._initial_learning_rate * (1 - 1 / (self._steps - self._step_decay) * (step - self._step_decay)),\n            false_fn=lambda: self._initial_learning_rate\n        ))\n        return self.current_learning_rate","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:20.22721Z","iopub.execute_input":"2022-07-06T18:16:20.227578Z","iopub.status.idle":"2022-07-06T18:16:20.239019Z","shell.execute_reply.started":"2022-07-06T18:16:20.227546Z","shell.execute_reply":"2022-07-06T18:16:20.237975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    #monet_generator = Generator() # transforms photos to Monet-esque paintings\n    #photo_generator = Generator() # transforms Monet paintings to be more like photos\n\n    #monet_discriminator = Discriminator() # differentiates real Monet paintings and generated Monet paintings\n    #photo_discriminator = Discriminator() # differentiates real photos and generated photos\n    \n    monet_generator = ResnetGenerator() # transforms photos to Monet-esque paintings\n    photo_generator = ResnetGenerator() # transforms Monet paintings to be more like photos\n\n    monet_discriminator = ConvDiscriminator() # differentiates real Monet paintings and generated Monet paintings\n    photo_discriminator = ConvDiscriminator() # differentiates real photos and generated photos\n    \n    \n    \n    \n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:20.240376Z","iopub.execute_input":"2022-07-06T18:16:20.240841Z","iopub.status.idle":"2022-07-06T18:16:21.877505Z","shell.execute_reply.started":"2022-07-06T18:16:20.240806Z","shell.execute_reply":"2022-07-06T18:16:21.876417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Since our generators are not trained yet, the generated Monet-esque photo does not show what is expected at this point.","metadata":{}},{"cell_type":"code","source":"to_monet = monet_generator(example_photo)\n\nplt.subplot(1, 2, 1)\nplt.title(\"Original Photo\")\nplt.imshow(example_photo[0] * 0.5 + 0.5)\n\nplt.subplot(1, 2, 2)\nplt.title(\"Monet-esque Photo\")\nplt.imshow(to_monet[0] * 0.5 + 0.5)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:21.883718Z","iopub.execute_input":"2022-07-06T18:16:21.884012Z","iopub.status.idle":"2022-07-06T18:16:22.330989Z","shell.execute_reply.started":"2022-07-06T18:16:21.883979Z","shell.execute_reply":"2022-07-06T18:16:22.330105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build the CycleGAN model\n\nWe will subclass a `tf.keras.Model` so that we can run `fit()` later to train our model. During the training step, the model transforms a photo to a Monet painting and then back to a photo. The difference between the original photo and the twice-transformed photo is the cycle-consistency loss. We want the original photo and the twice-transformed photo to be similar to one another.\n\nThe losses are defined in the next section.","metadata":{}},{"cell_type":"code","source":"class CycleGan(keras.Model):\n    def __init__(\n        self,\n        monet_generator,\n        photo_generator,\n        monet_discriminator,\n        photo_discriminator,\n        lambda_cycle=10,\n    ):\n        super(CycleGan, self).__init__()\n        self.m_gen = monet_generator\n        self.p_gen = photo_generator\n        self.m_disc = monet_discriminator\n        self.p_disc = photo_discriminator\n        self.lambda_cycle = lambda_cycle\n        \n    def compile(\n        self,\n        m_gen_optimizer,\n        p_gen_optimizer,\n        m_disc_optimizer,\n        p_disc_optimizer,\n        gen_loss_fn,\n        disc_loss_fn,\n        cycle_loss_fn,\n        identity_loss_fn\n    ):\n        super(CycleGan, self).compile()\n        self.m_gen_optimizer = m_gen_optimizer\n        self.p_gen_optimizer = p_gen_optimizer\n        self.m_disc_optimizer = m_disc_optimizer\n        self.p_disc_optimizer = p_disc_optimizer\n        self.gen_loss_fn = gen_loss_fn\n        self.disc_loss_fn = disc_loss_fn\n        self.cycle_loss_fn = cycle_loss_fn\n        self.identity_loss_fn = identity_loss_fn\n        \n    def train_step(self, batch_data):\n        real_monet, real_photo = batch_data\n        \n        with tf.GradientTape(persistent=True) as tape:\n            # photo to monet back to photo\n            fake_monet = self.m_gen(real_photo, training=True)\n            cycled_photo = self.p_gen(fake_monet, training=True)\n\n            # monet to photo back to monet\n            fake_photo = self.p_gen(real_monet, training=True)\n            cycled_monet = self.m_gen(fake_photo, training=True)\n\n            # generating itself\n            same_monet = self.m_gen(real_monet, training=True)\n            same_photo = self.p_gen(real_photo, training=True)\n\n            # discriminator used to check, inputing real images\n            disc_real_monet = self.m_disc(real_monet, training=True)\n            disc_real_photo = self.p_disc(real_photo, training=True)\n\n            # discriminator used to check, inputing fake images\n            disc_fake_monet = self.m_disc(fake_monet, training=True)\n            disc_fake_photo = self.p_disc(fake_photo, training=True)\n\n            # evaluates generator loss\n            monet_gen_loss = self.gen_loss_fn(disc_fake_monet)\n            photo_gen_loss = self.gen_loss_fn(disc_fake_photo)\n\n            # evaluates total cycle consistency loss\n            total_cycle_loss = self.cycle_loss_fn(real_monet, cycled_monet, self.lambda_cycle) + self.cycle_loss_fn(real_photo, cycled_photo, self.lambda_cycle)\n\n            # evaluates total generator loss\n            total_monet_gen_loss = monet_gen_loss + total_cycle_loss + self.identity_loss_fn(real_monet, same_monet, self.lambda_cycle)\n            total_photo_gen_loss = photo_gen_loss + total_cycle_loss + self.identity_loss_fn(real_photo, same_photo, self.lambda_cycle)\n\n            # evaluates discriminator loss\n            monet_disc_loss = self.disc_loss_fn(disc_real_monet, disc_fake_monet)\n            photo_disc_loss = self.disc_loss_fn(disc_real_photo, disc_fake_photo)\n\n        # Calculate the gradients for generator and discriminator\n        monet_generator_gradients = tape.gradient(total_monet_gen_loss,\n                                                  self.m_gen.trainable_variables)\n        photo_generator_gradients = tape.gradient(total_photo_gen_loss,\n                                                  self.p_gen.trainable_variables)\n\n        monet_discriminator_gradients = tape.gradient(monet_disc_loss,\n                                                      self.m_disc.trainable_variables)\n        photo_discriminator_gradients = tape.gradient(photo_disc_loss,\n                                                      self.p_disc.trainable_variables)\n\n        # Apply the gradients to the optimizer\n        self.m_gen_optimizer.apply_gradients(zip(monet_generator_gradients,\n                                                 self.m_gen.trainable_variables))\n\n        self.p_gen_optimizer.apply_gradients(zip(photo_generator_gradients,\n                                                 self.p_gen.trainable_variables))\n\n        self.m_disc_optimizer.apply_gradients(zip(monet_discriminator_gradients,\n                                                  self.m_disc.trainable_variables))\n\n        self.p_disc_optimizer.apply_gradients(zip(photo_discriminator_gradients,\n                                                  self.p_disc.trainable_variables))\n        \n        return {\n            \"monet_gen_loss\": total_monet_gen_loss,\n            \"photo_gen_loss\": total_photo_gen_loss,\n            \"monet_disc_loss\": monet_disc_loss,\n            \"photo_disc_loss\": photo_disc_loss\n        }","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.332769Z","iopub.execute_input":"2022-07-06T18:16:22.333137Z","iopub.status.idle":"2022-07-06T18:16:22.352134Z","shell.execute_reply.started":"2022-07-06T18:16:22.3331Z","shell.execute_reply":"2022-07-06T18:16:22.351201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define loss functions\n\nThe discriminator loss function below compares real images to a matrix of 1s and fake images to a matrix of 0s. The perfect discriminator will output all 1s for real images and all 0s for fake images. The discriminator loss outputs the average of the real and generated loss.","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    def discriminator_loss(real, generated):\n        real_loss = tf.keras.losses.BinaryCrossentropy(from_logits=True, reduction=tf.keras.losses.Reduction.NONE)(tf.ones_like(real), real)\n\n        generated_loss = tf.keras.losses.BinaryCrossentropy(from_logits=True, reduction=tf.keras.losses.Reduction.NONE)(tf.zeros_like(generated), generated)\n\n        total_disc_loss = real_loss + generated_loss\n\n        return total_disc_loss * 0.5","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.353752Z","iopub.execute_input":"2022-07-06T18:16:22.354619Z","iopub.status.idle":"2022-07-06T18:16:22.370269Z","shell.execute_reply.started":"2022-07-06T18:16:22.354584Z","shell.execute_reply":"2022-07-06T18:16:22.369193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The generator wants to fool the discriminator into thinking the generated image is real. The perfect generator will have the discriminator output only 1s. Thus, it compares the generated image to a matrix of 1s to find the loss.","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    def generator_loss(generated):\n        return tf.keras.losses.BinaryCrossentropy(from_logits=True, reduction=tf.keras.losses.Reduction.NONE)(tf.ones_like(generated), generated)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.371929Z","iopub.execute_input":"2022-07-06T18:16:22.372329Z","iopub.status.idle":"2022-07-06T18:16:22.383216Z","shell.execute_reply.started":"2022-07-06T18:16:22.372292Z","shell.execute_reply":"2022-07-06T18:16:22.382303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We want our original photo and the twice transformed photo to be similar to one another. Thus, we can calculate the cycle consistency loss be finding the average of their difference.","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    def calc_cycle_loss(real_image, cycled_image, LAMBDA):\n        loss1 = tf.reduce_mean(tf.abs(real_image - cycled_image))\n\n        return LAMBDA * loss1","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.384562Z","iopub.execute_input":"2022-07-06T18:16:22.384963Z","iopub.status.idle":"2022-07-06T18:16:22.393528Z","shell.execute_reply.started":"2022-07-06T18:16:22.384925Z","shell.execute_reply":"2022-07-06T18:16:22.392651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The identity loss compares the image with its generator (i.e. photo with photo generator). If given a photo as input, we want it to generate the same image as the image was originally a photo. The identity loss compares the input with the output of the generator.","metadata":{}},{"cell_type":"code","source":"with strategy.scope():\n    def identity_loss(real_image, same_image, LAMBDA):\n        loss = tf.reduce_mean(tf.abs(real_image - same_image))\n        return LAMBDA * 0.5 * loss","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.39501Z","iopub.execute_input":"2022-07-06T18:16:22.39545Z","iopub.status.idle":"2022-07-06T18:16:22.403112Z","shell.execute_reply.started":"2022-07-06T18:16:22.395413Z","shell.execute_reply":"2022-07-06T18:16:22.402198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train the CycleGAN\n\nLet's compile our model. Since we used `tf.keras.Model` to build our CycleGAN, we can just ude the `fit` function to train our model.","metadata":{}},{"cell_type":"code","source":"import re\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) for filename in filenames]\n    return np.sum(n)\n\nn_monet_samples = count_data_items(MONET_FILENAMES)\nn_photo_samples = count_data_items(PHOTO_FILENAMES)\n\nprint(f'Monet TFRecord files: {len(MONET_FILENAMES)}')\nprint(f'Monet image files: {n_monet_samples}')\nprint(f'Photo TFRecord files: {len(PHOTO_FILENAMES)}')\nprint(f'Photo image files: {n_photo_samples}')","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:17:43.725042Z","iopub.execute_input":"2022-07-06T18:17:43.725547Z","iopub.status.idle":"2022-07-06T18:17:43.749277Z","shell.execute_reply.started":"2022-07-06T18:17:43.725501Z","shell.execute_reply":"2022-07-06T18:17:43.746454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    #monet_generator_optimizer = tf.keras.optimizers.Adam(2e-4, beta_1=0.5)\n    #photo_generator_optimizer = tf.keras.optimizers.Adam(2e-4, beta_1=0.5)\n\n    #monet_discriminator_optimizer = tf.keras.optimizers.Adam(2e-4, beta_1=0.5)\n    #photo_discriminator_optimizer = tf.keras.optimizers.Adam(2e-4, beta_1=0.5)\n    \n    lr = 2e-4\n    beta_1 = 0.5\n    len_dataset = n_photo_samples\n    EPOCHS=200\n    \n    GA_lr_scheduler = LinearDecay(lr, EPOCHS * len_dataset, EPOCHS // 2 * len_dataset)\n    DA_lr_scheduler = LinearDecay(lr, EPOCHS * len_dataset, EPOCHS // 2 * len_dataset)\n    GB_lr_scheduler = LinearDecay(lr, EPOCHS * len_dataset, EPOCHS // 2 * len_dataset)\n    DB_lr_scheduler = LinearDecay(lr, EPOCHS * len_dataset, EPOCHS // 2 * len_dataset)\n    \n    monet_discriminator_optimizer = keras.optimizers.Adam(learning_rate=GA_lr_scheduler, beta_1=beta_1)\n    photo_discriminator_optimizer = keras.optimizers.Adam(learning_rate=DA_lr_scheduler, beta_1=beta_1)\n\n    monet_generator_optimizer = keras.optimizers.Adam(learning_rate=GB_lr_scheduler, beta_1=beta_1)\n    photo_generator_optimizer = keras.optimizers.Adam(learning_rate=DB_lr_scheduler, beta_1=beta_1)\n\n    \n","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:18:02.060842Z","iopub.execute_input":"2022-07-06T18:18:02.061272Z","iopub.status.idle":"2022-07-06T18:18:02.07319Z","shell.execute_reply.started":"2022-07-06T18:18:02.061239Z","shell.execute_reply":"2022-07-06T18:18:02.07217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    cycle_gan_model = CycleGan(\n        monet_generator, photo_generator, monet_discriminator, photo_discriminator\n    )\n\n    cycle_gan_model.compile(\n        m_gen_optimizer = monet_generator_optimizer,\n        p_gen_optimizer = photo_generator_optimizer,\n        m_disc_optimizer = monet_discriminator_optimizer,\n        p_disc_optimizer = photo_discriminator_optimizer,\n        gen_loss_fn = generator_loss,\n        disc_loss_fn = discriminator_loss,\n        cycle_loss_fn = calc_cycle_loss,\n        identity_loss_fn = identity_loss\n    )","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:18:05.905121Z","iopub.execute_input":"2022-07-06T18:18:05.905575Z","iopub.status.idle":"2022-07-06T18:18:05.930007Z","shell.execute_reply.started":"2022-07-06T18:18:05.905535Z","shell.execute_reply":"2022-07-06T18:18:05.928625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cycle_gan_model.fit(\n    tf.data.Dataset.zip((monet_ds, photo_ds)),\n    epochs=EPOCHS\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:18:08.175459Z","iopub.execute_input":"2022-07-06T18:18:08.175899Z","iopub.status.idle":"2022-07-06T19:14:23.713564Z","shell.execute_reply.started":"2022-07-06T18:18:08.175837Z","shell.execute_reply":"2022-07-06T19:14:23.712325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize our Monet-esque photos","metadata":{}},{"cell_type":"code","source":"_, ax = plt.subplots(5, 2, figsize=(12, 12))\nfor i, img in enumerate(photo_ds.take(5)):\n    prediction = monet_generator(img, training=False)[0].numpy()\n    prediction = (prediction * 127.5 + 127.5).astype(np.uint8)\n    img = (img[0] * 127.5 + 127.5).numpy().astype(np.uint8)\n\n    ax[i, 0].imshow(img)\n    ax[i, 1].imshow(prediction)\n    ax[i, 0].set_title(\"Input Photo\")\n    ax[i, 1].set_title(\"Monet-esque\")\n    ax[i, 0].axis(\"off\")\n    ax[i, 1].axis(\"off\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-06T19:39:06.131747Z","iopub.execute_input":"2022-07-06T19:39:06.132361Z","iopub.status.idle":"2022-07-06T19:39:07.235142Z","shell.execute_reply.started":"2022-07-06T19:39:06.132327Z","shell.execute_reply":"2022-07-06T19:39:07.233886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create submission file","metadata":{}},{"cell_type":"code","source":"import PIL\n! mkdir ../images","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.443041Z","iopub.status.idle":"2022-07-06T18:16:22.44373Z","shell.execute_reply.started":"2022-07-06T18:16:22.443475Z","shell.execute_reply":"2022-07-06T18:16:22.443497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 1\nfor img in photo_ds:\n    prediction = monet_generator(img, training=False)[0].numpy()\n    prediction = (prediction * 127.5 + 127.5).astype(np.uint8)\n    im = PIL.Image.fromarray(prediction)\n    im.save(\"../images/\" + str(i) + \".jpg\")\n    #print(i)\n    i += 1","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.444948Z","iopub.status.idle":"2022-07-06T18:16:22.445651Z","shell.execute_reply.started":"2022-07-06T18:16:22.445385Z","shell.execute_reply":"2022-07-06T18:16:22.445408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nshutil.make_archive(\"/kaggle/working/images\", 'zip', \"/kaggle/images\")","metadata":{"execution":{"iopub.status.busy":"2022-07-06T18:16:22.446921Z","iopub.status.idle":"2022-07-06T18:16:22.447609Z","shell.execute_reply.started":"2022-07-06T18:16:22.447354Z","shell.execute_reply":"2022-07-06T18:16:22.447376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}