{"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 numpy as np\nimport cv2\nimport os\nimport shutil\nimport torch\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"device = \", device)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-31T00:08:17.159073Z","iopub.execute_input":"2022-07-31T00:08:17.159461Z","iopub.status.idle":"2022-07-31T00:08:19.364397Z","shell.execute_reply.started":"2022-07-31T00:08:17.159385Z","shell.execute_reply":"2022-07-31T00:08:19.363386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Cycle GAN model with MMGeneration\n\nMMGeneration is https://mmgeneration.readthedocs.io/en/latest/tutorials/customize_runtime.html\n\nMMGeneration have many GAN models.\n","metadata":{}},{"cell_type":"markdown","source":"## Install MMCV\n","metadata":{}},{"cell_type":"code","source":"# install mmcv-full\n#!pip install mmcv-full==1.4.8 -f https://download.openmmlab.com/mmcv/dist/cu110/torch1.7.0/index.html\n!pip install /kaggle/input/mmcv-full-151-cu110/mmcv_full-1.5.1-cp37-cp37m-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:19.367303Z","iopub.execute_input":"2022-07-31T00:08:19.367777Z","iopub.status.idle":"2022-07-31T00:08:31.531222Z","shell.execute_reply.started":"2022-07-31T00:08:19.367735Z","shell.execute_reply":"2022-07-31T00:08:31.530080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Install MMGeneration","metadata":{}},{"cell_type":"code","source":"shutil.copytree('/kaggle/input/mmgeneration/', '/kaggle/mmgeneration')\n\n%cd /kaggle/mmgeneration/mmgeneration\n!pip install -v -e .  # or \"python setup.py develop\"","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:31.534205Z","iopub.execute_input":"2022-07-31T00:08:31.534610Z","iopub.status.idle":"2022-07-31T00:08:53.922878Z","shell.execute_reply.started":"2022-07-31T00:08:31.534567Z","shell.execute_reply":"2022-07-31T00:08:53.921675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creage config file","metadata":{}},{"cell_type":"code","source":"from mmcv import Config\ncfg = Config.fromfile('./configs/cyclegan/cyclegan_lsgan_id0_resnet_in_facades_b1x1_80k.py')\n","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:53.926044Z","iopub.execute_input":"2022-07-31T00:08:53.926677Z","iopub.status.idle":"2022-07-31T00:08:55.240778Z","shell.execute_reply.started":"2022-07-31T00:08:53.926628Z","shell.execute_reply":"2022-07-31T00:08:55.239539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Copy train data\n\nBecause train folder must be \"trainA\" and \"trainB\", I copy train data to output folder.\n\n","metadata":{}},{"cell_type":"code","source":"!mkdir /kaggle/data","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:55.242209Z","iopub.execute_input":"2022-07-31T00:08:55.242634Z","iopub.status.idle":"2022-07-31T00:08:56.062059Z","shell.execute_reply.started":"2022-07-31T00:08:55.242588Z","shell.execute_reply":"2022-07-31T00:08:56.060814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working_dir","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:56.065052Z","iopub.execute_input":"2022-07-31T00:08:56.065461Z","iopub.status.idle":"2022-07-31T00:08:56.738348Z","shell.execute_reply.started":"2022-07-31T00:08:56.065419Z","shell.execute_reply":"2022-07-31T00:08:56.737169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.copytree('/kaggle/input/gan-getting-started/monet_jpg/', '/kaggle/data/trainA')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:56.740429Z","iopub.execute_input":"2022-07-31T00:08:56.740830Z","iopub.status.idle":"2022-07-31T00:08:58.302155Z","shell.execute_reply.started":"2022-07-31T00:08:56.740788Z","shell.execute_reply":"2022-07-31T00:08:58.301150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.copytree('/kaggle/input/gan-getting-started/photo_jpg/', '/kaggle/data/trainB')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:08:58.303850Z","iopub.execute_input":"2022-07-31T00:08:58.304187Z","iopub.status.idle":"2022-07-31T00:09:42.422784Z","shell.execute_reply.started":"2022-07-31T00:08:58.304154Z","shell.execute_reply":"2022-07-31T00:09:42.421794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's have a look at the final config used for training\nprint(f'Config:\\n{cfg.pretty_text}')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:09:42.424279Z","iopub.execute_input":"2022-07-31T00:09:42.424624Z","iopub.status.idle":"2022-07-31T00:09:42.891854Z","shell.execute_reply.started":"2022-07-31T00:09:42.424590Z","shell.execute_reply":"2022-07-31T00:09:42.889812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from mmgen.apis import set_random_seed\n\n# データのパス\ncfg.data.train.dataroot = \"/kaggle/data/\"\ncfg.data.test.dataroot = \"/kaggle/data/\"\ncfg.data.val.dataroot = \"/kaggle/data/\"\ncfg.gpu_ids = range(0, 1)\ncfg.seed = 123\n\ncfg.work_dir = '/kaggle/woking_dir/'\ncfg.total_iters = 40000\ncfg.lr_config.start =20000\n\nprint(f'Config:\\n{cfg.pretty_text}')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:09:42.895648Z","iopub.execute_input":"2022-07-31T00:09:42.896249Z","iopub.status.idle":"2022-07-31T00:09:45.142120Z","shell.execute_reply.started":"2022-07-31T00:09:42.896208Z","shell.execute_reply":"2022-07-31T00:09:45.141080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"import argparse\nimport copy\nimport multiprocessing as mp\nimport os\nimport os.path as osp\nimport platform\nimport time\nimport warnings\n\nimport cv2\nimport mmcv\nimport torch\nfrom mmcv import Config, DictAction\nfrom mmcv.runner import get_dist_info, init_dist\nfrom mmcv.utils import get_git_hash\n\nfrom mmgen import __version__\nfrom mmgen.apis import set_random_seed, train_model\nfrom mmgen.datasets import build_dataset\nfrom mmgen.models import build_model\nfrom mmgen.utils import collect_env, get_root_logger","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:09:45.143410Z","iopub.execute_input":"2022-07-31T00:09:45.144310Z","iopub.status.idle":"2022-07-31T00:09:45.152315Z","shell.execute_reply.started":"2022-07-31T00:09:45.144274Z","shell.execute_reply":"2022-07-31T00:09:45.151054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_model(\n    cfg.model, train_cfg=cfg.train_cfg, test_cfg=cfg.test_cfg)\n\ndatasets = [build_dataset(cfg.data.train)]\n\ntimestamp = time.strftime('%Y%m%d_%H%M%S', time.localtime())\n\nmeta = dict()\n# log env info\nenv_info_dict = collect_env()\nenv_info = '\\n'.join([(f'{k}: {v}') for k, v in env_info_dict.items()])\ndash_line = '-' * 60 + '\\n'\n\nmeta['env_info'] = env_info\nmeta['config'] = cfg.pretty_text\n\ntrain_model(\n    model,\n    datasets,\n    cfg,\n    distributed=False,\n    timestamp=timestamp,\n    meta=meta)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:09:45.153523Z","iopub.execute_input":"2022-07-31T00:09:45.153903Z","iopub.status.idle":"2022-07-31T00:30:39.410265Z","shell.execute_reply.started":"2022-07-31T00:09:45.153868Z","shell.execute_reply":"2022-07-31T00:30:39.409043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Output","metadata":{}},{"cell_type":"code","source":"cfg2 = cfg\n\nfrom mmgen.datasets.pipelines import Compose\nfrom mmgen.models import BaseTranslationModel\n\nfrom mmcv.parallel import collate, scatter\nfrom mmcv.runner import load_checkpoint\nfrom mmcv.utils import is_list_of\n\ndef sample_img2img_model2(model, image_path, target_domain=None, **kwargs):\n    \"\"\"Sampling from translation models.\n\n    Args:\n        model (nn.Module): The loaded model.\n        image_path (str): File path of input image.\n        style (str): Target style of output image.\n    Returns:\n        Tensor: Translated image tensor.\n    \"\"\"\n    assert isinstance(model, BaseTranslationModel)\n\n    # get source domain and target domain\n    if target_domain is None:\n        target_domain = model._default_domain\n    source_domain = model.get_other_domains(target_domain)[0]\n\n    #cfg = model._cfg\n    cfg = cfg2\n    device = next(model.parameters()).device  # model device\n    # build the data pipeline\n    test_pipeline = Compose(cfg.test_pipeline)\n\n    # prepare data\n    data = dict()\n    # dirty code to deal with test data pipeline\n    data['pair_path'] = image_path\n    data[f'img_{source_domain}_path'] = image_path\n    data[f'img_{target_domain}_path'] = image_path\n\n    data = test_pipeline(data)\n    if device.type == 'cpu':\n        data = collate([data], samples_per_gpu=1)\n        data['meta'] = []\n    else:\n        data = scatter(collate([data], samples_per_gpu=1), [device])[0]\n\n    source_image = data[f'img_{source_domain}']\n    # forward the model\n    with torch.no_grad():\n        results = model(\n            source_image,\n            test_mode=True,\n            target_domain=target_domain,\n            **kwargs)\n    output = results['target']\n    return output","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:30:39.414676Z","iopub.execute_input":"2022-07-31T00:30:39.415069Z","iopub.status.idle":"2022-07-31T00:30:39.493993Z","shell.execute_reply.started":"2022-07-31T00:30:39.415027Z","shell.execute_reply":"2022-07-31T00:30:39.492789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ndel datasets\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:30:39.497428Z","iopub.execute_input":"2022-07-31T00:30:39.497934Z","iopub.status.idle":"2022-07-31T00:30:39.713600Z","shell.execute_reply.started":"2022-07-31T00:30:39.497900Z","shell.execute_reply":"2022-07-31T00:30:39.712476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from glob import glob\n\ntest_folder = \"/kaggle/data/trainB\"\n\ntest_images = glob(test_folder + \"/*.jpg\")\n\nm = len(test_images)\n\n!mkdir /kaggle/images_origin\n\nfor i,image_path in enumerate(test_images):\n    \n    translated_image = sample_img2img_model2(model, image_path, target_domain='mask')\n    translate_image = translated_image.cpu().numpy()[0]\n    translate_image = translate_image.transpose(1,2,0)\n    fname = image_path.split('/')[-1]\n    translate_image = (((translate_image - translate_image.min()) * 255)/ (translate_image.max() - translate_image.min())).astype(np.uint8)\n    \n    cv2.imwrite(\"/kaggle/images_origin/\" + fname,translate_image)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:30:39.715080Z","iopub.execute_input":"2022-07-31T00:30:39.715469Z","iopub.status.idle":"2022-07-31T00:33:22.493339Z","shell.execute_reply.started":"2022-07-31T00:30:39.715430Z","shell.execute_reply":"2022-07-31T00:33:22.492187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"%cd /kaggle/working","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:33:22.495225Z","iopub.execute_input":"2022-07-31T00:33:22.495925Z","iopub.status.idle":"2022-07-31T00:33:22.504157Z","shell.execute_reply.started":"2022-07-31T00:33:22.495882Z","shell.execute_reply":"2022-07-31T00:33:22.503175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nshutil.make_archive(\"/kaggle/working/images\", 'zip', \"/kaggle/images_origin\")","metadata":{"execution":{"iopub.status.busy":"2022-07-31T00:33:22.505910Z","iopub.execute_input":"2022-07-31T00:33:22.506617Z","iopub.status.idle":"2022-07-31T00:33:32.347318Z","shell.execute_reply.started":"2022-07-31T00:33:22.506579Z","shell.execute_reply":"2022-07-31T00:33:32.346391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}