{"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":"markdown","source":"I used AIN DAO resources for fft-conv-pytorch.\n - [proposal link](https://dao.ainetwork.ai/t/proposal-fft-convolution-development/94)\n - fft convolution source code: https://github.com/klae01/fft-conv-pytorch\n \nThe idea is that larger kernels will focus on the shape of the wave rather than the texture, resulting in a higher accuracy. This idea came from this [research paper](https://arxiv.org/abs/2203.06717)","metadata":{}},{"cell_type":"code","source":"!pip install timm  \n!pip install git+https://github.com/klae01/fft-conv-pytorch@v0.1.0","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:47:36.422524Z","iopub.execute_input":"2022-12-31T17:47:36.423393Z","iopub.status.idle":"2022-12-31T17:48:05.619176Z","shell.execute_reply.started":"2022-12-31T17:47:36.423293Z","shell.execute_reply":"2022-12-31T17:48:05.617936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc, glob, os, tqdm\nfrom threading import Semaphore\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom scipy.stats import norm\nfrom timm import create_model\nfrom fft_conv_pytorch.functional import fft_conv","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:05.622409Z","iopub.execute_input":"2022-12-31T17:48:05.623065Z","iopub.status.idle":"2022-12-31T17:48:07.96924Z","shell.execute_reply.started":"2022-12-31T17:48:05.623022Z","shell.execute_reply":"2022-12-31T17:48:07.96803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sema = Semaphore(8)\ndef normalize(X):\n    X = (X[..., None].view(X.real.dtype) ** 2).sum(-1)\n    POS = int(X.size * 0.99903)\n    EXP = norm.ppf((POS + 0.4) / (X.size + 0.215))\n    scale = np.partition(X.flatten(), POS, -1)[POS]\n    X /= scale / EXP.astype(scale.dtype) ** 2\n    return X","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:07.971037Z","iopub.execute_input":"2022-12-31T17:48:07.971557Z","iopub.status.idle":"2022-12-31T17:48:07.980747Z","shell.execute_reply.started":"2022-12-31T17:48:07.971524Z","shell.execute_reply":"2022-12-31T17:48:07.979121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dataload(filepath):\n    sema.acquire()\n    astime = np.full([2, 360, 5760], np.nan, dtype=np.float32)\n    with h5py.File(filepath, \"r\") as f:\n        fid, _ = os.path.splitext(os.path.split(filepath)[1])\n        HT = (np.asarray(f[fid][\"H1\"][\"timestamps_GPS\"]) / 1800).round().astype(np.int64)\n        LT = (np.asarray(f[fid][\"L1\"][\"timestamps_GPS\"]) / 1800).round().astype(np.int64)\n        MIN = min(HT.min(), LT.min()); HT -= MIN; LT -= MIN\n        H1 = normalize(np.asarray(f[fid][\"H1\"][\"SFTs\"], np.complex128))\n        valid = HT < 5760; astime[0][:, HT[valid]] = H1[:, valid]\n        L1 = normalize(np.asarray(f[fid][\"L1\"][\"SFTs\"], np.complex128))\n        valid = LT < 5760; astime[1][:, LT[valid]] = L1[:, valid]\n#     gc.collect()\n    return fid, astime, H1.mean(), L1.mean()","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:07.982185Z","iopub.execute_input":"2022-12-31T17:48:07.982765Z","iopub.status.idle":"2022-12-31T17:48:07.993421Z","shell.execute_reply.started":"2022-12-31T17:48:07.982727Z","shell.execute_reply":"2022-12-31T17:48:07.992379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LargeKernel_debias(nn.Conv2d):\n    def forward(self, input: torch.Tensor):\n        finput = input.flatten(0, 1)[:, None]\n        target = abs(self.weight)\n        target = target / target.sum((-1, -2), True)\n        joined_kernel = torch.cat([self.weight, target], 0)\n        reals = target.new_zeros(\n            [1, 1] + [s + p * 2 for p, s in zip(self.padding, input.shape[-2:])]\n        )\n        reals[\n            [slice(None)] * 2 + [slice(p, -p) if p != 0 else slice(None) for p in self.padding]\n        ].fill_(1)\n        output, power = fft_conv(\n            finput, joined_kernel, padding=self.padding\n        ).chunk(2, 1)\n        ratio = torch.div(*fft_conv(reals, joined_kernel).chunk(2, 1))\n        output.sub_(power.mul_(ratio))\n        return output.unflatten(0, input.shape[:2]).flatten(1, 2)","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:07.996473Z","iopub.execute_input":"2022-12-31T17:48:07.996837Z","iopub.status.idle":"2022-12-31T17:48:08.00651Z","shell.execute_reply.started":"2022-12-31T17:48:07.996801Z","shell.execute_reply":"2022-12-31T17:48:08.005429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess(num, input, H1, L1):\n    input = torch.from_numpy(input).to(\"cuda\", non_blocking=True)\n    rescale = torch.tensor([[H1, L1]]).to(\"cuda\", non_blocking=True)\n    tta = (\n        torch.randn(\n            [num, *input.shape, 2], device=input.device, dtype=torch.float32\n        )\n        .square_()\n        .sum(-1)\n    )\n    tta *= rescale[..., None, None] / 2\n    valid = ~torch.isnan(input); tta[:, valid] = input[valid].float()\n    return tta","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:08.008101Z","iopub.execute_input":"2022-12-31T17:48:08.008523Z","iopub.status.idle":"2022-12-31T17:48:08.018314Z","shell.execute_reply.started":"2022-12-31T17:48:08.008485Z","shell.execute_reply":"2022-12-31T17:48:08.017333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(path):\n    model = create_model(\n        \"tf_efficientnetv2_b0\",\n        in_chans=32,\n        num_classes=2,\n    )\n    state_dict = torch.load(path)\n    C, _, H, W = state_dict[\"conv_stem.2.weight\"].shape\n    model.conv_stem = nn.Sequential(\n        nn.Identity(),\n        nn.AvgPool2d((1, 9), (1, 8), (0, 4), count_include_pad=False),\n        LargeKernel_debias(1, C, [H, W], 1, [H//2, W//2], 1, 1, False),\n        model.conv_stem,\n    )\n    model.load_state_dict(state_dict)\n    model.cuda().eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:08.019555Z","iopub.execute_input":"2022-12-31T17:48:08.02004Z","iopub.status.idle":"2022-12-31T17:48:08.029888Z","shell.execute_reply.started":"2022-12-31T17:48:08.019972Z","shell.execute_reply":"2022-12-31T17:48:08.02895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef inference(model, path):\n    file_path = glob.glob(os.path.join(path, \"*.hdf5\"))\n    FID, RES = [], []\n    with ThreadPoolExecutor(2) as pool:\n        for fid, input, H1, L1 in tqdm.tqdm(pool.map(dataload, sorted(file_path)), total=len(file_path)):\n            sema.release()\n            FID += [fid]\n            RES += [\n                sum(\n                    model(\n                        preprocess(32, input, H1, L1)\n                    ).softmax(-1)[..., 1].mean(0)\n                    for _ in range(2)\n                ) / 2\n            ]\n    return FID, torch.stack(RES, 0).cpu().float().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:08.03137Z","iopub.execute_input":"2022-12-31T17:48:08.031706Z","iopub.status.idle":"2022-12-31T17:48:08.04159Z","shell.execute_reply.started":"2022-12-31T17:48:08.031673Z","shell.execute_reply":"2022-12-31T17:48:08.040558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    model = get_model(\"../input/g2netmodel/group0/model_best.pth\")\n    fid, infer = inference(\n        model, \"../input/g2net-detecting-continuous-gravitational-waves/test\"\n    )\n    result = pd.DataFrame.from_dict({\"id\": fid, \"target\": infer})\n    result.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-12-31T17:48:08.043024Z","iopub.execute_input":"2022-12-31T17:48:08.043444Z","iopub.status.idle":"2022-12-31T20:02:33.770151Z","shell.execute_reply.started":"2022-12-31T17:48:08.043411Z","shell.execute_reply":"2022-12-31T20:02:33.76915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nfor weight in model.conv_stem[2].weight.flatten(0, 1).detach().cpu().numpy():\n    plt.matshow(weight, cmap=\"seismic\", vmin=-0.7, vmax=0.7)\n    plt.title(f\"{weight.min():.3f}/{weight.mean():.3f}/{np.median(weight):.3f}/{weight.max():.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-31T20:02:33.771594Z","iopub.execute_input":"2022-12-31T20:02:33.772617Z","iopub.status.idle":"2022-12-31T20:02:38.794326Z","shell.execute_reply.started":"2022-12-31T20:02:33.772578Z","shell.execute_reply":"2022-12-31T20:02:38.793197Z"},"trusted":true},"execution_count":null,"outputs":[]}]}