{"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":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2022-05-07T03:56:35.407655Z","iopub.execute_input":"2022-05-07T03:56:35.407891Z","iopub.status.idle":"2022-05-07T03:56:35.425012Z","shell.execute_reply.started":"2022-05-07T03:56:35.407837Z","shell.execute_reply":"2022-05-07T03:56:35.424452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport os\nimport numpy as np\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nfrom glob import glob\nfrom tensorflow.python.framework import graph_util","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:16.801496Z","iopub.execute_input":"2022-05-07T01:21:16.802104Z","iopub.status.idle":"2022-05-07T01:21:16.989829Z","shell.execute_reply.started":"2022-05-07T01:21:16.802035Z","shell.execute_reply":"2022-05-07T01:21:16.98902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_h = 256\nimg_w = 256\nimg_shape = (img_h, img_w, 3)\ndirs = \"../input\"\ndataset_name = \"summer2winter-yosemite\"\nbatch_size = 1","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:16.99112Z","iopub.execute_input":"2022-05-07T01:21:16.991637Z","iopub.status.idle":"2022-05-07T01:21:16.996143Z","shell.execute_reply.started":"2022-05-07T01:21:16.991341Z","shell.execute_reply":"2022-05-07T01:21:16.99539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Data_Generator():\n    \"\"\"数据生成器\"\"\"\n    def __init__(self, dirs, name, img_shape, batch_size, is_testing=False):\n        \"\"\"初始化数据集参数\"\"\"\n        self.dirs = dirs\n        self.name = name\n        self.img_shape = img_shape\n        self.batch_size = batch_size\n        self.is_testing = is_testing\n        self.h, self.w = self.img_shape[0], self.img_shape[1]\n    \n    def img_read(self, path):\n        \"\"\"读取一张图片\"\"\"\n        img = cv2.imread(path)\n        img = img[...,::-1]\n        img = cv2.resize(img, (self.h, self.w), interpolation=cv2.INTER_LINEAR)\n        return img\n    \n    def get_img_paths(self, is_testing=False):\n        \"\"\"取得所有图片的路径\"\"\"\n        data_type = \"train\" if not is_testing else \"test\"\n        imgs_A_log = \"%s/%s/%sA/*\" % (self.dirs, self.name, data_type)\n        imgs_B_log = \"%s/%s/%sB/*\" % (self.dirs, self.name, data_type)\n        imgs_A_path = glob(imgs_A_log)\n        imgs_B_path = glob(imgs_B_log)\n        return imgs_A_path, imgs_B_path\n    \n    def load_batch_imgs(self):\n        \"\"\"加载一个batch的数据，作为生成器\"\"\"\n        A_paths, B_paths = self.get_img_paths(self.is_testing)\n        max_batch = min(len(A_paths), len(B_paths)) // self.batch_size\n        total = max_batch * self.batch_size\n        A_paths = np.random.choice(A_paths, total, replace=False)\n        B_paths = np.random.choice(B_paths, total, replace=False)\n        A_paths = np.reshape(A_paths, (-1, self.batch_size))\n        B_paths = np.reshape(B_paths, (-1, self.batch_size))\n        for batch_imgs_A, batch_imgs_B in zip(A_paths, B_paths):\n            imgs_A, imgs_B = [], []\n            for img_A, img_B in zip(batch_imgs_A, batch_imgs_B):\n                imgs_A.append(self.img_read(img_A))\n                imgs_B.append(self.img_read(img_B))\n            imgs_A = np.array(imgs_A, dtype=np.float32) / 127.5 -1.\n            imgs_B = np.array(imgs_B, dtype=np.float32) / 127.5 -1.\n            yield imgs_A, imgs_B\n            \n    def load_data(self):\n        A_paths, B_paths = self.get_img_paths(self.is_testing)\n        A_path = np.random.choice(A_paths, 1, replace=False)\n        B_path = np.random.choice(B_paths, 1, replace=False)\n        img_A = self.img_read(A_path[0]).astype(np.float32) / 127.5 -1\n        img_B = self.img_read(B_path[0]).astype(np.float32) / 127.5 -1\n        return img_A, img_B","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:16.997969Z","iopub.execute_input":"2022-05-07T01:21:16.998488Z","iopub.status.idle":"2022-05-07T01:21:17.014882Z","shell.execute_reply.started":"2022-05-07T01:21:16.99844Z","shell.execute_reply":"2022-05-07T01:21:17.014139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_generator(inputs, scope, reuse=False):\n    \"\"\"创建A到B的生成网络\"\"\"\n    def residual_block(inputs, filters, kernel_size):\n        \"\"\"残差块\"\"\"\n        pad = kernel_size // 2\n        x = tf.pad(inputs, [[0, 0], [pad, pad], [pad, pad], [0, 0]])\n        x = tf.layers.conv2d(x, filters, kernel_size, 1)\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n\n        x = tf.pad(inputs, [[0, 0], [pad, pad], [pad, pad], [0, 0]])\n        x = tf.layers.conv2d(x, filters, kernel_size, 1)\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.add(x, inputs)\n        x = tf.nn.relu(x)\n        return x\n    \n    with tf.variable_scope(scope, reuse=reuse):\n        x = tf.pad(inputs, [[0, 0], [3, 3], [3, 3], [0, 0]])\n        x = tf.layers.conv2d(x, 64, 7, 1)\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n        \n        # 下采样\n        x = tf.pad(x, [[0, 0], [1, 1], [1, 1], [0, 0]])\n        x = tf.layers.conv2d(x, 128, 3, 2)\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n        \n        x = tf.pad(x, [[0, 0], [1, 1], [1, 1], [0, 0]])\n        x = tf.layers.conv2d(x, 256, 3, 2)\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n        \n        # 残差部分\n        for i in range(9):\n            x = residual_block(x, 256, 3)\n            \n        # 上采样\n        x = tf.layers.conv2d_transpose(x, filters=128, kernel_size=3, strides=2, padding=\"same\")\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n        \n        x = tf.layers.conv2d_transpose(x, filters=64, kernel_size=3, strides=2, padding=\"same\")\n        x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.relu(x)\n        \n        x = tf.pad(x, [[0, 0], [3, 3], [3, 3], [0, 0]])\n        x = tf.layers.conv2d(x, 3, 7, 1)\n        x = tf.nn.tanh(x, name=\"generator\")\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:17.016835Z","iopub.execute_input":"2022-05-07T01:21:17.017118Z","iopub.status.idle":"2022-05-07T01:21:17.034517Z","shell.execute_reply.started":"2022-05-07T01:21:17.01705Z","shell.execute_reply":"2022-05-07T01:21:17.033657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_discriminator(inputs, scope, reuse):\n    \"\"\"创建判别模型\"\"\"\n    def conv2d(inputs, filters, kernel_size, strides, normal=True):\n        x = tf.layers.conv2d(inputs, filters, kernel_size, strides, padding=\"same\")\n        if normal:\n            x = tf.contrib.layers.instance_norm(x)\n        x = tf.nn.leaky_relu(x)\n        return x\n    with tf.variable_scope(scope, reuse=reuse):\n        x = conv2d(inputs, 64, 4, 2, False)\n        x = conv2d(x, 128, 4, 2)\n        x = conv2d(x, 256, 4, 2)\n        x = conv2d(x, 512, 4, 2)\n        x = tf.layers.conv2d(x, 1, 3, 1, padding=\"same\")\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:17.036148Z","iopub.execute_input":"2022-05-07T01:21:17.036711Z","iopub.status.idle":"2022-05-07T01:21:17.045636Z","shell.execute_reply.started":"2022-05-07T01:21:17.036662Z","shell.execute_reply":"2022-05-07T01:21:17.044742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_A = tf.placeholder(dtype=tf.float32, shape=(None,)+img_shape, name=\"X_A\")  # A域图片占位符\nX_B = tf.placeholder(dtype=tf.float32, shape=(None,)+img_shape, name=\"X_B\")  # B域图片占位符\nfake_B = get_generator(X_A, \"gen_A2B\", False)  # 得到A2B的数据流图\nfake_A = get_generator(X_B, \"gen_B2A\", False)  # 得到B2A的数据流图\nre_A = get_generator(fake_B, \"gen_B2A\", True)  # 从fakeB重构为A\nre_B = get_generator(fake_A, \"gen_A2B\", True)  # 从fakeA重构为B\nid_A = get_generator(X_A, \"gen_B2A\", True)  # 要求A域不会被B2A改变\nid_B = get_generator(X_B, \"gen_A2B\", True)  # 要求B域不会被A2B改变\nd_fake_A = get_discriminator(fake_A, \"d_A\", False)  # 辨别fakeA\nd_fake_B = get_discriminator(fake_B, \"d_B\", False)  # 辨别fakeB\nd_real_A = get_discriminator(X_A, \"d_A\", True)  # 辨别真实A图片\nd_real_B = get_discriminator(X_B, \"d_B\", True)  # 辨别真实B图片","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:17.04692Z","iopub.execute_input":"2022-05-07T01:21:17.047616Z","iopub.status.idle":"2022-05-07T01:21:28.939746Z","shell.execute_reply.started":"2022-05-07T01:21:17.047172Z","shell.execute_reply":"2022-05-07T01:21:28.938924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lambda_cycle = 5\nlambda_id = 2.5\nlambda_d = 0.5","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.941Z","iopub.execute_input":"2022-05-07T01:21:28.941278Z","iopub.status.idle":"2022-05-07T01:21:28.947253Z","shell.execute_reply.started":"2022-05-07T01:21:28.941232Z","shell.execute_reply":"2022-05-07T01:21:28.946426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def genA2B_loss(re_B, id_B, X_B, d_fake_B):\n    \"\"\"genA2B损失函数\"\"\"\n    valid_B = np.ones(shape=(batch_size,)+(16, 16, 1))\n    re_loss = tf.reduce_mean(tf.abs(re_B-X_B))\n    id_loss = tf.reduce_mean(tf.abs(id_B-X_B))\n    fake_loss = tf.reduce_mean(tf.square(d_fake_B-valid_B))\n    return lambda_cycle * re_loss + lambda_id * id_loss + lambda_d * fake_loss","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.94861Z","iopub.execute_input":"2022-05-07T01:21:28.949057Z","iopub.status.idle":"2022-05-07T01:21:28.956211Z","shell.execute_reply.started":"2022-05-07T01:21:28.948838Z","shell.execute_reply":"2022-05-07T01:21:28.955288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def genB2A_loss(re_A, id_A, X_A, d_fake_A):\n    \"\"\"genB2A损失函数\"\"\"\n    valid_A = np.ones(shape=(batch_size,)+(16, 16, 1))\n    re_loss = tf.reduce_mean(tf.abs(re_A-X_A))\n    id_loss = tf.reduce_mean(tf.abs(id_A-X_A))\n    fake_loss = tf.reduce_mean(tf.square(d_fake_A-valid_A))\n    return lambda_cycle * re_loss + lambda_id * id_loss + lambda_d * fake_loss","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.957485Z","iopub.execute_input":"2022-05-07T01:21:28.958047Z","iopub.status.idle":"2022-05-07T01:21:28.964208Z","shell.execute_reply.started":"2022-05-07T01:21:28.957998Z","shell.execute_reply":"2022-05-07T01:21:28.963609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def discriminator_A_loss(d_fake_A, d_real_A):\n    \"\"\"d_A损失函数\"\"\"\n    valid = np.ones(shape=(batch_size,)+(16, 16, 1))\n    fake = np.zeros(shape=(batch_size,)+(16, 16, 1))\n    fake_loss = tf.reduce_mean(tf.square(d_fake_A-fake))\n    real_loss = tf.reduce_mean(tf.square(d_real_A-valid))\n    return fake_loss + real_loss","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.965537Z","iopub.execute_input":"2022-05-07T01:21:28.966127Z","iopub.status.idle":"2022-05-07T01:21:28.97581Z","shell.execute_reply.started":"2022-05-07T01:21:28.966052Z","shell.execute_reply":"2022-05-07T01:21:28.974944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def discriminator_B_loss(d_fake_B, d_real_B):\n    \"\"\"d_A损失函数\"\"\"\n    valid = np.ones(shape=(batch_size,)+(16, 16, 1))\n    fake = np.zeros(shape=(batch_size,)+(16, 16, 1))\n    fake_loss = tf.reduce_mean(tf.square(d_fake_B-fake))\n    real_loss = tf.reduce_mean(tf.square(d_real_B-valid))\n    return fake_loss + real_loss","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.977108Z","iopub.execute_input":"2022-05-07T01:21:28.9774Z","iopub.status.idle":"2022-05-07T01:21:28.983774Z","shell.execute_reply.started":"2022-05-07T01:21:28.977339Z","shell.execute_reply":"2022-05-07T01:21:28.982916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"A2B_loss = genA2B_loss(re_B, id_B, X_B, d_fake_B)\nB2A_loss = genB2A_loss(re_A, id_A, X_A, d_fake_A)\nD_A_loss = discriminator_A_loss(d_fake_A, d_real_A)\nD_B_loss = discriminator_B_loss(d_fake_B, d_real_B)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:28.986215Z","iopub.execute_input":"2022-05-07T01:21:28.986718Z","iopub.status.idle":"2022-05-07T01:21:29.027611Z","shell.execute_reply.started":"2022-05-07T01:21:28.986666Z","shell.execute_reply":"2022-05-07T01:21:29.026952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optmizier = tf.train.AdamOptimizer(0.0002, 0.5)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:29.028746Z","iopub.execute_input":"2022-05-07T01:21:29.029206Z","iopub.status.idle":"2022-05-07T01:21:29.033104Z","shell.execute_reply.started":"2022-05-07T01:21:29.029159Z","shell.execute_reply":"2022-05-07T01:21:29.032383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"D_A_train_op = optmizier.minimize(D_A_loss, var_list=[var for var in tf.trainable_variables() if var.name.startswith(\"d_A\")])\nD_B_train_op = optmizier.minimize(D_B_loss, var_list=[var for var in tf.trainable_variables() if var.name.startswith(\"d_B\")])\nA2B_train_op = optmizier.minimize(A2B_loss, var_list=[var for var in tf.trainable_variables() if var.name.startswith(\"gen_A2B\")])\nB2A_train_op = optmizier.minimize(B2A_loss, var_list=[var for var in tf.trainable_variables() if var.name.startswith(\"gen_B2A\")])","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:29.034339Z","iopub.execute_input":"2022-05-07T01:21:29.034829Z","iopub.status.idle":"2022-05-07T01:21:43.420004Z","shell.execute_reply.started":"2022-05-07T01:21:29.034779Z","shell.execute_reply":"2022-05-07T01:21:43.419197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sample_images(epoch, img_A, img_B, gen_A, gen_B):\n        os.makedirs('trans_images/%s' % dataset_name, exist_ok=True)\n        r, c = 2, 2\n        \n        img_A = np.expand_dims(img_A, axis=0)\n        img_B = np.expand_dims(img_B, axis=0)\n        gen_imgs = np.concatenate([img_A, gen_B, img_B, gen_A])\n        gen_imgs = 0.5 * gen_imgs + 0.5\n\n        titles = ['Original', 'Translated', 'Reconstructed']\n        fig, axs = plt.subplots(r, c)\n        cnt = 0\n        for i in range(r):\n            for j in range(c):\n                axs[i,j].imshow(gen_imgs[cnt])\n                axs[i, j].set_title(titles[j])\n                axs[i,j].axis('off')\n                cnt += 1\n        fig.savefig(\"trans_images/%s/%d.png\" % (dataset_name, epoch))\n        plt.close()","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:43.421392Z","iopub.execute_input":"2022-05-07T01:21:43.421675Z","iopub.status.idle":"2022-05-07T01:21:43.429971Z","shell.execute_reply.started":"2022-05-07T01:21:43.421628Z","shell.execute_reply":"2022-05-07T01:21:43.429181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = tf.ConfigProto()\nconfig.gpu_options.allow_growth = True","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:43.431276Z","iopub.execute_input":"2022-05-07T01:21:43.431859Z","iopub.status.idle":"2022-05-07T01:21:43.440188Z","shell.execute_reply.started":"2022-05-07T01:21:43.431572Z","shell.execute_reply":"2022-05-07T01:21:43.439309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with tf.Session(config=config) as sess:\n    epochs = 41\n    data_loader = Data_Generator(dirs, dataset_name, img_shape, batch_size, False)\n    sess.run(tf.global_variables_initializer())\n    for epoch in range(epochs):\n        dA_losses, dB_losses, genA2B_losses, genB2A_losses = [], [], [], []\n        for i, (img_A, img_B) in enumerate(data_loader.load_batch_imgs()):\n            feed_dict = {X_A: img_A, X_B: img_B}\n            _, genA2B_loss = sess.run([A2B_train_op, A2B_loss], feed_dict=feed_dict)\n            _, genB2A_loss = sess.run([B2A_train_op, B2A_loss], feed_dict=feed_dict)\n            _, dA_loss = sess.run([D_A_train_op, D_A_loss], feed_dict=feed_dict)\n            _, dB_loss = sess.run([D_B_train_op, D_B_loss], feed_dict=feed_dict)\n            dA_losses.append(dA_loss)\n            dB_losses.append(dB_loss)\n            genA2B_losses.append(genA2B_loss)\n            genB2A_losses.append(genB2A_loss)\n            if i % 50 == 0:\n                str_log = \"Epoch: %d \\t Batch: %d \\t d_loss: %f \\t g_loss: %f\"\n                print(str_log % (epoch, i, np.mean(dA_losses+dB_losses)/2, np.mean(genA2B_losses+genB2A_losses)/2))\n        img_A, img_B = Data_Generator(dirs, dataset_name, img_shape, 1, False).load_data()\n        gen_B = sess.run(fake_B, feed_dict={X_A: np.expand_dims(img_A, axis=0)})\n        gen_A = sess.run(fake_A, feed_dict={X_B: np.expand_dims(img_B, axis=0)})\n        sample_images(epoch, img_A, img_B, gen_A, gen_B)\n        \n        if epoch % 5 == 0 and epoch !=0:\n            pb_file_path = \"train_models/epoch_%d\" % epoch\n            os.makedirs(pb_file_path, exist_ok=True)\n            with tf.gfile.FastGFile(pb_file_path+'/A2B.pb', mode='wb') as f:\n                A2B_graph = graph_util.convert_variables_to_constants(sess, sess.graph_def, ['gen_A2B/generator'])\n                f.write(A2B_graph.SerializeToString())\n            with tf.gfile.FastGFile(pb_file_path+'/B2A.pb', mode='wb') as f:\n                B2A_graph = graph_util.convert_variables_to_constants(sess, sess.graph_def, ['gen_B2A/generator'])\n                f.write(B2A_graph.SerializeToString())","metadata":{"execution":{"iopub.status.busy":"2022-05-07T01:21:43.441495Z","iopub.execute_input":"2022-05-07T01:21:43.442056Z","iopub.status.idle":"2022-05-07T01:26:18.679548Z","shell.execute_reply.started":"2022-05-07T01:21:43.442007Z","shell.execute_reply":"2022-05-07T01:26:18.678318Z"},"trusted":true},"execution_count":null,"outputs":[]}]}