{"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":"!pip install pytorch_pretrained_vit","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:20.514851Z","iopub.execute_input":"2021-08-01T22:41:20.515198Z","iopub.status.idle":"2021-08-01T22:41:30.374141Z","shell.execute_reply.started":"2021-08-01T22:41:20.515120Z","shell.execute_reply":"2021-08-01T22:41:30.373207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q nnAudio","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:33.488350Z","iopub.execute_input":"2021-08-01T22:41:33.488736Z","iopub.status.idle":"2021-08-01T22:41:40.881320Z","shell.execute_reply.started":"2021-08-01T22:41:33.488694Z","shell.execute_reply":"2021-08-01T22:41:40.880425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom scipy.fftpack import fft, dct, fftshift\nimport os\nfrom tqdm import tqdm\nimport torch.nn.functional as F\nfrom scipy import stats \nfrom scipy import signal\nfrom scipy.signal import spectrogram,butter\n\nfrom pytorch_pretrained_vit import ViT\nfrom torchvision import transforms\n\nfrom nnAudio.Spectrogram import CQT1992v2","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:45.965021Z","iopub.execute_input":"2021-08-01T22:41:45.965380Z","iopub.status.idle":"2021-08-01T22:41:48.290901Z","shell.execute_reply.started":"2021-08-01T22:41:45.965344Z","shell.execute_reply":"2021-08-01T22:41:48.290038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"frame = pd.read_csv('/kaggle/input/g2net-gravitational-wave-detection/training_labels.csv')","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:52.590004Z","iopub.execute_input":"2021-08-01T22:41:52.590314Z","iopub.status.idle":"2021-08-01T22:41:52.968604Z","shell.execute_reply.started":"2021-08-01T22:41:52.590285Z","shell.execute_reply":"2021-08-01T22:41:52.967742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i=5\nidx = frame.iloc[i]['id']\nprint(frame.iloc[i]['target'])\ndirectory = os.path.join('/kaggle/input/g2net-gravitational-wave-detection/','train', idx[0],idx[1],idx[2],idx+'.npy')\n# print(frame)\nsample = np.load(directory)\nfilte = CQT1992v2(sr=2048, fmin=20, fmax=1024, hop_length=32)\n# waves = np.hstack(sample)\n# print(waves.shape)\nwaves = sample[0] / np.max(sample[0])\n\nwaves = torch.from_numpy(waves).float()\nz= filte(waves).permute(1,2,0)\nz = resize(z,(384,384))\nprint(z.shape)\ny = filte(waves)\n\n# plt.imshow(z.squeeze())\nplt.imshow(y.squeeze())\n\n# plt.plot(Sxx)\n# print(z)","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:53.716083Z","iopub.execute_input":"2021-08-01T22:41:53.716397Z","iopub.status.idle":"2021-08-01T22:41:53.970889Z","shell.execute_reply.started":"2021-08-01T22:41:53.716369Z","shell.execute_reply":"2021-08-01T22:41:53.969143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from skimage.transform import resize\nclass GWavesDataSet(Dataset):\n\n    def __init__(self, csv_file, root_dir = os.path.join('/kaggle/input/g2net-gravitational-wave-detection/','train')):\n\n        self.gwaves = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.mean = np.load('../input/meanstd-of-g2n/means.npy')\n        self.std = np.load('../input/meanstd-of-g2n/std.npy')\n        self.spectrogram = CQT1992v2(sr=2048, fmin=20, fmax=1024, hop_length=16)\n\n        \n    def __len__(self):\n        return len(self.gwaves)\n\n    def __getitem__(self, idx):\n        \n        file_id = frame.iloc[idx]['id']\n        wave_forms = np.load(os.path.join(self.root_dir, file_id[0], file_id[1], file_id[2],file_id+'.npy')).astype('float32')\n#         waves = np.hstack(wave_forms)\n        ligo_han = wave_forms[0]\n        ligo_liv = wave_forms[1]\n        vir = wave_forms[2]\n        ligo_han = ligo_han/ np.max(ligo_han)\n        ligo_liv = ligo_liv / np.max(ligo_liv)\n        vir = vir / np.max(vir)\n        ligo_han = self.spectrogram(torch.from_numpy(ligo_han).float())\n        ligo_liv = self.spectrogram(torch.from_numpy(ligo_liv).float())\n        vir = self.spectrogram(torch.from_numpy(vir).float())\n#         print(type(vir))\n#         waves = waves / np.max(waves)\n#         waves = torch.from_numpy(waves).float()\n        wave_forms = np.transpose(np.concatenate((ligo_han,ligo_liv,vir)),(1,2,0))\n        wave_forms = torch.from_numpy(resize(wave_forms, (384, 384)).transpose((2,0,1))).float()\n        target = int(frame.iloc[idx]['target'])\n\n        sample  = {'waveforms': wave_forms, 'target':torch.FloatTensor([target])}\n        return sample","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:57.935790Z","iopub.execute_input":"2021-08-01T22:41:57.936137Z","iopub.status.idle":"2021-08-01T22:41:57.948068Z","shell.execute_reply.started":"2021-08-01T22:41:57.936108Z","shell.execute_reply":"2021-08-01T22:41:57.947157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = GWavesDataSet('/kaggle/input/g2net-gravitational-wave-detection/training_labels.csv')\ndataset[154828]['waveforms'].shape","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:41:59.132516Z","iopub.execute_input":"2021-08-01T22:41:59.132919Z","iopub.status.idle":"2021-08-01T22:41:59.717305Z","shell.execute_reply.started":"2021-08-01T22:41:59.132885Z","shell.execute_reply":"2021-08-01T22:41:59.716494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def torch_nanmean(x):\n    num = torch.where(torch.isnan(x), torch.full_like(x, 0), torch.full_like(x, 1)).sum()\n    value = torch.where(torch.isnan(x), torch.full_like(x, 0), x).sum()\n    return value / num\n\ndef accuracy(output, target):\n    \"\"\"Computes the accuracy for multiple binary predictions\"\"\"\n    pred = output >= 0.5\n    truth = target >= 0.5\n    acc = pred.eq(truth).sum() / target.numel()\n    return acc\n","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:00.978032Z","iopub.execute_input":"2021-08-01T22:42:00.978396Z","iopub.status.idle":"2021-08-01T22:42:00.985035Z","shell.execute_reply.started":"2021-08-01T22:42:00.978367Z","shell.execute_reply":"2021-08-01T22:42:00.984076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViT('B_16_imagenet1k', pretrained=True)\nmodel.fc = nn.Linear(in_features=768, out_features=1, bias=True)\n# model = model.cuda()\nmodel.train()\nfor name, params in model.named_parameters():\n#     print(name)\n    if name == 'fc.weight' or name == 'fc.bias':\n        params.requires_grad= True\n        print(name)\n        print(params.requires_grad)\n    else:\n#         print('hi')\n        params.requires_grad= False\n#         print(params.required_grad)\n        \nmodel = model.cuda()\n\n# print(model)\n","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:01.273671Z","iopub.execute_input":"2021-08-01T22:42:01.274022Z","iopub.status.idle":"2021-08-01T22:42:12.499755Z","shell.execute_reply.started":"2021-08-01T22:42:01.273989Z","shell.execute_reply":"2021-08-01T22:42:12.498850Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class G2N(nn.Module):\n#     def __init__(self):\n#         super(G2N,self).__init__()\n    \n        \n#         self.conv1 = nn.Conv2d(1, 32, 3,1,1)\n#         self.conv1_bn = nn.BatchNorm2d(32)\n\n#         self.conv2 = nn.Conv2d(32, 32, 3,1,1)\n#         self.conv2_bn = nn.BatchNorm2d(32)\n\n#         self.conv3 = nn.Conv2d(32, 64, 3,1,1)\n#         self.conv3_bn = nn.BatchNorm2d(64)\n        \n#         self.conv4 = nn.Conv2d(64, 64, 3,1,1)\n#         self.conv4_bn = nn.BatchNorm2d(64)\n        \n#         self.conv4 = nn.Conv2d(64, 64, 3,1,1)\n#         self.conv4_bn = nn.BatchNorm2d(64)\n        \n#         self.ff1 = nn.Linear(69*193*64,512)\n#         self.ff1_bn = nn.BatchNorm1d(num_features=512)\n#         self.ff2 = nn.Linear(512,1024)\n#         self.ff2_bn = nn.BatchNorm1d(num_features=1024)\n#         self.ff3 = nn.Linear(1024,2048)\n#         self.ff3_bn = nn.BatchNorm1d(num_features=2048)\n#         self.ff4 = nn.Linear(2048,512)\n#         self.ff4_bn = nn.BatchNorm1d(num_features=512)\n#         self.ff5 = nn.Linear(512,32)\n#         self.ff5_bn = nn.BatchNorm1d(num_features=32)\n#         self.ff6 = nn.Linear(32,1)\n        \n\n        \n#     def forward(self, sample):\n#         x = F.relu(self.conv1_bn(self.conv1(sample)))\n#         x = F.relu(self.conv2_bn(self.conv2(x)))\n#         x = F.relu(self.conv3_bn(self.conv3(x)))\n#         x = F.relu(self.conv4_bn(self.conv4(x)))\n        \n#         x = x.reshape(-1,64*193*69)\n                   \n#         x = F.relu(self.ff1_bn(self.ff1(x)))\n#         x = F.relu(self.ff2(x))\n#         x = F.relu(self.ff3(x))\n#         x = F.relu(self.ff4_bn(self.ff4(x)))\n#         x = F.relu(self.ff5(x))\n#         x = torch.sigmoid(self.ff6(x))\n        \n#         return x        ","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:12.501245Z","iopub.execute_input":"2021-08-01T22:42:12.501563Z","iopub.status.idle":"2021-08-01T22:42:12.507113Z","shell.execute_reply.started":"2021-08-01T22:42:12.501525Z","shell.execute_reply":"2021-08-01T22:42:12.506386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = G2N()\ndataloader = torch.utils.data.DataLoader(dataset, batch_size=8,\n                                          shuffle=True, num_workers=8)\n# len(dataloader)\n# z = list()\n# for data in tqdm(dataloader):\n#     z.append(data['waveforms'])\n# d = torch.stack(z)\n# d = d.reshape(35000*16,300)\n# mean = torch.mean(d,1).numpy()\n# std = torch.std(d,1).numpy()\n# print(mean)\n# print(std)\n# np.save('means.npy',mean)\n# np.save('std.npy',std)","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:12.508854Z","iopub.execute_input":"2021-08-01T22:42:12.509323Z","iopub.status.idle":"2021-08-01T22:42:12.527754Z","shell.execute_reply.started":"2021-08-01T22:42:12.509282Z","shell.execute_reply":"2021-08-01T22:42:12.526955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\n# model = G2N().cuda()\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.1)","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:12.529400Z","iopub.execute_input":"2021-08-01T22:42:12.529823Z","iopub.status.idle":"2021-08-01T22:42:12.543106Z","shell.execute_reply.started":"2021-08-01T22:42:12.529782Z","shell.execute_reply":"2021-08-01T22:42:12.542329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model.load_state_dict(torch.load('../input/attention-is-all-space-time-need/check_point.pth'))\nimport os\n# os.mkdir(\"/kaggle/working/models\")\nfor epoch in range(50):  # loop over the dataset multiple times\n#     torch.save(model.state_dict(),'/kaggle/working/models/check_pointyy.pth')\n    progrss_loader = tqdm(dataloader)\n    running_loss = 0.0\n    for i, data in enumerate(progrss_loader):\n        # get the inputs; data is a list of [inputs, labels]\n#         print(data['target'].shape)\n        inputs, labels = data['waveforms'].cuda(), data['target'].cuda()\n\n        # zero the parameter gradients\n        optimizer.zero_grad()\n\n        # forward + backward + optimize\n        outputs = model(inputs)\n        labels = labels\n#         print(labels.shape)\n\n#         print(outputs.shape)\n        loss = criterion(outputs, labels)\n        to_logits  = torch.sigmoid(outputs)\n        acc = accuracy(to_logits.squeeze(1), labels.squeeze(1))\n        loss.backward()\n        optimizer.step()\n\n        # print statistics\n        running_loss += loss.item()\n        if (i+1) % 4 == 0:    # print every 2000 mini-batches\n            progrss_loader.set_description('[accuracy: %d, epoch: %5d] loss: %.3f' %\n                  (acc*100, epoch+1, running_loss / 4))\n#             print('[%d, %5d] loss: %.3f' %\n#                   (epoch + 1, i + 1, running_loss / 4))\n            running_loss = 0.0\nprint('Finished Training')","metadata":{"execution":{"iopub.status.busy":"2021-08-01T22:42:14.072096Z","iopub.execute_input":"2021-08-01T22:42:14.072416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"raw","source":"","metadata":{"execution":{"iopub.status.busy":"2021-07-28T12:21:52.673711Z","iopub.execute_input":"2021-07-28T12:21:52.674083Z"}}},{"cell_type":"code","source":"print(model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}