{"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":"from pathlib import Path","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-03T09:17:14.536003Z","iopub.execute_input":"2022-11-03T09:17:14.53661Z","iopub.status.idle":"2022-11-03T09:17:14.543681Z","shell.execute_reply.started":"2022-11-03T09:17:14.536567Z","shell.execute_reply":"2022-11-03T09:17:14.541985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_images = list(Path(\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\").rglob(\"*.JPEG\"))","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:17:20.127375Z","iopub.execute_input":"2022-11-03T09:17:20.128662Z","iopub.status.idle":"2022-11-03T09:17:46.743603Z","shell.execute_reply.started":"2022-11-03T09:17:20.128599Z","shell.execute_reply":"2022-11-03T09:17:46.741427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nimport cv2\nimport numpy as np","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:18:12.143525Z","iopub.execute_input":"2022-11-03T09:18:12.144129Z","iopub.status.idle":"2022-11-03T09:18:12.548893Z","shell.execute_reply.started":"2022-11-03T09:18:12.144081Z","shell.execute_reply":"2022-11-03T09:18:12.547156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(data_images))","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:24:01.669711Z","iopub.execute_input":"2022-11-03T09:24:01.670273Z","iopub.status.idle":"2022-11-03T09:24:01.679451Z","shell.execute_reply.started":"2022-11-03T09:24:01.670234Z","shell.execute_reply":"2022-11-03T09:24:01.677425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 15000\nimage = Image.open(data_images[i])\nimage.save(\"sharp1.png\")\nimage","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:19:47.031623Z","iopub.execute_input":"2022-11-03T09:19:47.032109Z","iopub.status.idle":"2022-11-03T09:19:47.194848Z","shell.execute_reply.started":"2022-11-03T09:19:47.032073Z","shell.execute_reply":"2022-11-03T09:19:47.193445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nGRADIENT_PENALTY_WEIGHT = 10\n\nEPOCHS = 10\nMODEL_DIR = 'models'\nMODEL_C = MODEL_DIR + '/colorization.pt'\nMODEL_D = MODEL_DIR + '/discriminator.pt'\n\nTRAIN_PATH = 'sample/train'\nTEST_PATH = 'sample/test'\n\nOUTPUT_PATH = 'Output/'\n\nBATCH_SIZE = 2\nTEST_BATCH_SIZE = 4\n\n# -1 for no log\nCHECK_PER = 100\n\nLR = 2e-5\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:51:49.503951Z","iopub.execute_input":"2022-11-03T09:51:49.504544Z","iopub.status.idle":"2022-11-03T09:51:49.513381Z","shell.execute_reply.started":"2022-11-03T09:51:49.504503Z","shell.execute_reply":"2022-11-03T09:51:49.512119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport os\nimport pandas as pd\n\nimport configs.config as config\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torch.optim import Adam\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nfrom utils.utils import *\n\nfrom src.colorization import Colorization\nfrom src.ColorizeDataloader import ColorizeDataLoader\nfrom src.discriminator import Discriminator\n\n\ndef model(train_data, test_data, epochs, version=0.0):\n    \"\"\"\n    Create the model and train it for the given epochs\n    :param train_data: train data\n    :param test_data: test data\n    :param epochs: number of epochs\n    :param version: version of the model\n    \"\"\"\n    # data loader\n    train_dataloader = ColorizeDataLoader(train_data)\n    train_dataloader = DataLoader(\n        train_dataloader, batch_size=config.BATCH_SIZE,\n        shuffle=True, num_workers=2, drop_last=True)\n\n    test_dataloader = ColorizeDataLoader(test_data)\n    test_dataloader = DataLoader(\n        test_dataloader, batch_size=config.BATCH_SIZE,\n        shuffle=True, num_workers=2, drop_last=True)\n\n    # Load the discriminator model and the colorization model   \n    discriminator = Discriminator(input_size=224).to(config.DEVICE)\n    discriminator.apply(initialize_weights)\n    colorization_model = Colorization(input_size=224).to(config.DEVICE)\n    vgg_model_f = models.vgg16(pretrained=True).to(config.DEVICE)\n    vgg_model_f.requires_grad_(False)\n\n    # positive_real = torch.ones(size=(config.BATCH_SIZE, 1, 28, 28), requires_grad=True).to(config.DEVICE)\n    # negative_real = (-positive_real).to(config.DEVICE)\n\n    optimizer_g = Adam(\n        colorization_model.parameters(), lr=config.LR, betas=(0.5, 0.999)\n    )\n\n    optimizer_d = Adam(\n        discriminator.parameters(), lr=config.LR, betas=(0.5, 0.999)\n    )\n\n    # init loss function\n    KKLDivergence = nn.KLDivLoss()\n    MSE = nn.MSELoss()\n\n    for epoch in range(epochs):\n        print(f'EPOCH {epoch} / {epochs}')\n        print('-' * 30)\n\n        for idx, (trainL, trainAB, _) in enumerate(tqdm(train_dataloader)):\n            trainL = trainL.to(config.DEVICE)\n            trainAB = trainAB.to(config.DEVICE)\n\n            l_3 = torch.cat([trainL, trainL, trainL], dim=1)\n            pred_class_vgg = F.softmax(vgg_model_f(l_3))\n            # ----------------- Train the generator -----------------\n            optimizer_g.zero_grad()\n            pred_AB, pred_class_c = colorization_model(l_3)\n            pred_LAB_C = torch.cat([trainL, pred_AB], dim=1)\n            with torch.no_grad():\n                dis_C = discriminator(pred_LAB_C)\n            KLD_loss = KKLDivergence(\n                F.softmax(pred_class_c).detach().float(), \n                pred_class_vgg.detach().float()\n                ) * 0.003\n            MSE_loss = MSE(pred_AB.float(), trainAB.float())\n            W_loss = wasserstein_loss(dis_C, True) * 0.1\n            g_loss = KLD_loss + MSE_loss + W_loss\n\n            # ----------------- Train the discriminator -----------------\n            for param in discriminator.parameters():\n                param.requires_grad = True\n            optimizer_d.zero_grad()\n            pred_LAB_D = torch.cat([trainL, pred_AB], dim=1)\n            dis_pred = discriminator(pred_LAB_D)\n            dis_pred = dis_pred.mean()\n\n            true_LAB_D = torch.cat([trainL, trainAB], dim=1)\n            dis_true = discriminator(true_LAB_D)\n            dis_true = dis_true.mean()\n            \n            weights = torch.randn((trainAB.size(0),1,1,1), device=config.DEVICE)\n            averaged_samples = (weights * trainAB) + ((1 - weights) * pred_AB)\n            averaged_samples = torch.autograd.Variable(averaged_samples, requires_grad=True)\n            avg_img = torch.cat([trainL, averaged_samples], dim=1)\n            dis_avg = discriminator(avg_img)\n\n            W_loss_true = wasserstein_loss(dis_true, False)\n            W_loss_pred = wasserstein_loss(dis_pred, True)\n            gp_loss_avg = partial_gp_loss(dis_avg, averaged_samples, config.GRADIENT_PENALTY_WEIGHT)\n            d_loss = W_loss_true + W_loss_pred + gp_loss_avg\n            with torch.autograd.set_detect_anomaly(True):\n                g_loss.backward(retain_graph=True)\n                d_loss.backward()\n                optimizer_g.step()\n                optimizer_d.step()\n            # ----------------- Log the trainning process  -----------------\n            if config.CHECK_PER!=-1:\n                if idx % config.CHECK_PER == 0:\n                    print('\\n')\n                    print(f\"Epoch {epoch} - Batch {idx} - Loss G: {g_loss} - Loss D: {d_loss}\")\n    # create MODEL_DIR if not exist\n    if not os.path.exists(config.MODEL_DIR):\n        os.makedirs(config.MODEL_DIR)\n    torch.save(discriminator.state_dict(), config.MODEL_D)\n    torch.save(colorization_model.state_dict(), config.MODEL_C)\n\n\n\ndef train():\n    train_path = config.TRAIN_PATH\n    test_path = config.TEST_PATH\n    epochs = config.EPOCHS\n\n    print('Start training...')\n    print('-' * 30)\n\n    model(train_path, test_path, epochs)\n\n    print('-' * 30)\n    print('Training done!')\n    print('-' * 30)\n\n    print('Start testing...')\n    print('-' * 30)\n\n    test()\n    \n    print('Testing done!')\n\n    print('-' * 30)\n    print('All done!')\n\n\ndef sample_images(test_data, colorizationModel):\n    \"\"\"\n    Sample images after training\n        :param test_data: test data\n        :param colorizationModel: colorization model\n    \"\"\"\n    print('Sampling images')\n    for idx, (gray, ori_ab, _) in enumerate(tqdm(test_data)):\n        l_3 = torch.cat([gray, gray, gray], dim=1).to(config.DEVICE)\n        # torch required no grad\n        with torch.no_grad():\n            colored, _ = colorizationModel(l_3)\n\n        gray = gray.detach().cpu().numpy()\n        ori_ab = ori_ab.detach().cpu().numpy()\n        colored = colored.detach().cpu().numpy()\n        for i in range(config.BATCH_SIZE):\n            original_result_red = reconstruct(deprocess(gray)[i], deprocess(colored)[i])\n            #print('originalResult_red shape: ', original_result_red.shape)\n            cv2.imwrite(config.OUTPUT_PATH + str(idx) + '.png', original_result_red)\n    print('Sampling images done')\n\n\ndef test():\n    \"\"\"\n    Test the model\n    \"\"\"\n    path = config.MODEL_C\n    test_dataloader = ColorizeDataLoader(config.TEST_PATH)\n    test_dataloader = DataLoader(\n        test_dataloader, batch_size=config.BATCH_SIZE,\n        shuffle=True, num_workers=2, drop_last=True)\n    # test dataloader working correctly\n    for idx, (grey_img, color_img, original_images_shape) in enumerate(test_dataloader): \n        print(f\"{idx} / {len(test_dataloader)}\")\n        print(f\"gray shape: {grey_img.shape}\")\n        print(f\"ori_ab shape: {color_img.shape}\")\n        print(f\"{original_images_shape}\")\n        break\n    colorizationModel = Colorization(input_size=224).to(config.DEVICE)\n    if config.DEVICE == 'cpu':\n        colorizationModel.load_state_dict(torch.load(path, map_location=torch.device('cpu')))\n    else:\n        colorizationModel.load_state_dict(torch.load(path))\n    colorizationModel.eval()\n    sample_images(test_dataloader, colorizationModel)\n\ndef test_case_1(device):\n    print('test case 1 start')\n    positive_real = torch.ones(size=(config.BATCH_SIZE, 1))\n    negative_real = -positive_real\n    dummy_y = torch.zeros(size=(config.BATCH_SIZE, 1))\n    print(positive_real.shape)\n    print(negative_real.shape)\n    print(dummy_y.shape)\n    print('test case 1 complete')\n\n\ndef run_all_test_case(device):\n    print('test case 1')\n    test_case_1(device)\n\n\nclass Trainer:\n    def __init__(self):\n        self.device =config.DEVICE\n\n    @staticmethod\n    def train():\n        train()\n\n\nclass Tester:\n    def __init__(self):\n        self.device =config.DEVICE\n\n    def test(self):\n        test()","metadata":{"execution":{"iopub.status.busy":"2022-11-03T09:53:23.341962Z","iopub.execute_input":"2022-11-03T09:53:23.34269Z","iopub.status.idle":"2022-11-03T09:53:23.41795Z","shell.execute_reply.started":"2022-11-03T09:53:23.34264Z","shell.execute_reply":"2022-11-03T09:53:23.415531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}