{"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 os\nimport json\nimport random\nimport collections\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-08-26T08:00:27.361216Z","iopub.execute_input":"2021-08-26T08:00:27.361626Z","iopub.status.idle":"2021-08-26T08:00:28.352761Z","shell.execute_reply.started":"2021-08-26T08:00:27.361547Z","shell.execute_reply":"2021-08-26T08:00:28.351968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TPU setup","metadata":{}},{"cell_type":"code","source":"!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n!python pytorch-xla-env-setup.py --version nightly --apt-packages libomp5 libopenblas-dev","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch_xla\nimport torch_xla.debug.metrics as met\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.utils.utils as xu\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.test.test_utils as test_utils\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = xm.xla_device()\nt1 = torch.randn(2, 2).to(device)\nt2 = torch.randn(2, 2).to(device)\nprint(t1+t2)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import libraries","metadata":{}},{"cell_type":"code","source":"!pip install efficientnet_pytorch -qq\n\n!pip install -q nnAudio -qq\nimport torch\nfrom nnAudio.Spectrogram import CQT1992v2, CQT2010v2\n\nimport time\n\nimport torch\nfrom torch import nn\nfrom torch.utils import data as torch_data\nfrom sklearn import model_selection as sk_model_selection\nfrom torch.nn import functional as torch_functional\nfrom torch.autograd import Variable\nimport efficientnet_pytorch\nfrom tqdm.auto import tqdm\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nfrom torchaudio.functional import lfilter\nfrom torch.fft import fft, rfft, ifft\nimport numpy as np\nimport torchvision\n\n\n\nfrom sklearn.metrics import roc_auc_score\n\nfrom sklearn.model_selection import StratifiedKFold\n\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:28.354224Z","iopub.execute_input":"2021-08-26T08:00:28.35458Z","iopub.status.idle":"2021-08-26T08:00:47.230117Z","shell.execute_reply.started":"2021-08-26T08:00:28.354543Z","shell.execute_reply":"2021-08-26T08:00:47.229147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#sys.path.append('../input/pytorch-swa')\n#import swa","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:47.23355Z","iopub.execute_input":"2021-08-26T08:00:47.233806Z","iopub.status.idle":"2021-08-26T08:00:47.236878Z","shell.execute_reply.started":"2021-08-26T08:00:47.233779Z","shell.execute_reply":"2021-08-26T08:00:47.236104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv(\"../input/g2net-gravitational-wave-detection/sample_submission.csv\")\ntrain_df = pd.read_csv(\"../input/g2net-gravitational-wave-detection/training_labels.csv\")\ntrain_df_pred = pd.read_csv(\"../input/train-pred-cqt-v10/train_preds_CQT_V10.csv\")","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:47.238846Z","iopub.execute_input":"2021-08-26T08:00:47.239436Z","iopub.status.idle":"2021-08-26T08:00:49.152587Z","shell.execute_reply.started":"2021-08-26T08:00:47.239359Z","shell.execute_reply":"2021-08-26T08:00:49.15171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['preds'] = train_df_pred['target']\nweight = 0.5\ntrain_df['soft_target'] = train_df['preds']*weight + train_df['target']*(1-weight)\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.153784Z","iopub.execute_input":"2021-08-26T08:00:49.154103Z","iopub.status.idle":"2021-08-26T08:00:49.207266Z","shell.execute_reply.started":"2021-08-26T08:00:49.154071Z","shell.execute_reply":"2021-08-26T08:00:49.206526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_pred","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.208486Z","iopub.execute_input":"2021-08-26T08:00:49.208806Z","iopub.status.idle":"2021-08-26T08:00:49.223456Z","shell.execute_reply.started":"2021-08-26T08:00:49.208772Z","shell.execute_reply":"2021-08-26T08:00:49.2225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define config","metadata":{}},{"cell_type":"code","source":"# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nclass CFG:\n    \n    TRAIN = True\n    \n    WARM_START = True\n    \n    EPOCHS = 7\n    \n    # batch size\n    BATCH = 128\n    \n    lr = 1e-2\n    \n    n_fold = 5\n    fold = 0\n    \n    # scheduler_params\n    scheduler='CosineAnnealingLR'\n    T_max=EPOCHS*2 # CosineAnnealingLR\n    T_0=3 # CosineAnnealingWarmRestarts\n    min_lr=1e-6\n    schedulerStepFreq = 3500*128/BATCH/2\n    \n    \n    # Parameters CWT\n    cwt_params = {'fs':2048, 'lower_freq':10, 'upper_freq': 500, \n                  'n_scales':81, 'wavelet_width':1, 'stride':12, 'border_crop':0, 'train_width':True}\n    \n    cqt_params = {'sr':2048, 'fmin':20, 'fmax':512, 'hop_length':32, 'n_bins':69, 'bins_per_octave': 8}\n    # cqt_params = {'sr':2048, 'fmin':20, 'fmax':512, 'hop_length':32, 'bins_per_octave': 25, 'norm':1}\n    \n    BPfilter = True\n    \n    RESIZE = [128,128] #False\n    \n    \n    \n    \n    # Post Proc Option\n    PREPROC = 'Q_transform'\n    \n    # scale:linear or log\n    SCALE = 'linear'\n    \n    DEBUG = False\n    \n    SMALL_TRAIN_SET = False\n    \n    seed = 42\n    \n    model_name = 'tf_efficientnet_b0' #'tf_efficientnet_b4' #'efficientnet-b7'\n    pretrained = False\n    unfreezeStep = 100 # set 0 for no freezing\n    \n    useSoftLabels = True\n    \n    \nif CFG.DEBUG:\n    CFG.EPOCHS = 2\n    train_df = train_df.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)\nelif CFG.SMALL_TRAIN_SET:\n    CFG.EPOCHS = 4\n    train_df = train_df.sample(n=CFG.BATCH*500, random_state=CFG.seed).reset_index(drop=True)\nelif CFG.WARM_START:\n    CFG.EPOCHS = 3\n    CFG.lr = 2e-3\n\nif not CFG.pretrained and not CFG.WARM_START:\n    unfreezeStep = 10\n    \n    ","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.225067Z","iopub.execute_input":"2021-08-26T08:00:49.225783Z","iopub.status.idle":"2021-08-26T08:00:49.278313Z","shell.execute_reply.started":"2021-08-26T08:00:49.225718Z","shell.execute_reply":"2021-08-26T08:00:49.277491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\ndef set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\n\nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.281364Z","iopub.execute_input":"2021-08-26T08:00:49.281699Z","iopub.status.idle":"2021-08-26T08:00:49.293348Z","shell.execute_reply.started":"2021-08-26T08:00:49.281667Z","shell.execute_reply":"2021-08-26T08:00:49.292547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data retrieving and related functions","metadata":{}},{"cell_type":"code","source":"from scipy import signal \n\nfrom scipy import signal\n\nbHP, aHP = signal.butter(8, (25, 500), btype='bandpass', fs= 2048)\ndef filterSig(waves, a=aHP, b=bHP, axis = 1):\n    '''Apply a 20Hz high pass filter to the three events'''\n    if not CFG.BPfilter:\n        return waves\n    return signal.filtfilt(b, a, waves, axis = axis) #lfilter introduces a larger spike around 20hz\n\nclass DataRetriever(torch_data.Dataset):\n    def __init__(self, paths, targets):\n        self.paths = paths\n        self.targets = targets\n          \n    def __len__(self):\n        return len(self.paths)\n    \n    def __get_qtransform(self, x):\n        image = x / np.max(x,axis=1,keepdims = True)\n        if CFG.DEBUG:\n            image = filterSig(image).copy()\n        # image = image / np.max(np.abs(image),axis=1,keepdims = True)\n        # image is [chan x time]\n        image = torch.tensor(image).float()\n        return image\n\n    \n    def __getitem__(self, index):\n        #file_path = convert_image_id_2_path(self.paths[index])\n        file_path = self.paths[index]\n        x = np.load(file_path)\n        image = self.__get_qtransform(x)\n        \n        y = torch.tensor(self.targets[index], dtype=torch.float)\n            \n        return {\"X\": image, \"y\": y}\n    \n    \nclass TestDataRetriever(torch_data.Dataset):\n    def __init__(self, paths):\n        self.paths = paths\n        \n        self.q_transform = CQT1992v2(\n            sr=2048, fmin=20, fmax=1024, hop_length=32\n        ) if CFG.PREPROC == 'Q_transform' else None\n        \n          \n          \n    def __len__(self):\n        return len(self.paths)\n    \n    def __get_qtransform(self, x):\n        image = x / np.max(x,axis=1,keepdims = True)\n        if CFG.DEBUG:\n            image = filterSig(image).copy()\n        # image = image / np.max(np.abs(image),axis=1,keepdims = True)\n        # image is [chan x time]\n        image = torch.tensor(image).float()\n        return image\n    \n    def __getitem__(self, index):\n        # file_path = convert_image_id_2_path(self.paths[index], is_train=False)\n        file_path = self.paths[index]\n        x = np.load(file_path)\n        image = self.__get_qtransform(x)\n            \n        return {\"X\": image, \"id\": self.paths[index]}","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.295259Z","iopub.execute_input":"2021-08-26T08:00:49.295618Z","iopub.status.idle":"2021-08-26T08:00:49.312208Z","shell.execute_reply.started":"2021-08-26T08:00:49.295585Z","shell.execute_reply":"2021-08-26T08:00:49.311342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:\n    Fold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n    for n, (train_index, val_index) in enumerate(Fold.split(train_df, train_df['target'])):\n        train_df.loc[val_index, 'fold'] = int(n)\n    train_df['fold'] = train_df['fold'].astype(int)\n    display(train_df.groupby(['fold', 'target']).size())\n\n    df_train = train_df.loc[train_df['fold'] != CFG.fold,:]\n    df_valid= train_df.loc[train_df['fold'] == CFG.fold,:]\nelse:\n    df_train, df_valid = sk_model_selection.train_test_split(\n    train_df, \n    test_size=0.2, \n    random_state=42, \n    stratify=train_df[\"target\"],\n    )\n\n\n","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.314349Z","iopub.execute_input":"2021-08-26T08:00:49.31507Z","iopub.status.idle":"2021-08-26T08:00:49.745022Z","shell.execute_reply.started":"2021-08-26T08:00:49.315031Z","shell.execute_reply":"2021-08-26T08:00:49.744163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/train/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ndef get_test_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/test/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ndf_train['file_path'] = df_train['id'].apply(get_train_file_path)\ndf_valid['file_path'] = df_valid['id'].apply(get_train_file_path)\n\nsubmission['file_path'] = submission['id'].apply(get_test_file_path)\n\ntrain_data_retriever = DataRetriever(\n    df_train['file_path'].values, \n    df_train[\"target\"].values, \n)\n\ntrain_data_retriever_soft = DataRetriever(\n    df_train['file_path'].values, \n    df_train[\"soft_target\"].values, \n)\n\nvalid_data_retriever = DataRetriever(\n    df_valid['file_path'].values, \n    df_valid[\"target\"].values,\n)\n\ntest_data_retriever = TestDataRetriever(\n    submission[\"file_path\"].values, \n)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:49.746348Z","iopub.execute_input":"2021-08-26T08:00:49.746686Z","iopub.status.idle":"2021-08-26T08:00:50.471585Z","shell.execute_reply.started":"2021-08-26T08:00:49.74665Z","shell.execute_reply":"2021-08-26T08:00:50.470704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = torch_data.DataLoader(\n    train_data_retriever,\n    batch_size=CFG.BATCH,\n    shuffle=True,\n    num_workers=12,\n)\n\ntrain_loader_soft = torch_data.DataLoader(\n    train_data_retriever_soft,\n    batch_size=CFG.BATCH,\n    shuffle=True,\n    num_workers=12,\n)\n\nvalid_loader = torch_data.DataLoader(\n    valid_data_retriever, \n    batch_size=CFG.BATCH*4,\n    shuffle=False,\n    num_workers=8,\n)\n\ntest_loader = torch_data.DataLoader(\n    test_data_retriever,\n    batch_size=CFG.BATCH*4,\n    shuffle=False,\n    num_workers=8,\n)\n\n","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.472825Z","iopub.execute_input":"2021-08-26T08:00:50.473332Z","iopub.status.idle":"2021-08-26T08:00:50.480677Z","shell.execute_reply.started":"2021-08-26T08:00:50.473293Z","shell.execute_reply":"2021-08-26T08:00:50.479873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"\"\"\" Triplet Attention Module\nImplementation of triplet attention module from https://arxiv.org/abs/2010.03045\n(slightly) Modified from official implementation: https://github.com/LandskapeAI/triplet-attention\nOriginal license:\nMIT License\nCopyright (c) 2020 LandskapeAI\nPermission is hereby granted, free of charge, to any person obtaining a copy\nof this software and associated documentation files (the \"Software\"), to deal\nin the Software without restriction, including without limitation the rights\nto use, copy, modify, merge, publish, distribute, sublicense, and/or sell\ncopies of the Software, and to permit persons to whom the Software is\nfurnished to do so, subject to the following conditions:\nThe above copyright notice and this permission notice shall be included in all\ncopies or substantial portions of the Software.\nTHE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\nIMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\nFITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\nAUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\nLIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\nOUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\nSOFTWARE.\n\"\"\"\n\nimport torch\nfrom torch import nn as nn\nfrom timm.models.layers import ConvBnAct\n\nclass ZPool(nn.Module):\n    def forward(self, x):\n        return torch.cat( (torch.max(x,1)[0].unsqueeze(1), torch.mean(x,1).unsqueeze(1)), dim=1 )\n\nclass AttentionGate(nn.Module):\n    def __init__(self, kernel_size=7):\n        super(AttentionGate, self).__init__()\n        self.zpool = ZPool()\n        self.conv = ConvBnAct(2, 1, kernel_size=kernel_size, stride=1, padding=(kernel_size-1) // 2, apply_act=False)\n    \n    def forward(self, x):\n        x_out = self.conv(self.zpool(x))\n        scale = torch.sigmoid_(x_out) \n        return x * scale\n\nclass TripletAttention(nn.Module):\n    def __init__(self, no_spatial=False):\n        super(TripletAttention, self).__init__()\n        self.cw = AttentionGate()\n        self.hc = AttentionGate()\n        self.hw = nn.Identity() if no_spatial else AttentionGate()\n        self.no_spatial = no_spatial\n\n    def forward(self, x):\n        x_perm1 = x.permute(0,2,1,3).contiguous()\n        x_out1 = self.cw(x_perm1)\n        x_out11 = x_out1.permute(0,2,1,3).contiguous()\n        x_perm2 = x.permute(0,3,2,1).contiguous()\n        x_out2 = self.hc(x_perm2)\n        x_out21 = x_out2.permute(0,3,2,1).contiguous()\n        x_out = self.hw(x)\n\n        x_out = (1/2) * (x_out11 + x_out21) if self.no_spatial else (1/3) * (x_out + x_out11 + x_out21)\n        return x_out","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.481997Z","iopub.execute_input":"2021-08-26T08:00:50.482353Z","iopub.status.idle":"2021-08-26T08:00:50.49667Z","shell.execute_reply.started":"2021-08-26T08:00:50.482319Z","shell.execute_reply":"2021-08-26T08:00:50.495899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Different heads\n\nclass BasicHead(nn.Module):   \n    def __init__(self,n_features):\n        super().__init__()\n        self.classifier = nn.Sequential(\n          nn.Dropout(0.5),\n          nn.Linear(in_features=n_features, out_features=256, bias=True),\n          nn.ReLU(),\n          # nn.Dropout(0.5), # p is probability of zeroing\n          nn.Linear(in_features=256, out_features=1, bias=True),\n        )\n        \n    def forward(self,x):\n        return self.classifier(x)\n    \nclass MultiDropoutHead(nn.Module):\n    def __init__(self,n_features):\n        super().__init__()\n        if False:\n            self.classifier = nn.Sequential(\n              nn.Linear(in_features=n_features, out_features=256, bias=True),\n              nn.ReLU(),\n              # nn.Dropout(0.5), # p is probability of zeroing\n              nn.Linear(in_features=256, out_features=1, bias=True),\n            )\n        else:\n            self.classifier = nn.Linear(in_features=n_features, out_features=1, bias=True)\n        self.dropout = lambda p: nn.Dropout(p)\n        \n    def forward(self,x):\n        return torch.mean(torch.stack([\n            self.classifier(self.dropout(p)(x))\n            for p in np.linspace(0.3,0.7, 5)\n        ], dim=0), dim=0)\n    \n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM,self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        x = self.gem(x, p=self.p, eps=self.eps)\n        x = torch.flatten(x,start_dim=1,end_dim=-1)\n        return x\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return torch.nn.functional.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.497843Z","iopub.execute_input":"2021-08-26T08:00:50.498407Z","iopub.status.idle":"2021-08-26T08:00:50.514108Z","shell.execute_reply.started":"2021-08-26T08:00:50.498373Z","shell.execute_reply":"2021-08-26T08:00:50.513387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model(CFG.model_name, pretrained=CFG.pretrained,features_only = True,out_indices = [1,2,3,4])\nmodel.feature_info.channels()\no = model(torch.randn(2, 3, 128,128))\nfor x in o:\n  print(x.shape)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.515397Z","iopub.execute_input":"2021-08-26T08:00:50.51573Z","iopub.status.idle":"2021-08-26T08:00:50.870212Z","shell.execute_reply.started":"2021-08-26T08:00:50.515696Z","shell.execute_reply":"2021-08-26T08:00:50.869174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model= timm.create_model(CFG.model_name, pretrained=CFG.pretrained,features_only = True)\n#model","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.871729Z","iopub.execute_input":"2021-08-26T08:00:50.872079Z","iopub.status.idle":"2021-08-26T08:00:50.875553Z","shell.execute_reply.started":"2021-08-26T08:00:50.872041Z","shell.execute_reply":"2021-08-26T08:00:50.874636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef whiten(signal):\n    # From here: https://www.kaggle.com/kevinmcisaac/g2net-spectral-whitening\n    length = signal.size(2)\n    hann = torch.hann_window(length, periodic=True, dtype=float).view(1,1,-1)\n    spec = fft(signal* hann, dim = 2)\n    mag = torch.sqrt(torch.real(spec*torch.conj(spec))) \n\n    return torch.real(ifft(spec/mag)) * np.sqrt(length/2)\n\ndef batch_preprocessing(X):\n    # X = whiten(X)\n    X = X.numpy()\n    if CFG.BPfilter:        \n        X = filterSig(X,axis=2).copy()\n    X = torch.tensor(X).float()\n    return X      \n        \n\nHead = MultiDropoutHead#BasicHead\n\nmodel_no = 1\n\nclass FeatureExtractorModel(nn.Module):\n    \n    def __init__(self):\n        super().__init__()\n        self.model = timm.create_model(CFG.model_name, pretrained=CFG.pretrained,features_only = True, out_indices = [2,3,4])\n        self.prepooler = nn.Sequential(nn.Conv2d(self.model.feature_info.channels()[-1], 1280, kernel_size=(1, 1), stride=(1, 1), bias=False), \n                                  nn.BatchNorm2d(1280, eps=0.001, momentum=0.1, affine=True, track_running_stats=True))\n        self.n_features = np.sum(self.model.feature_info.channels())+1280\n        self.poolers = []\n        for i in range(4):\n            self.poolers.append(nn.Sequential(nn.SiLU(),TripletAttention(),GeM()))\n        self.poolers = nn.ModuleList(self.poolers)\n        \n    def forward(self, x):\n        features = self.model(x)\n        features.append(self.prepooler(features[-1]))       \n        out = []\n        for i,feat in enumerate(features):\n            out.append(torch.squeeze(self.poolers[i](feat)))\n        return torch.cat(out,dim=1)\n            \n            \n        \n        \ndef Backbone():\n    if 'tf_efficientnet' in CFG.model_name:\n        if False:\n            model = timm.create_model(CFG.model_name, pretrained=CFG.pretrained)\n            n_features = model.classifier.in_features\n            model.classifier = nn.Identity()\n            if True:\n                model.global_pool = nn.Sequential(TripletAttention(),\n                                                  GeM())\n        else:\n            model = FeatureExtractorModel()\n            n_features = model.n_features\n            \n    elif 'efficientnet-' in CFG.model_name:\n        model = efficientnet_pytorch.EfficientNet.from_pretrained(CFG.model_name)\n        n_features = model._fc.in_features\n        model._fc = nn.Identity()\n    elif 'rexnet_' in CFG.model_name:\n        model = timm.create_model(CFG.model_name, pretrained=CFG.pretrained)\n        n_features = model.head.fc.in_features\n        model.head.fc = nn.Identity()\n        \n    return model, n_features\n\nif model_no == 1:\n    print('Selecting multi channel model ... ')\n    class Model(nn.Module):\n        def __init__(self, get_spectrogram = False):\n            self.get_spectrogram = get_spectrogram\n            super().__init__()\n            self.q_transform = CQT1992v2(\n                **CFG.cqt_params\n            )\n            if not self.get_spectrogram:\n                \n                self.model, n_features = Backbone()\n                self.head = Head(n_features)\n            if CFG.RESIZE:\n                self.resize = torchvision.transforms.Resize(CFG.RESIZE)\n\n        def freezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = False\n\n        def unfreezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = True\n\n        def forward(self, x):\n            # reshape from [batch by chan by time] [(batch x chan) by time]\n            batch_size = x.size(0)\n\n            x = torch.divide(x,torch.max(torch.abs(x),dim=2,keepdims = True)[0])\n            x = torch.reshape(x,(batch_size*3,-1))\n            x = self.q_transform(x)\n            # x = x[:,0:-1,0:-1]\n            x = x[:,0:-1,35:120]\n            size = list(x.size())\n            x = torch.reshape(x,(batch_size,3,size[1],size[2]))\n            if CFG.RESIZE:\n                x = self.resize(x)\n            if CFG.SCALE == 'log':\n                x = (torch.log10(x) + 1.)/1.5\n            else:\n                x = torch.clamp(x,max=2.5)-1\n                #x = torch.divide(x,torch.mean(x,dim=2,keepdims = True))\n            \n            \n            \n            if self.get_spectrogram:\n                return x\n            x = self.model(x)\n            out = self.head(x)\n            return out\nelif model_no == 2:        \n    print('Selecting single channel model ... ')\n    \n    class Model(nn.Module):\n        def __init__(self, get_spectrogram = False):\n            self.get_spectrogram = get_spectrogram\n            super().__init__()\n            self.q_transform = CQT1992v2(\n                **CFG.cqt_params\n            )\n            if not self.get_spectrogram:\n                self.model, n_features = Backbone()\n                self.head = Head(n_features)\n\n        def freezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = False\n\n        def unfreezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = True\n\n        def forward(self, x):\n            # reshape from [batch by chan by time] [(batch x chan) by time]\n            batch_size = x.size(0)\n\n            x = torch.divide(x,torch.max(torch.abs(x),dim=2,keepdims = True)[0])\n            x = torch.reshape(x,(batch_size*3,-1))\n            x = self.q_transform(x)\n            x = x[:,0:-1,0:-1]\n            if CFG.SCALE == 'log':\n                x = (torch.log10(x) + 1.)/1.5\n            else:\n                x = torch.clamp(x,max=2.5)-1\n            if self.get_spectrogram:\n                size = list(x.size())\n                x = torch.reshape(x,(batch_size,3,size[1],size[2]))\n                return x\n            \n            x = torch.unsqueeze(x,1)\n            \n            # x_mean = torch.mean(x,dim=1,keepdims = True)\n            # x = torch.stack([x,x_mean],dim=1)\n            \n            x = self.model(x)\n            size = list(x.size())\n            x = torch.reshape(x,(batch_size,-1,size[1]))\n            x = torch.max(x,dim=1,keepdims = False)[0]\n            out = self.head(x)\n            out = out\n            return out\n        \nelif model_no == 3:\n    print('Selecting experimental model ... ')\n    class Model(nn.Module):\n        def __init__(self, get_spectrogram = False):\n            self.get_spectrogram = get_spectrogram\n            super().__init__()\n            self.q_transform = CQT1992v2(\n                **CFG.cqt_params\n            )\n            if not self.get_spectrogram:\n                self.model, n_features = Backbone()\n                self.head = Head(2*n_features)\n\n        def freezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = False\n\n        def unfreezeModel(self):\n            for param in self.model.parameters():\n                param.requires_grad = True\n\n        def forward(self, x):\n            # reshape from [batch by chan by time] [(batch x chan) by time]\n            batch_size = x.size(0)\n\n            x = torch.divide(x,torch.max(torch.abs(x),dim=2,keepdims = True)[0])\n            x = torch.reshape(x,(batch_size*3,-1))\n            x = self.q_transform(x)\n            x = x[:,0:-1,0:-1]\n            if CFG.SCALE == 'log':\n                x = (torch.log10(x) + 1.)/1.5\n            else:\n                x = torch.clamp(x,max=2.5)-1\n            if self.get_spectrogram:\n                size = list(x.size())\n                x = torch.reshape(x,(batch_size,3,size[1],size[2]))\n                return x\n            \n            x = torch.unsqueeze(x,1)\n            \n            size = list(x.size())\n            x = torch.reshape(x,(batch_size,3,size[2],size[3]))\n            x_mean = torch.mean(x,dim=1,keepdims = True)\n            x = torch.cat([x,x_mean],dim=1)\n            x = torch.reshape(x,(batch_size*4,1,size[2],size[3]))\n            \n            x = self.model(x)\n            size = list(x.size())\n            x = torch.reshape(x,(batch_size,-1,size[1]))\n            x_mean = x[:,3,:].squeeze(1)\n            x = torch.max(x,dim=1,keepdims = False)[0]\n            # this will be [batch by (2 x n_features)]\n            x = torch.cat([x,x_mean],dim=1)\n            \n            out = self.head(x)\n            out = out\n            return out\n    \n\n\n","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.876979Z","iopub.execute_input":"2021-08-26T08:00:50.87732Z","iopub.status.idle":"2021-08-26T08:00:50.922301Z","shell.execute_reply.started":"2021-08-26T08:00:50.877269Z","shell.execute_reply":"2021-08-26T08:00:50.921349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nif True:\n    \n    modelTmp = Model(get_spectrogram = True)\n    for step, batch in enumerate(train_loader,1):\n        X = batch[\"X\"]\n        X = batch_preprocessing(X)\n        modelTmp.to(device)\n        X = X.to(device)\n        targets = batch[\"y\"].to(device)\n        outputs = modelTmp(X)\n        n = np.random.randint(32)\n        import matplotlib.pyplot as plt\n        tmp = outputs[n].cpu().numpy()\n        for i in range(3):\n            plt.figure()\n            plt.title('Target = '+ str(targets[n]))\n            plt.imshow((tmp[i,:,:]).squeeze())\n            plt.colorbar()\n        plt.figure()\n        plt.title('Target = '+ str(targets[n]))\n        plt.imshow(np.mean(tmp[:,:,:],axis=0).squeeze())\n        plt.colorbar()\n        print(tmp.shape)\n        break","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:00:50.923622Z","iopub.execute_input":"2021-08-26T08:00:50.923984Z","iopub.status.idle":"2021-08-26T08:01:00.340518Z","shell.execute_reply.started":"2021-08-26T08:00:50.923948Z","shell.execute_reply":"2021-08-26T08:01:00.339504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss related functions","metadata":{}},{"cell_type":"code","source":"class LossMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n\n    def update(self, val):\n        self.n += 1\n        # incremental update\n        self.avg = val / self.n + (self.n - 1) / self.n * self.avg\n\n        \nclass AccMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n        \n    def update(self, y_true, y_pred):\n        y_true = y_true.cpu().round().numpy().astype(int)\n        y_pred = y_pred.cpu().numpy() >= 0\n        last_n = self.n\n        self.n += len(y_true)\n        true_count = np.sum(y_true == y_pred)\n        # incremental update\n        self.avg = true_count / self.n + last_n / self.n * self.avg\n        \n","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:01:00.341995Z","iopub.execute_input":"2021-08-26T08:01:00.342351Z","iopub.status.idle":"2021-08-26T08:01:00.350494Z","shell.execute_reply.started":"2021-08-26T08:01:00.342312Z","shell.execute_reply":"2021-08-26T08:01:00.349466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer related functions","metadata":{}},{"cell_type":"code","source":"def get_scheduler(optimizer):\n    if CFG.scheduler=='ReduceLROnPlateau':\n        scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n    elif CFG.scheduler=='CosineAnnealingLR':\n        scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n    elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n        scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n    return scheduler\n\nclass Trainer:\n    def __init__(\n        self, \n        model, \n        device, \n        optimizer, \n        criterion, \n        loss_meter, \n        score_meter,\n        use_swa = False\n    ):\n        self.model = model\n        # freeze model by default\n        self.model.freezeModel()\n        \n        self.device = device\n        self.use_swa = use_swa\n        self.optimizer = swa.SWA(optimizer) if self.use_swa else optimizer\n        self.criterion = criterion\n        self.loss_meter = loss_meter\n        self.score_meter = score_meter\n        self.scheduler = get_scheduler(optimizer)\n        self.learning_rate = self.scheduler.get_lr()\n        \n        self.best_valid_score = -np.inf\n        self.n_patience = 0\n        \n        self.messages = {\n            \"epoch\": \"[Epoch {}: {}] loss: {:.5f}, score: {:.5f}, auc_score: {:.5f}, time: {} s\",\n            \"checkpoint\": \"The score improved from {:.5f} to {:.5f}. Save model to '{}'\",\n            \"patience\": \"\\nValid score didn't improve last {} epochs.\"\n        }\n        self.training_step = 0\n        self.prevbatch = []\n        self.epoch = -1 \n    \n    def fit(self, epochs, train_loader, valid_loader, save_path, patience,train_loader_soft = False):        \n        for n_epoch in range(1, epochs + 1):\n            \n            self.epoch = n_epoch\n            \n            self.info_message(\"EPOCH: {}\", n_epoch)\n            \n            if self.epoch==1 or (not train_loader_soft):\n                train_loss, train_score, train_time = self.train_epoch(train_loader)\n            else:\n                train_loss, train_score, train_time = self.train_epoch(train_loader_soft)\n                \n            valid_loss, valid_score, valid_time, valid_rocauc = self.valid_epoch(valid_loader)\n            \n            self.info_message(\n                self.messages[\"epoch\"], \"Train\", n_epoch, train_loss, train_score, 0, train_time\n            )\n            \n            self.info_message(\n                self.messages[\"epoch\"], \"Valid\", n_epoch, valid_loss, valid_score, valid_rocauc, valid_time\n            )\n            \n\n            if self.best_valid_score < valid_score:\n                self.info_message(\n                    self.messages[\"checkpoint\"], self.best_valid_score, valid_score, save_path\n                )\n                self.best_valid_score = valid_score\n                self.save_model(n_epoch, save_path)\n                self.n_patience = 0\n            else:\n                self.n_patience += 1\n            \n            if self.n_patience >= patience:\n                self.info_message(self.messages[\"patience\"], patience)\n                break\n        if self.use_swa:\n            self.optimizer.bn_update(train_loader, self.model)\n            self.optimizer.swap_swa_sgd()\n        \n    def train_epoch(self, train_loader):\n        self.model.train()\n        t = time.time()\n        train_loss = self.loss_meter()\n        train_score = self.score_meter()\n        \n        for step, batch in enumerate(tqdm(train_loader),1):\n            \n            if self.training_step == CFG.unfreezeStep:\n                self.model.unfreezeModel()\n            \n            X = batch[\"X\"]\n            if self.prevbatch:\n                prevX = self.prevbatch[\"X\"]\n                prevY = self.prevbatch['y']\n                rndNum = np.random.rand()\n                if rndNum<0.5:\n                    # only keep prevX where there is no wave\n                    prevX = torch.where(prevY.view(-1,1,1)>0.5,X,prevX)\n                    # weight for prevX is at most 0.5, and not replaced\n                    # when there is a wave\n                    X = (1-rndNum)*X + rndNum*prevX \n\n            self.prevbatch = batch.copy()  \n            X = batch_preprocessing(X)\n            X = X.to(self.device)\n            targets = batch[\"y\"].to(self.device)\n            self.optimizer.zero_grad()\n            outputs = self.model(X).squeeze(1)\n            \n            loss = self.criterion(outputs, targets)\n            loss.backward()\n\n            train_loss.update(loss.detach().item())\n            train_score.update(targets, outputs.detach())\n\n            self.optimizer.step()\n            \n            _loss, _score = train_loss.avg, train_score.avg\n            \n            message = 'Train Step {}/{}, train_loss: {:.5f}, train_score: {:.5f}, learning_rate: {:.5f}/{:.5f}'\n            self.info_message(message, step, len(train_loader), _loss, _score, self.learning_rate[0], self.learning_rate[1],end=\"\\r\")\n            self.training_step += 1\n            \n            if self.training_step%CFG.schedulerStepFreq==0:\n                if isinstance(self.scheduler, CosineAnnealingLR):\n                    self.scheduler.step()\n                elif isinstance(self.scheduler, CosineAnnealingWarmRestarts):\n                    self.scheduler.step()\n                self.learning_rate = self.scheduler.get_lr()\n        # print('\\n Updated learning rate: '+ str(self.scheduler.get_lr()))\n        \n        return train_loss.avg, train_score.avg, int(time.time() - t)\n    \n    def valid_epoch(self, valid_loader,returnPred = False):\n        self.model.eval()\n        t = time.time()\n        valid_loss = self.loss_meter()\n        valid_score = self.score_meter()\n        \n        for step, batch in enumerate(valid_loader, 1):\n            y_pred = []\n            tgts = []\n            with torch.no_grad():\n                X = batch[\"X\"]  \n                X = batch_preprocessing(X)\n                X = X.to(self.device)\n                targets = batch[\"y\"].to(self.device)\n\n                outputs = self.model(X).squeeze(1)\n                loss = self.criterion(outputs, targets)\n\n                valid_loss.update(loss.detach().item())\n                valid_score.update(targets, outputs)\n                outputs = outputs\n                y_pred.extend(torch.sigmoid(outputs).cpu().numpy().squeeze())\n                tgts.extend(batch[\"y\"].numpy())\n                    \n            rocauc = roc_auc_score(tgts,y_pred)\n            _loss, _score = valid_loss.avg, valid_score.avg\n            message = 'Valid Step {}/{}, valid_loss: {:.5f}, valid_score: {:.5f},valid_roc_auc: {:.5f}'\n            self.info_message(message, step, len(valid_loader), _loss, _score, rocauc, end=\"\\r\")\n        if not returnPred:\n            return valid_loss.avg, valid_score.avg, int(time.time() - t), rocauc\n        else:\n            return y_pred, tgts\n    \n    def test_eval(self,test_loader):\n        y_pred = []\n        ids = []\n        for e, batch in enumerate(test_loader):\n            print(f\"{e}/{len(test_loader)}\", end=\"\\r\")\n            with torch.no_grad():\n                X = batch[\"X\"]\n                X = batch_preprocessing(X)\n                X = X.to(self.device)\n                outputs = self.model(X)\n                y_pred.extend(torch.sigmoid(outputs).cpu().numpy().squeeze())\n                ids.extend(batch[\"id\"])\n        return y_pred, ids\n    \n    def save_model(self, n_epoch, save_path):\n        torch.save(\n            {\n                \"model_state_dict\": self.model.state_dict(),\n                \"optimizer_state_dict\": self.optimizer.state_dict(),\n                \"best_valid_score\": self.best_valid_score,\n                \"n_epoch\": n_epoch,\n            },\n            save_path,\n        )\n    \n    @staticmethod\n    def info_message(message, *args, end=\"\\n\"):\n        print(message.format(*args), end=end)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:01:00.351954Z","iopub.execute_input":"2021-08-26T08:01:00.352306Z","iopub.status.idle":"2021-08-26T08:01:00.38908Z","shell.execute_reply.started":"2021-08-26T08:01:00.352254Z","shell.execute_reply":"2021-08-26T08:01:00.388246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel = Model()\nmodel.to(device)\nif (not CFG.TRAIN) or CFG.WARM_START:\n    startCheckpoint = torch.load(\"../input/CQTPT2V6/best-model.pth\")\n    checkpoint = torch.load(\"../input/CQTPT2V6/best-model.pth\")\n    model.load_state_dict(startCheckpoint[\"model_state_dict\"])\n\noptimizer = torch.optim.Adam([{\"params\": model.model.parameters(), \"lr\": CFG.lr},\n                              {\"params\": model.head.parameters(), \"lr\": CFG.lr/10}], \n                             lr=CFG.lr)\ncriterion = torch_functional.binary_cross_entropy_with_logits\n\ntrainer = Trainer(\n    model, \n    device, \n    optimizer, \n    criterion, \n    LossMeter, \n    AccMeter\n)\n\nif CFG.TRAIN:\n    history = trainer.fit(\n        CFG.EPOCHS, \n        train_loader, \n        valid_loader, \n        \"best-model.pth\", \n        400,\n        train_loader_soft = train_loader_soft if CFG.useSoftLabels else False\n    )\n    \n    y_pred_val,tgts = trainer.valid_epoch(valid_loader,returnPred = True)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:01:00.390482Z","iopub.execute_input":"2021-08-26T08:01:00.39117Z","iopub.status.idle":"2021-08-26T08:06:26.137014Z","shell.execute_reply.started":"2021-08-26T08:01:00.391133Z","shell.execute_reply":"2021-08-26T08:06:26.133717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.scatter(tgts,y_pred_val,1)\nplt.xlabel('targets')\nplt.ylabel('predictions')","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:06:26.138035Z","iopub.status.idle":"2021-08-26T08:06:26.138444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.TRAIN:\n    checkpoint = torch.load(\"best-model.pth\")\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nmodel.eval();","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:06:26.139407Z","iopub.status.idle":"2021-08-26T08:06:26.139915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ngc.collect()\n\n\n\ny_pred, ids = trainer.test_eval(test_loader)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:06:26.141045Z","iopub.status.idle":"2021-08-26T08:06:26.141636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(\"../input/g2net-gravitational-wave-detection/sample_submission.csv\")\nsubmission = pd.DataFrame({\"id\": submission['id'].values, \"target\": y_pred})\nsubmission.to_csv(\"model_submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:06:26.142776Z","iopub.status.idle":"2021-08-26T08:06:26.143333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2021-08-26T08:06:26.144522Z","iopub.status.idle":"2021-08-26T08:06:26.14507Z"},"trusted":true},"execution_count":null,"outputs":[]}]}