{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch.nn import Parameter, functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations\nimport multiprocessing\nimport os\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\n\n\nuse_tpu = False\n\nif use_tpu:\n    VERSION = \"20200325\"  #@param [\"1.5\" , \"20200325\", \"nightly\"]\n    !curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n    !pip install torchvision\n    !pip install torch==1.4.0\n    !pip install torchaudio==0.4.0\n    %matplotlib inline\n    !python pytorch-xla-env-setup.py --version $VERSION\n    import torch_xla\n    import torch_xla.core.xla_model as xm\n    import torch_xla.distributed.xla_multiprocessing as xmp\n    import torch_xla.distributed.parallel_loader as pl\n\n    os.environ['XLA_USE_BF16']= \"1\"\n    os.environ['XLA_TENSOR_ALLOCATOR_MAXSIZE'] = \"100000000\"\nelse:\n    torch.backends.cudnn.benchmark = True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"image_paths = os.listdir(\"../input/flickrfaceshq-dataset-ffhq\")\nimage_paths = [f\"../input/flickrfaceshq-dataset-ffhq/{el}\" for el in image_paths]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.environ['XLA_USE_BF16'] = \"1\"\nos.environ['XLA_TENSOR_ALLOCATOR_MAXSIZE'] = '100000000'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, image_paths):\n        super().__init__()\n        self.image_paths = image_paths\n        self.transform = albumentations.Compose([\n            albumentations.RandomCrop(height=256, width=256)        \n        ])\n        self.cutout_1 = albumentations.Cutout(num_holes=12, \n                                            max_h_size=12, \n                                            max_w_size=12,\n                                            p=1.0,\n                                            fill_value=1.0)\n        self.cutout_2 = albumentations.Cutout(num_holes=12, \n                                            max_h_size=24, \n                                            max_w_size=24,\n                                            p=1.0,\n                                            fill_value=1.0)\n        self.cutout_3 = albumentations.Cutout(num_holes=12, \n                                            max_h_size=36, \n                                            max_w_size=36,\n                                            p=1.0,\n                                            fill_value=1.0)        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, index):\n        image = cv2.imread(self.image_paths[index])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) / 255.0\n        if image.shape[0] < 256 or image.shape[1] < 256:\n            image = cv2.resize(image, (max(image.shape[0], 256), max(image.shape[1], 256)))\n            \n        image = self.transform(image=image)['image']\n        image_cutout = self.cutout_1(image=image)['image']\n        image_cutout = self.cutout_2(image=image_cutout)['image']\n        image_cutout = self.cutout_3(image=image_cutout)['image']\n        image = np.expand_dims(image, 0)\n        image_cutout = np.expand_dims(image_cutout, 0)\n        # TODO find way to generate exact mask (there could be pixel of value 1.0 already in the image)\n        mask = (image_cutout == 1.0) * 1.0\n        return image, image_cutout, mask\n\ndef reduce_sum(x, axis=None, keepdim=False):\n    if not axis:\n        axis = range(len(x.shape))\n    for i in sorted(axis, reverse=True):\n        x = torch.sum(x, dim=i, keepdim=keepdim)\n    return x\n    \ndef reduce_mean(x, axis=None, keepdim=False):\n    if not axis:\n        axis = range(len(x.shape))\n    for i in sorted(axis, reverse=True):\n        x = torch.mean(x, dim=i, keepdim=keepdim)\n    return x4\n    \ndef same_padding(images, ksizes, strides, rates):\n    assert len(images.size()) == 4\n    batch_size, channel, rows, cols = images.size()\n    out_rows = (rows + strides[0] - 1) // strides[0]\n    out_cols = (cols + strides[1] - 1) // strides[1]\n    effective_k_row = (ksizes[0] - 1) * rates[0] + 1\n    effective_k_col = (ksizes[1] - 1) * rates[1] + 1\n    padding_rows = max(0, (out_rows-1)*strides[0]+effective_k_row-rows)\n    padding_cols = max(0, (out_cols-1)*strides[1]+effective_k_col-cols)\n    # Pad the input\n    padding_top = int(padding_rows / 2.)\n    padding_left = int(padding_cols / 2.)\n    padding_bottom = padding_rows - padding_top\n    padding_right = padding_cols - padding_left\n    paddings = (padding_left, padding_right, padding_top, padding_bottom)\n    images = torch.nn.ZeroPad2d(paddings)(images)\n    return images\n    \ndef extract_image_patches(images, ksizes, strides, rates, padding='same'):\n    \"\"\"\n    Extract patches from images and put them in the C output dimension.\n    :param padding:\n    :param images: [batch, channels, in_rows, in_cols]. A 4-D Tensor with shape\n    :param ksizes: [ksize_rows, ksize_cols]. The size of the sliding window for\n     each dimension of images\n    :param strides: [stride_rows, stride_cols]\n    :param rates: [dilation_rows, dilation_cols]\n    :return: A Tensor\n    \"\"\"\n    assert len(images.size()) == 4\n    assert padding in ['same', 'valid']\n    batch_size, channel, height, width = images.size()\n\n    if padding == 'same':\n        images = same_padding(images, ksizes, strides, rates)\n    elif padding == 'valid':\n        pass\n    else:\n        raise NotImplementedError('Unsupported padding type: {}.\\\n                Only \"same\" or \"valid\" are supported.'.format(padding))\n\n    unfold = torch.nn.Unfold(kernel_size=ksizes,\n                             dilation=rates,\n                             padding=0,\n                             stride=strides)\n    patches = unfold(images)\n    return patches  # [N, C*k*k, L], L is the total number of such blocks\n\nclass ContextualAttention(nn.Module):\n    def __init__(self, ksize=3, stride=1, rate=1, fuse_k=3, softmax_scale=10,\n                 fuse=True, use_cuda=False, device_ids=None):\n        super(ContextualAttention, self).__init__()\n        self.ksize = ksize\n        self.stride = stride\n        self.rate = rate\n        self.fuse_k = fuse_k\n        self.softmax_scale = softmax_scale\n        self.fuse = fuse\n        self.use_cuda = use_cuda\n        self.device_ids = device_ids\n\n    def forward(self, f, b, mask=None):\n        \"\"\" Contextual attention layer implementation.\n        Contextual attention is first introduced in publication:\n            Generative Image Inpainting with Contextual Attention, Yu et al.\n        Args:\n            f: Input feature to match (foreground).\n            b: Input feature for match (background).\n            mask: Input mask for b, indicating patches not available.\n            ksize: Kernel size for contextual attention.\n            stride: Stride for extracting patches from b.\n            rate: Dilation for matching.\n            softmax_scale: Scaled softmax for attention.\n        Returns:\n            torch.tensor: output\n        \"\"\"\n        # get shapes\n        raw_int_fs = list(f.size())   # b*c*h*w\n        raw_int_bs = list(b.size())   # b*c*h*w\n\n        # extract patches from background with stride and rate\n        kernel = 2 * self.rate\n        # raw_w is extracted for reconstruction\n        raw_w = extract_image_patches(b, ksizes=[kernel, kernel],\n                                      strides=[self.rate*self.stride,\n                                               self.rate*self.stride],\n                                      rates=[1, 1],\n                                      padding='same') # [N, C*k*k, L]\n        # raw_shape: [N, C, k, k, L]\n        raw_w = raw_w.view(raw_int_bs[0], raw_int_bs[1], kernel, kernel, -1)\n        raw_w = raw_w.permute(0, 4, 1, 2, 3)    # raw_shape: [N, L, C, k, k]\n        raw_w_groups = torch.split(raw_w, 1, dim=0)\n\n        # downscaling foreground option: downscaling both foreground and\n        # background for matching and use original background for reconstruction.\n        f = F.interpolate(f, scale_factor=1./self.rate, mode='nearest', recompute_scale_factor=True)\n        b = F.interpolate(b, scale_factor=1./self.rate, mode='nearest', recompute_scale_factor=True)\n        int_fs = list(f.size())     # b*c*h*w\n        int_bs = list(b.size())\n        f_groups = torch.split(f, 1, dim=0)  # split tensors along the batch dimension\n        # w shape: [N, C*k*k, L]\n        w = extract_image_patches(b, ksizes=[self.ksize, self.ksize],\n                                  strides=[self.stride, self.stride],\n                                  rates=[1, 1],\n                                  padding='same')\n        # w shape: [N, C, k, k, L]\n        w = w.view(int_bs[0], int_bs[1], self.ksize, self.ksize, -1)\n        w = w.permute(0, 4, 1, 2, 3)    # w shape: [N, L, C, k, k]\n        w_groups = torch.split(w, 1, dim=0)\n\n        # process mask\n        if mask is None:\n            mask = torch.zeros([int_bs[0], 1, int_bs[2], int_bs[3]])\n            if self.use_cuda:\n                mask = mask.cuda()\n        else:\n            mask = F.interpolate(mask, scale_factor=1./(4*self.rate), mode='nearest', recompute_scale_factor=True)\n        int_ms = list(mask.size())\n        # m shape: [N, C*k*k, L]\n        m = extract_image_patches(mask, ksizes=[self.ksize, self.ksize],\n                                  strides=[self.stride, self.stride],\n                                  rates=[1, 1],\n                                  padding='same')\n        # m shape: [N, C, k, k, L]\n        m = m.view(int_ms[0], int_ms[1], self.ksize, self.ksize, -1)\n        m = m.permute(0, 4, 1, 2, 3)    # m shape: [N, L, C, k, k]\n        m = m[0]    # m shape: [L, C, k, k]\n        # mm shape: [L, 1, 1, 1]\n        mm = (reduce_mean(m, axis=[1, 2, 3], keepdim=True)==0.).to(torch.float32)\n        mm = mm.permute(1, 0, 2, 3) # mm shape: [1, L, 1, 1]\n\n        y = []\n        offsets = []\n        k = self.fuse_k\n        scale = self.softmax_scale    # to fit the PyTorch tensor image value range\n        fuse_weight = torch.eye(k).view(1, 1, k, k)  # 1*1*k*k\n        if self.use_cuda:\n            fuse_weight = fuse_weight.cuda()\n\n        for xi, wi, raw_wi in zip(f_groups, w_groups, raw_w_groups):\n            '''\n            O => output channel as a conv filter\n            I => input channel as a conv filter\n            xi : separated tensor along batch dimension of front; (B=1, C=128, H=32, W=32)\n            wi : separated patch tensor along batch dimension of back; (B=1, O=32*32, I=128, KH=3, KW=3)\n            raw_wi : separated tensor along batch dimension of back; (B=1, I=32*32, O=128, KH=4, KW=4)\n            '''\n            # conv for compare\n            escape_NaN = torch.FloatTensor([1e-4])\n            if self.use_cuda:\n                escape_NaN = escape_NaN.cuda()\n            wi = wi[0]  # [L, C, k, k]\n            max_wi = torch.sqrt(reduce_sum(torch.pow(wi, 2) + escape_NaN, axis=[1, 2, 3], keepdim=True))\n            wi_normed = wi / max_wi\n            # xi shape: [1, C, H, W], yi shape: [1, L, H, W]\n            xi = same_padding(xi, [self.ksize, self.ksize], [1, 1], [1, 1])  # xi: 1*c*H*W\n            yi = F.conv2d(xi, wi_normed, stride=1)   # [1, L, H, W]\n            # conv implementation for fuse scores to encourage large patches\n            if self.fuse:\n                # make all of depth to spatial resolution\n                yi = yi.view(1, 1, int_bs[2]*int_bs[3], int_fs[2]*int_fs[3])  # (B=1, I=1, H=32*32, W=32*32)\n                yi = same_padding(yi, [k, k], [1, 1], [1, 1])\n                yi = F.conv2d(yi, fuse_weight, stride=1)  # (B=1, C=1, H=32*32, W=32*32)\n                yi = yi.contiguous().view(1, int_bs[2], int_bs[3], int_fs[2], int_fs[3])  # (B=1, 32, 32, 32, 32)\n                yi = yi.permute(0, 2, 1, 4, 3)\n                yi = yi.contiguous().view(1, 1, int_bs[2]*int_bs[3], int_fs[2]*int_fs[3])\n                yi = same_padding(yi, [k, k], [1, 1], [1, 1])\n                yi = F.conv2d(yi, fuse_weight, stride=1)\n                yi = yi.contiguous().view(1, int_bs[3], int_bs[2], int_fs[3], int_fs[2])\n                yi = yi.permute(0, 2, 1, 4, 3).contiguous()\n            yi = yi.view(1, int_bs[2] * int_bs[3], int_fs[2], int_fs[3])  # (B=1, C=32*32, H=32, W=32)\n            # softmax to match\n            yi = yi * mm\n            yi = F.softmax(yi*scale, dim=1)\n            yi = yi * mm  # [1, L, H, W]\n\n            offset = torch.argmax(yi, dim=1, keepdim=True)  # 1*1*H*W\n\n            if int_bs != int_fs:\n                # Normalize the offset value to match foreground dimension\n                times = float(int_fs[2] * int_fs[3]) / float(int_bs[2] * int_bs[3])\n                offset = ((offset + 1).float() * times - 1).to(torch.int64)\n            offset = torch.cat([offset//int_fs[3], offset%int_fs[3]], dim=1)  # 1*2*H*W\n\n            # deconv for patch pasting\n            wi_center = raw_wi[0]\n            # yi = F.pad(yi, [0, 1, 0, 1])    # here may need conv_transpose same padding\n            yi = F.conv_transpose2d(yi, wi_center, stride=self.rate, padding=1) / 4.  # (B=1, C=128, H=64, W=64)\n            y.append(yi)\n            offsets.append(offset)\n\n        y = torch.cat(y, dim=0)  # back to the mini-batch\n        y.contiguous().view(raw_int_fs)\n\n        offsets = torch.cat(offsets, dim=0)\n        offsets = offsets.view(int_fs[0], 2, *int_fs[2:])\n\n        # case1: visualize optical flow: minus current position\n        h_add = torch.arange(int_fs[2]).view([1, 1, int_fs[2], 1]).expand(int_fs[0], -1, -1, int_fs[3])\n        w_add = torch.arange(int_fs[3]).view([1, 1, 1, int_fs[3]]).expand(int_fs[0], -1, int_fs[2], -1)\n        ref_coordinate = torch.cat([h_add, w_add], dim=1)\n        if self.use_cuda:\n            ref_coordinate = ref_coordinate.cuda()\n\n        offsets = offsets - ref_coordinate\n        # flow = pt_flow_to_image(offsets)\n\n        flow = torch.from_numpy(flow_to_image(offsets.permute(0, 2, 3, 1).cpu().data.numpy())) / 255.\n        flow = flow.permute(0, 3, 1, 2)\n        if self.use_cuda:\n            flow = flow.cuda()\n        # case2: visualize which pixels are attended\n        # flow = torch.from_numpy(highlight_flow((offsets * mask.long()).cpu().data.numpy()))\n\n        if self.rate != 1:\n            flow = F.interpolate(flow, scale_factor=self.rate*4, mode='nearest', recompute_scale_factor=True)\n\n        return y, flow\n\nclass GatedConv2d(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, activation=nn.LeakyReLU(0.2), padding_mode='replicate', **kwargs):\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels * 2, kernel_size, stride, padding, padding_mode=padding_mode, **kwargs)\n        self.activation = activation\n        self.out_channels = out_channels\n\n    def forward(self, x):\n        x = self.conv(x)\n        x, gate = x[:,:self.out_channels], x[:,self.out_channels:]\n        return self.activation(x) * torch.sigmoid(gate)\n    \nclass GatedConvTranspose2d(nn.Module):\n    # Not setting leakyRelu here led to dying neuron problem\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, activation=nn.LeakyReLU(0.2), **kwargs):\n        # TODO can be expressed in a single convolution with splitting afterwards\n        super().__init__()\n        self.conv = nn.ConvTranspose2d(in_channels, out_channels * 2, kernel_size, stride, padding, **kwargs)\n        self.activation = activation\n        self.out_channels = out_channels\n\n    def forward(self, x):\n        x = self.conv(x)\n        x, gate = x[:,:self.out_channels], x[:,self.out_channels:]\n        return self.activation(x) * torch.sigmoid(gate)\n\nclass CoarseNet(torch.nn.Module):\n    def __init__(self, in_channels, activation=nn.LeakyReLU(0.2), cnum=40):\n        super().__init__()\n        \n        self.l1 = GatedConv2d(in_channels + 1, cnum, 5, 1, 2, activation=activation)\n        self.l2 = GatedConv2d(cnum, 2 * cnum, 3, 2, 1, activation=activation)\n        self.l3 = GatedConv2d(2 * cnum, 2 * cnum, 3, 1, 1, activation=activation)\n        self.l4 = GatedConv2d(2 * cnum, 4 * cnum, 3, 2, 1, activation=activation)\n        self.l5 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        self.l6 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n\n        self.l7 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 2, dilation=2, activation=activation)\n        self.l8 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 4, dilation=4, activation=activation)\n        self.l9 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 8, dilation=8, activation=activation)\n        self.l10 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 16, dilation=16, activation=activation)\n        self.l11 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n\n        self.l12 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n\n        self.l13 = GatedConvTranspose2d(4 * cnum, 2 * cnum, 3, 2, 1, activation=activation, output_padding=1)\n        self.l14 = GatedConv2d(2 * cnum, 2 * cnum, 3, 1, 1, activation=activation)\n        self.l15 = GatedConvTranspose2d(2 * cnum, cnum, 3, 2, 1, activation=activation, output_padding=1)\n        self.l16 = GatedConv2d(cnum, cnum // 2, 3, 1, 1, activation=activation)\n        self.l17 = GatedConv2d(cnum // 2, in_channels, 3, 1, 1, activation=lambda x: x)\n        \n    def forward(self, x, mask):\n        x = torch.cat([x, mask], dim=1)\n        x = self.l1(x)\n        x = self.l2(x)\n        x = self.l3(x)\n        x = self.l4(x)\n        x = self.l5(x)\n        x = self.l6(x)\n        x = self.l7(x)\n        x = self.l8(x)\n        x = self.l9(x)\n        x = self.l10(x)\n        x = self.l11(x)\n        x = self.l12(x)\n        x = self.l13(x)\n        x = self.l14(x)\n        x = self.l15(x)\n        x = self.l16(x)\n        x = self.l17(x)\n        return x\n    \nclass RefineNet(torch.nn.Module):\n    def __init__(self, n_in_channel, device, activation=nn.LeakyReLU(0.2), cnum=40):\n        super().__init__()\n        \n        self.l1_at = GatedConv2d(n_in_channel + 1, cnum, 5, 1, 2, activation=activation)\n        self.l2_at = GatedConv2d(cnum, 2 * cnum, 3, 2, 1, activation=activation)\n        self.l3_at = GatedConv2d(2 * cnum, 2 * cnum, 3, 1, 1, activation=activation)\n        self.l4_at = GatedConv2d(2 * cnum, 4 * cnum, 3, 2, 1, activation=activation)\n        self.l5_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        self.l6_at = SelfAttention(4 * cnum)\n        # self.l6_at = ContextualAttention(ksize=3, stride=1, rate=2, fuse_k=3, softmax_scale=10, fuse=True, use_cuda=True)\n        self.l7_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        self.l8_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        \n        self.l1_no_at = GatedConv2d(n_in_channel + 1, cnum, 5, 1, 2, activation=activation)\n        self.l2_no_at = GatedConv2d(cnum, 2 * cnum, 3, 2, 1, activation=activation)\n        self.l3_no_at = GatedConv2d(2 * cnum, 2 * cnum, 3, 1, 1, activation=activation)\n        self.l4_no_at = GatedConv2d(2 * cnum, 4 * cnum, 3, 2, 1, activation=activation)\n        self.l5_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        self.l6_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n        self.l7_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 2, dilation=2, activation=activation)\n        self.l8_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 4, dilation=4, activation=activation)\n        self.l9_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 8, dilation=8, activation=activation)\n        self.l10_no_at = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 16, dilation=16, activation=activation)\n\n        self.up_1 = GatedConv2d(2 * 4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n\n        self.up_2 = GatedConv2d(4 * cnum, 4 * cnum, 3, 1, 1, activation=activation)\n\n        self.up_3 = GatedConvTranspose2d(4 * cnum, 2 * cnum, 3, 2, 1, activation=activation, output_padding=1)\n        self.up_4 = GatedConv2d(2 * cnum, 2 * cnum, 3, 1, 1, activation=activation)\n        self.up_5 = GatedConvTranspose2d(2 * cnum, cnum, 3, 2, 1, activation=activation, output_padding=1)\n        self.up_6 = GatedConv2d(cnum, cnum // 2, 3, 1, 1, activation=activation)\n        self.up_7 = GatedConv2d(cnum // 2, n_in_channel, 3, 1, 1, activation=lambda x: x)\n        \n    def non_attention_branch(self, x):\n        x = self.l1_no_at(x) \n        x = self.l2_no_at(x)\n        x = self.l3_no_at(x)\n        x = self.l4_no_at(x)\n        x = self.l5_no_at(x)\n        x = self.l6_no_at(x)\n        x = self.l7_no_at(x)\n        x = self.l8_no_at(x)\n        x = self.l9_no_at(x)\n        x = self.l10_no_at(x)\n        return x\n        \n    def attention_branch(self, x, mask):\n        x = self.l1_at(x)\n        x = self.l2_at(x)\n        x = self.l3_at(x)\n        x = self.l4_at(x)\n        x = self.l5_at(x)\n        # mask = F.interpolate(mask, (x.shape[2], x.shape[3]))\n        # x, offset_flow = self.l6_at(x, x, mask)\n        x = self.l6_at(x)\n        x = self.l7_at(x)\n        x = self.l8_at(x)\n        return x\n        \n    def forward(self, x, mask):\n        x = torch.cat([x, mask], dim=1)\n        x_at = self.attention_branch(x, mask)\n        x_no_at = self.non_attention_branch(x)\n        x = torch.cat([x_at, x_no_at], dim=1)\n        x = self.up_1(x)\n        x = self.up_2(x)\n        x = self.up_3(x)\n        x = self.up_4(x)\n        x = self.up_5(x)\n        x = self.up_6(x)\n        x = self.up_7(x)\n        return x\n    \nclass CoarseRefineNet(nn.Module):\n    def __init__(self, in_channels, device):\n        super().__init__()\n        self.coarse_net = CoarseNet(in_channels)\n        self.refine_net = RefineNet(in_channels, device)\n    \n    def call_coarse(self, x, mask):\n        x_coarse = self.coarse_net(x, mask)\n        return x_coarse\n    \n    def call_refined(self, x, x_real, mask):\n        x_coarse_filled = x * mask + (1 - mask) * x_real\n        x_refined = self.refine_net(x_coarse_filled, mask)\n        return x_refined\n        \n    def forward(self, x, mask):\n        x_coarse = self.coarse_net(x, mask)\n        x_coarse_filled = x_coarse * mask + (1 - mask) * x\n        x_refined = self.refine_net(x_coarse_filled, mask)\n        return x_coarse, x_refined\n\nclass PatchDiscriminatorBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, use_norm=True, activation=nn.LeakyReLU(0.2)):\n        super().__init__()\n        self.conv_1 = nn.utils.spectral_norm(nn.Conv2d(in_channels, out_channels, 5, 2, 2))\n        self.activation = activation\n        self.use_norm = use_norm\n        if use_norm:\n            # this is the same as layer_norm\n            self.norm = nn.GroupNorm(1, out_channels)\n        \n    def forward(self, x):\n        x = self.conv_1(x)\n        x = self.activation(x)\n        if self.use_norm:\n            x = self.norm(x)\n        return x\n    \n\nclass PatchCritic(nn.Module):\n    def __init__(self, in_channels, activation=nn.LeakyReLU(0.2), cnum=64):\n        super().__init__()\n        self.b1 = PatchDiscriminatorBlock(in_channels + 1, cnum, use_norm=False)\n        self.b2 = PatchDiscriminatorBlock(cnum, cnum * 2, use_norm=False)\n        self.b3 = PatchDiscriminatorBlock(cnum * 2, cnum * 4, use_norm=False)\n        self.b4 = PatchDiscriminatorBlock(cnum * 4, cnum * 4, use_norm=False)\n        self.b5 = PatchDiscriminatorBlock(cnum * 4, cnum * 4, use_norm=False)\n        self.b6 = PatchDiscriminatorBlock(cnum * 4, cnum * 4, use_norm=False, activation=lambda x: x)\n        \n    def forward(self, x, mask):\n        x = torch.cat([x, mask], dim=1)\n        x = self.b1(x)\n        x = self.b2(x)\n        x = self.b3(x)\n        x = self.b4(x)\n        x = self.b5(x)\n        x = self.b6(x)\n        x = x.view((x.shape[0], -1))\n        return x\n        \n    \nclass SelfAttention(torch.nn.Module):\n    def __init__(self, in_dim, return_attention=False):\n        super().__init__()\n        self.in_dim = in_dim\n        self.return_attention = return_attention\n        self.query_conv = torch.nn.Conv2d(in_channels=in_dim, out_channels=in_dim//8, kernel_size=1)\n        self.key_conv = torch.nn.Conv2d(in_channels=in_dim, out_channels=in_dim//8, kernel_size=1)\n        self.value_conv = torch.nn.Conv2d(in_channels=in_dim, out_channels=in_dim, kernel_size=1)\n        self.gamma = torch.nn.Parameter(torch.zeros(1))\n        self.softmax = torch.nn.Softmax(dim=-1)\n        self.init_weights()\n\n    def init_weights(self):\n        for m in self.modules():\n            if isinstance(m, torch.nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n\n    def forward(self, x):\n        B, C, H, W = x.size()\n        proj_query = self.query_conv(x).view(B, -1, W*H).permute(0, 2, 1)\n        proj_key = self.key_conv(x).view(B, -1, W*H)\n        energy = torch.bmm(proj_query, proj_key)\n        attention = self.softmax(energy)\n        proj_value = self.value_conv(x).view(B, -1, W*H)\n\n        out = torch.bmm(proj_value, attention.permute(0, 2, 1))\n        out = out.view(B, C, H, W)\n        out = self.gamma * out + x\n\n        if self.return_attention:\n            return out, attention\n        return out\n    \ndef complete_img(real_img, fake_img, mask):\n    return mask * fake_img + (1 - mask) * real_img\n        \ndef wasserstein_loss(y_pred):\n    \"\"\" Minimisation of the earth-mover distance\n    \"\"\"\n    return torch.mean(y_pred)\n\n\ndef hinge_loss(y_pred, y_true):\n    \"\"\" Hinge loss, used for large margin classifiers. \n    If y = 1 and y_pred >= 1, the loss is 0, otherwise its 1 - y_pred\n    If y = 0 and y_pred <= -1, the loss is 0, otherwise its 1 + y_pred\n    y_pred: predictions of the model of shape (m, 1). Can be outside of range [0, 1]\n    y_true: binary labels of shape (m, 1), can be either 0 or 1.\n    \"\"\"\n    y_true = torch.where(y_true == 0, -1.0 * torch.ones_like(y_true).float(), y_true.float())\n    l = torch.max(torch.zeros_like(y_true), 1 - y_true * y_pred)\n    return torch.mean(l)\n\n\"\"\"def adv_loss_disc(y_pred_real, y_pred_fake):\n    return torch.mean(F.relu(1 - y_pred_real)) + torch.mean(F.relu(1 + y_pred_fake))\"\"\"\n\ndef adv_loss_disc(y_pred_real, y_pred_fake):\n    return -torch.mean(y_pred_real) + torch.mean(y_pred_fake)\n    \ndef adv_loss_gen(y_pred_fake):\n    return - torch.mean(y_pred_fake)\n\nl1_loss = torch.nn.L1Loss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def renormalise(img):\n    return (img + 1.0) * 0.5\n\ndef plot_progress(x, x_blacked_out, fake_imgs, mask):\n    \"\"\"x = renormalise(x)\n    x_blacked_out = renormalise(x_blacked_out)\n    fake_imgs = renormalise(fake_imgs)\"\"\"\n    img_reconstructed = mask[0] * fake_imgs[0] + (1 - mask[0]) * x[0]\n    fig, ax = plt.subplots(nrows=2, ncols=2, figsize=(20, 20))\n    ax[0, 0].imshow(torch.clamp(x[0], 0, 1).cpu().reshape((256, 256)), cmap=\"gray\", vmin=0.0, vmax=1.0)\n    ax[0, 1].imshow(torch.clamp(x_blacked_out[0, 0], 0, 1).cpu().reshape((256, 256)), cmap=\"gray\", vmin=0.0, vmax=1.0)\n    ax[1, 0].imshow(torch.clamp(fake_imgs[0], 0, 1).detach().cpu().reshape((256, 256)), cmap=\"gray\", vmin=0.0, vmax=1.0)\n    ax[1, 1].imshow(torch.clamp(img_reconstructed, 0, 1).detach().cpu().reshape((256, 256)), cmap=\"gray\", vmin=0.0, vmax=1.0)\n    plt.show()\n\ndef train_fn(index, flags):\n    \n    if flags['use_tpu']:\n        # Sets a common random seed - both for initialization and ensuring graph is the same\n        torch.manual_seed(flags['seed'])\n        # Acquires the (unique) Cloud TPU core corresponding to this process's index\n        device = xm.xla_device()  \n        ds_train = ImageDataset(image_paths)\n        train_sampler = torch.utils.data.distributed.DistributedSampler(\n            ds_train,\n            num_replicas=xm.xrt_world_size(),\n            rank=xm.get_ordinal(),\n            shuffle=True\n        )\n        train_loader = DataLoader(\n              ds_train,\n              batch_size=20, \n              sampler=train_sampler,\n              num_workers=multiprocessing.cpu_count(),\n              drop_last=True\n        )\n        save_fn = xm.save\n    else:\n        device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        ds_train = ImageDataset(image_paths)\n        train_loader = DataLoader(\n              ds_train,\n              batch_size=16,\n              num_workers=multiprocessing.cpu_count(),\n              shuffle=True,\n              drop_last=True\n        )\n        save_fn = torch.save\n\n    gen = CoarseRefineNet(1, device).to(device)\n    disc = PatchCritic(1).to(device)\n    optim_coarse = torch.optim.Adam(gen.coarse_net.parameters(), lr=0.0001, betas=(0.5, 0.999))\n    optim_fine = torch.optim.Adam(gen.refine_net.parameters(), lr=0.0001, betas=(0.5, 0.999))\n    optim_disc = torch.optim.Adam(disc.parameters(), lr=0.0004, betas=(0.5, 0.999))\n\n    gen.load_state_dict(torch.load(\"../input/coarse-refine-attention-patchgan-on-tpu-2/gen_model\"))\n    disc.load_state_dict(torch.load(\"../input/coarse-refine-attention-patchgan-on-tpu-2/disc_model\"))\n    optim_coarse.load_state_dict(torch.load(\"../input/coarse-refine-attention-patchgan-on-tpu-2/optim_coarse\"))\n    optim_fine.load_state_dict(torch.load(\"../input/coarse-refine-attention-patchgan-on-tpu-2/optim_fine\"))\n    optim_disc.load_state_dict(torch.load(\"../input/coarse-refine-attention-patchgan-on-tpu-2/optim_disc\"))\n\n    for epoch in range(1, flags['num_epochs'] + 1):\n        if flags['use_tpu']:\n            train_sampler.set_epoch(epoch)\n            train_iterator = pl.ParallelLoader(train_loader, [device]).per_device_loader(device)\n        else:\n            train_iterator = train_loader\n        gen_coarse_rec_losses, gen_fine_rec_losses, gen_fine_adv_losses, gen_disc_losses, disc_losses = [], [], [], [], []\n        for i, (x, x_blacked_out, mask) in enumerate(tqdm(train_iterator)):\n            x, x_blacked_out, mask = x.to(device).float(), x_blacked_out.to(device).float(), mask.to(device).float()\n            with torch.no_grad():\n                fake_imgs_coarse, fake_imgs_refined = gen(x_blacked_out, mask) \n            y_pred_real = disc(x, mask)\n            y_pred_fake = disc(complete_img(x, fake_imgs_refined, mask), mask)\n            disc_loss = adv_loss_disc(y_pred_real, y_pred_fake)\n            optim_disc.zero_grad()\n            disc_loss.backward()\n            nn.utils.clip_grad_norm_(disc.parameters(), flags['grad_clip_value'])\n            if flags['use_tpu']:\n                xm.optimizer_step(optim_disc)\n            else:\n                optim_disc.step()\n            disc_losses.append(disc_loss.item())\n              \n            if i % 3 == 0:\n                # train coarse generator\n                fake_imgs_coarse = gen.call_coarse(x_blacked_out, mask) \n                gen_coarse_rec_loss = l1_loss(fake_imgs_coarse, x)\n\n                optim_coarse.zero_grad()\n                gen_coarse_rec_loss.backward()\n                nn.utils.clip_grad_norm_(gen.coarse_net.parameters(), flags['grad_clip_value'])\n                if flags['use_tpu']:\n                    xm.optimizer_step(optim_coarse)\n                else:\n                    optim_coarse.step()\n                gen_coarse_rec_losses.append(gen_coarse_rec_loss.item())\n\n                # train refine generator\n                fake_imgs_refined = gen.call_refined(fake_imgs_coarse.detach(), x, mask)\n                y_pred_fake = disc(complete_img(x, fake_imgs_refined, mask), mask)\n                gen_fine_rec_loss = l1_loss(fake_imgs_refined, x)\n                gen_fine_adv_loss = adv_loss_gen(y_pred_fake)\n                gen_fine_loss = gen_fine_rec_loss + 1 * gen_fine_adv_loss\n\n\n                optim_fine.zero_grad()\n                gen_fine_loss.backward()\n                nn.utils.clip_grad_norm_(gen.refine_net.parameters(), flags['grad_clip_value'])\n                if flags['use_tpu']:\n                    xm.optimizer_step(optim_fine)\n                else:\n                    optim_fine.step()\n\n                gen_fine_rec_losses.append(gen_fine_rec_loss.item())\n                gen_fine_adv_losses.append(gen_fine_adv_loss.item())\n            \n            if flags['use_tpu']:\n                if index == 0 and i != 0:\n                    if i % 5 == 0:\n                        print(f\"{epoch}/{flags['num_epochs']}, ITERATION {i} gen_coarse_rec_loss: {round(np.mean(gen_coarse_rec_losses), 6)} gen_fine_rec_loss: {round(np.mean(gen_fine_rec_losses), 6)} gen_fine_adv_loss: {round(np.mean(gen_fine_adv_losses), 6)} disc_loss: {round(np.mean(disc_losses), 6)}\")\n                        plot_progress(x, x_blacked_out, fake_imgs_refined, mask)\n            else:\n                if i % 100  == 0 and i != 0:\n                    print(f\"{epoch}/{flags['num_epochs']}, ITERATION {i} gen_coarse_rec_loss: {round(np.mean(gen_coarse_rec_losses), 6)} gen_fine_rec_loss: {round(np.mean(gen_fine_rec_losses), 6)} gen_fine_adv_loss: {round(np.mean(gen_fine_adv_losses), 6)} disc_loss: {round(np.mean(disc_losses), 6)}\")\n                    plot_progress(x, x_blacked_out, fake_imgs_refined, mask)\n\n                if i % 250 == 0 and i != 0:\n                    save_fn(gen.state_dict(), \"./gen_model\")\n                    save_fn(optim_coarse.state_dict(), './optim_coarse')\n                    save_fn(optim_fine.state_dict(), './optim_fine')\n                    save_fn(disc.state_dict(), \"./disc_model\")\n                    save_fn(optim_disc.state_dict(), './optim_disc')\n        \n        save_fn(gen.state_dict(), \"./gen_model\")\n        save_fn(optim_coarse.state_dict(), './optim_coarse')\n        save_fn(optim_fine.state_dict(), './optim_fine')\n        save_fn(disc.state_dict(), \"./disc_model\")\n        save_fn(optim_disc.state_dict(), './optim_disc')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"flags = {}\n\nflags['num_epochs'] = 7\nflags['seed'] = 1234\nflags['grad_clip_value'] = 1.0\nflags['use_tpu'] = use_tpu\n\nif use_tpu:\n    xmp.spawn(train_fn, args=(flags,), nprocs=8, start_method='fork')\nelse:\n    train_fn(0, flags)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}