{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":8797547,"sourceType":"datasetVersion","datasetId":5290011},{"sourceId":9008063,"sourceType":"datasetVersion","datasetId":5426967},{"sourceId":9008416,"sourceType":"datasetVersion","datasetId":5427227}],"dockerImageVersionId":28772,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!git clone https://github.com/weepiess/StyleFlow-Content-Fixed-I2I.git\nfrom os import chdir\nchdir(\"StyleFlow-Content-Fixed-I2I\")","metadata":{"execution":{"iopub.status.busy":"2024-07-24T04:04:38.035274Z","iopub.execute_input":"2024-07-24T04:04:38.035668Z","iopub.status.idle":"2024-07-24T04:04:38.040515Z","shell.execute_reply.started":"2024-07-24T04:04:38.035613Z","shell.execute_reply":"2024-07-24T04:04:38.039399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nvgg_path = \"/kaggle/working/StyleFlow-Content-Fixed-I2I/vgg_model/vgg_normalised.pth\"\nraw_vgg = nn.Sequential(\n    nn.Conv2d(3, 3, (1, 1)), \n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(3, 64, (3, 3)),\n    nn.ReLU(),  # relu1-1\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(64, 64, (3, 3)),\n    nn.ReLU(),  # relu1-2\n    nn.MaxPool2d((2, 2), (2, 2), (0, 0), ceil_mode=True),\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(64, 128, (3, 3)),\n    nn.ReLU(),  # relu2-1\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(128, 128, (3, 3)),\n    nn.ReLU(),  # relu2-2\n    nn.MaxPool2d((2, 2), (2, 2), (0, 0), ceil_mode=True),\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(128, 256, (3, 3)),\n    nn.ReLU(),  # relu3-1\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(256, 256, (3, 3)),\n    nn.ReLU(),  # relu3-2\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(256, 256, (3, 3)),\n    nn.ReLU(),  # relu3-3\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(256, 256, (3, 3)),\n    nn.ReLU(),  # relu3-4\n    nn.MaxPool2d((2, 2), (2, 2), (0, 0), ceil_mode=True),\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(256, 512, (3, 3)),\n    nn.ReLU(),  # relu4-1, this is the last layer used\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu4-2\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu4-3\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu4-4\n    nn.MaxPool2d((2, 2), (2, 2), (0, 0), ceil_mode=True),\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu5-1\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu5-2\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU(),  # relu5-3\n    nn.ReflectionPad2d((1, 1, 1, 1)),\n    nn.Conv2d(512, 512, (3, 3)),\n    nn.ReLU()  # relu5-4\n)\nraw_vgg.load_state_dict(torch.load(vgg_path))\n\nclass Net(nn.Module):\n    def __init__(self, encoder, keep_ratio=1.0):\n        super(Net, self).__init__()\n        enc_layers = list(encoder.children())\n        self.enc_1 = nn.Sequential(*enc_layers[:4])  # input -> relu1_1\n        self.enc_2 = nn.Sequential(*enc_layers[4:11])  # relu1_1 -> relu2_1\n        self.enc_3 = nn.Sequential(*enc_layers[11:18])  # relu2_1 -> relu3_1\n        self.enc_4 = nn.Sequential(*enc_layers[18:31])  # relu3_1 -> relu4_1 31\n\n        self.mse_loss = nn.MSELoss()\n        self.keep_ratio = keep_ratio\n\n        # fix the encoder\n        for name in ['enc_1', 'enc_2', 'enc_3', 'enc_4']:\n            for param in getattr(self, name).parameters():\n                param.requires_grad = False\n\n    # extract relu1_1, relu2_1, relu3_1, relu4_1 from input image\n    def encode_with_intermediate(self, input):\n        results = [input]\n        for i in range(4):\n            func = getattr(self, 'enc_{:d}'.format(i + 1))\n            results.append(func(results[-1]))\n        return results[1:]\n\n    # extract relu4_1 from input image\n    def encode(self, input):\n        for i in range(4):\n            input = getattr(self, 'enc_{:d}'.format(i + 1))(input)\n        return input\n\nvgg = Net(raw_vgg)","metadata":{"execution":{"iopub.status.busy":"2024-07-24T04:52:15.476693Z","iopub.execute_input":"2024-07-24T04:52:15.47703Z","iopub.status.idle":"2024-07-24T04:52:15.757406Z","shell.execute_reply.started":"2024-07-24T04:52:15.476979Z","shell.execute_reply":"2024-07-24T04:52:15.756649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from cv2 import imread, resize\nimport numpy as np\n\ndef read_im(path, target_size=(224, 224)):\n    np_im = resize(imread(path)/255.0, target_size)\n    np_im = np.moveaxis(np_im, -1, 0)\n    return torch.tensor(np_im).unsqueeze(0).type(torch.float)\n\ndef viz_map(feats):\n    # plot all 64 maps in an 8x8 squares\n    square = 8\n    ix = 1\n    f = plt.figure(figsize=(16,16))\n    for _ in range(square):\n        for _ in range(square):\n            # specify subplot and turn of axis\n            ax = pyplot.subplot(square, square, ix)\n            ax.set_xticks([])\n            ax.set_yticks([])\n            # plot filter channel in grayscale\n            pyplot.imshow(feats[:, :, ix-1], cmap='viridis')\n            ix += 1\n    # show the figure\n    pyplot.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-24T05:08:59.628638Z","iopub.execute_input":"2024-07-24T05:08:59.62897Z","iopub.status.idle":"2024-07-24T05:08:59.639989Z","shell.execute_reply.started":"2024-07-24T05:08:59.628922Z","shell.execute_reply":"2024-07-24T05:08:59.639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot \nimport matplotlib.pyplot as plt\nfrom numpy import expand_dims\n\nimg = read_im(\"/kaggle/input/gta-ds/00012.png\")\nfeats = vgg.encode_with_intermediate(img)\nfeats_dict = { f\"relu_{idx}_1\":feat[0].numpy() for idx, feat in enumerate(feats, 1) }\nfor lay_name, feat in feats_dict.items():\n    feat = np.moveaxis(feat, 0, -1)\n    print(f\"{lay_name} viz {feat.shape}: \")\n    viz_map(feat)\n    print(\"\\n======================\\n\\n\")","metadata":{"execution":{"iopub.status.busy":"2024-07-24T05:09:43.234737Z","iopub.execute_input":"2024-07-24T05:09:43.235083Z","iopub.status.idle":"2024-07-24T05:09:54.109477Z","shell.execute_reply.started":"2024-07-24T05:09:43.235033Z","shell.execute_reply":"2024-07-24T05:09:54.108692Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import einsum\n\nim1 = read_im(\"/kaggle/input/gta-ds/00012.png\")\nim2 = read_im(\"/kaggle/input/cts-ds/aachen_000011_000019_leftImg8bit.png\")\nfeats1, feats2 = vgg.encode_with_intermediate(im1), vgg.encode_with_intermediate(im2)\nfeats1_dict = { f\"relu_{idx}_1\":feat[0].numpy() for idx, feat in enumerate(feats1, 1) }\nfeats2_dict = { f\"relu_{idx}_1\":feat[0].numpy() for idx, feat in enumerate(feats2, 1) }\n\ndef corr_viz(feat1, feat2):\n    feat1, feat2 = rearrange(feat1, \"b c h w -> b c (h w)\"), rearrange(feat2, \"b c h w -> b c (h w)\")\n    corr = einsum(\"b c a, b g a -> b c g\", feat1, feat2)\n    \n    \n\nfor lay_name in feats1_dict.keys():\n    feat1, feat2 = feats1_dict[lay_name], feats2_dict[lay_name]\n    feat1, feat2 = np.moveaxis(feat1, 0, -1), np.moveaxis(feat2, 0, -1)\n    print(f\"{lay_name} viz {feat1.shape}: \")\n    corr_viz(feat1, feat2)\n    print(\"\\n======================\\n\\n\")","metadata":{"execution":{"iopub.status.busy":"2024-07-24T05:19:47.097982Z","iopub.execute_input":"2024-07-24T05:19:47.098377Z","iopub.status.idle":"2024-07-24T05:19:47.111373Z","shell.execute_reply.started":"2024-07-24T05:19:47.098314Z","shell.execute_reply":"2024-07-24T05:19:47.110244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#corr = einsum(\"b c h w, b g h w -> b c g\", feats[0], feats[0])\n#corr.shape\n#from einops import rearrange\n\ncorr2.shape\n#corr2.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-24T05:30:55.70706Z","iopub.execute_input":"2024-07-24T05:30:55.707448Z","iopub.status.idle":"2024-07-24T05:30:55.714005Z","shell.execute_reply.started":"2024-07-24T05:30:55.707384Z","shell.execute_reply":"2024-07-24T05:30:55.712909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}}]}