{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        #print(os.path.join(dirname, filename))\n        pass\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-09T13:59:30.401571Z","iopub.execute_input":"2022-08-09T13:59:30.402214Z","iopub.status.idle":"2022-08-09T13:59:30.457392Z","shell.execute_reply.started":"2022-08-09T13:59:30.402081Z","shell.execute_reply":"2022-08-09T13:59:30.456115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#imports and init stuff\nimport os\nimport sys\nimport copy\nimport torch\nimport numpy as np\nfrom torch import nn\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom zipfile import ZipFile\nfrom torch.nn import LayerNorm\nimport torch.nn.functional as F\nfrom torchvision import transforms\nfrom collections import OrderedDict\nfrom torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize\nfrom sklearn.decomposition import PCA\nimport pickle as pk\n\ntry:\n    from torchvision.transforms import InterpolationMode\n    BICUBIC = InterpolationMode.BICUBIC\nexcept ImportError: BICUBIC = Image.BICUBIC","metadata":{"execution":{"iopub.status.busy":"2022-08-09T13:59:39.360815Z","iopub.execute_input":"2022-08-09T13:59:39.361972Z","iopub.status.idle":"2022-08-09T13:59:40.928372Z","shell.execute_reply.started":"2022-08-09T13:59:39.361923Z","shell.execute_reply":"2022-08-09T13:59:40.927122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#activation function for clip\nclass QuickGELU(nn.Module):\n    def forward(self, x: torch.Tensor):\n        return x * torch.sigmoid(1.702 * x)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:07:28.676926Z","iopub.execute_input":"2022-08-09T14:07:28.677334Z","iopub.status.idle":"2022-08-09T14:07:28.683676Z","shell.execute_reply.started":"2022-08-09T14:07:28.677302Z","shell.execute_reply":"2022-08-09T14:07:28.682191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#classes for clip \nclass ResidualAttentionBlock(nn.Module):\n    def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):\n        super().__init__()\n        self.attn = nn.MultiheadAttention(d_model, n_head)\n        self.ln_1 = LayerNorm(d_model)\n        self.mlp = nn.Sequential(OrderedDict([\n            (\"c_fc\", nn.Linear(d_model, d_model * 4)),\n            (\"gelu\", QuickGELU()),\n            (\"c_proj\", nn.Linear(d_model * 4, d_model))\n        ]))\n        self.ln_2 = LayerNorm(d_model)\n        self.attn_mask = attn_mask\n\n    def attention(self, x: torch.Tensor):\n        self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None\n        return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]\n\n    def forward(self, x: torch.Tensor):\n        x = x + self.attention(self.ln_1(x))\n        x = x + self.mlp(self.ln_2(x))\n        return x\n\nclass Transformer(nn.Module):\n    def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):\n        super().__init__()\n        self.width = width\n        self.layers = layers\n        self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])\n\n    def forward(self, x: torch.Tensor):\n        return self.resblocks(x)\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):\n        super().__init__()\n        self.input_resolution = input_resolution\n        self.output_dim = output_dim\n        self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)\n        scale = width ** -0.5\n        self.class_embedding = nn.Parameter(scale * torch.randn(width))\n        self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))\n        self.ln_pre = LayerNorm(width)\n        self.transformer = Transformer(width, layers, heads)\n        self.ln_post = LayerNorm(width)\n        self.proj = nn.Parameter(scale * torch.randn(width, output_dim))        \n        \n    def forward(self, x: torch.Tensor):\n        x = self.conv1(x)  # shape = [*, width, grid, grid]\n        x = x.reshape(x.shape[0], x.shape[1], -1)  # shape = [*, width, grid ** 2]\n        x = x.permute(0, 2, 1)  # shape = [*, grid ** 2, width]\n        x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1)  # shape = [*, grid ** 2 + 1, width]\n        x = x + self.positional_embedding.to(x.dtype)\n        x = self.ln_pre(x)\n        x = x.permute(1, 0, 2)\n        x = self.transformer(x)\n        x = x.permute(1, 0, 2)\n        x = self.ln_post(x[:, 0, :])\n        if self.proj is not None:\n            x = x @ self.proj\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:07:29.762996Z","iopub.execute_input":"2022-08-09T14:07:29.763398Z","iopub.status.idle":"2022-08-09T14:07:29.798403Z","shell.execute_reply.started":"2022-08-09T14:07:29.763366Z","shell.execute_reply":"2022-08-09T14:07:29.797225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#PCA params load\npca = pk.load(open(\"/kaggle/input/gimpcaweights/pca_weights.pkl\",'rb'))","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:07:32.321377Z","iopub.execute_input":"2022-08-09T14:07:32.321787Z","iopub.status.idle":"2022-08-09T14:07:32.327936Z","shell.execute_reply.started":"2022-08-09T14:07:32.321752Z","shell.execute_reply":"2022-08-09T14:07:32.326888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#pca torch model\nclass Model_PCA(nn.Module):\n    def __init__(self, dim1=768, dim2=64):                                                        \n        super(Model_PCA, self).__init__()      \n            \n        self.mean = nn.Parameter( torch.zeros((dim1, ),     dtype=torch.float32),  requires_grad = False)\n        self.comp = nn.Parameter( torch.zeros((dim2, dim1), dtype=torch.float32),  requires_grad = False)\n \n    def set(self, mean, comp):  # numpy\n        self.mean.copy_( torch.tensor(mean, dtype=torch.float32) )\n        self.comp.copy_( torch.tensor(comp, dtype=torch.float32) )\n\n    def forward(self, x):               \n        x =  x - self.mean\n        print(x.shape, self.comp.T.shape)\n        x =  torch.mm(x, self.comp.T)        \n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:10:49.328029Z","iopub.execute_input":"2022-08-09T14:10:49.329776Z","iopub.status.idle":"2022-08-09T14:10:49.340214Z","shell.execute_reply.started":"2022-08-09T14:10:49.329708Z","shell.execute_reply":"2022-08-09T14:10:49.338909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model_pca = Model_PCA()  \n#model_pca.set(pca.mean_, pca.components_)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#full model for usage\nclass ClipEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.feature_extractor = VisionTransformer(heads = 16,\n                                                    layers = 24,\n                                                    width = 1024,\n                                                    patch_size = 14,\n                                                    output_dim = 768,\n                                                    input_resolution = 336)\n        self.feature_extractor.load_state_dict(\n            torch.load('../input/openaiclip-weights/ViT_L_14_336px_vision_model.pt'))\n        \n        self.reducer_pca = Model_PCA()\n    \n    def forward(self, x):\n        # ----- clip -----\n        x = transforms.functional.resize(x, size = [336, 336])\n        x = x / 255.0\n        x = transforms.functional.normalize(x, \n                                            mean = [0.48145466, 0.4578275, 0.40821073], \n                                            std = [0.26862954, 0.26130258, 0.27577711])        \n        \n        x = self.feature_extractor(x) # / 2.0 + self.feature_extractor(x1.flip([-1])) / 2.0\n        \n        x = self.reducer_pca(x)\n        \n        x = torch.nn.functional.normalize(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:11:38.937384Z","iopub.execute_input":"2022-08-09T14:11:38.937805Z","iopub.status.idle":"2022-08-09T14:11:38.948720Z","shell.execute_reply.started":"2022-08-09T14:11:38.937771Z","shell.execute_reply":"2022-08-09T14:11:38.947454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pca.mean_.shape, pca.components_.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:11:41.355540Z","iopub.execute_input":"2022-08-09T14:11:41.356679Z","iopub.status.idle":"2022-08-09T14:11:41.363609Z","shell.execute_reply.started":"2022-08-09T14:11:41.356633Z","shell.execute_reply":"2022-08-09T14:11:41.362650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model object creation and initialization\nmodel = ClipEncoder()\nmodel.reducer_pca.set(pca.mean_, pca.components_)\nmodel.eval()\nprint('Model Created')","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:11:42.512934Z","iopub.execute_input":"2022-08-09T14:11:42.513360Z","iopub.status.idle":"2022-08-09T14:11:46.752601Z","shell.execute_reply.started":"2022-08-09T14:11:42.513324Z","shell.execute_reply":"2022-08-09T14:11:46.751372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#some test samples for usage test\nprint('Input: ', torch.randn(1, 3, 4545, 2324).shape, \n      'Output: ', model(torch.randn(1, 3, 4545, 2324)).shape)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:11:48.606721Z","iopub.execute_input":"2022-08-09T14:11:48.607889Z","iopub.status.idle":"2022-08-09T14:11:53.544021Z","shell.execute_reply.started":"2022-08-09T14:11:48.607829Z","shell.execute_reply":"2022-08-09T14:11:53.542275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#saving a model for inference\nsaved_model = torch.jit.script(model)\nsaved_model.save('saved_model.pt')\nwith ZipFile('submission.zip','w') as zip:           \n    zip.write('./saved_model.pt', arcname='saved_model.pt') ","metadata":{"execution":{"iopub.status.busy":"2022-08-09T14:12:02.114109Z","iopub.execute_input":"2022-08-09T14:12:02.114508Z","iopub.status.idle":"2022-08-09T14:12:09.310014Z","shell.execute_reply.started":"2022-08-09T14:12:02.114475Z","shell.execute_reply":"2022-08-09T14:12:09.308801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}