{"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":"# Starter using the Vision Transformer (ViT)\n\nWhat Transformer does:\n- dividing the spectrogram into patches\n- build patch embeddings\n- attention between different patches\n\nSince the default ViT needs 16x16 patches, the final images are padded...\n\nReference:\n\n- Yasufumi Nakama's (@yasufuminakama) spectrogram preprocessing notebooks and datasets:\n    * Train: [Notebook](https://www.kaggle.com/yasufuminakama/g2net-spectrogram-generation-train), [Dataset](https://www.kaggle.com/yasufuminakama/g2net-n-mels-128-train-images)\n    * Test: [Notebook](https://www.kaggle.com/yasufuminakama/g2net-spectrogram-generation-test), [Dataset](https://www.kaggle.com/yasufuminakama/g2net-n-mels-128-test-images)\n- @xhlulu 's pipeline: https://www.kaggle.com/xhlulu/g2net-rnn-starter-from-spectrogram","metadata":{}},{"cell_type":"code","source":"!pip install -q vit-keras","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:37:56.854143Z","iopub.execute_input":"2021-09-10T18:37:56.85459Z","iopub.status.idle":"2021-09-10T18:38:07.363627Z","shell.execute_reply.started":"2021-09-10T18:37:56.854486Z","shell.execute_reply":"2021-09-10T18:38:07.362193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nfrom tensorflow.keras import layers","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-09-10T18:38:08.244863Z","iopub.execute_input":"2021-09-10T18:38:08.245212Z","iopub.status.idle":"2021-09-10T18:38:14.516666Z","shell.execute_reply.started":"2021-09-10T18:38:08.245177Z","shell.execute_reply":"2021-09-10T18:38:14.515338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from vit_keras import vit","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:14.518948Z","iopub.execute_input":"2021-09-10T18:38:14.519596Z","iopub.status.idle":"2021-09-10T18:38:14.795138Z","shell.execute_reply.started":"2021-09-10T18:38:14.519509Z","shell.execute_reply":"2021-09-10T18:38:14.793981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLD = 0\nN_SPLITS = 5","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:14.796704Z","iopub.execute_input":"2021-09-10T18:38:14.79739Z","iopub.status.idle":"2021-09-10T18:38:14.803877Z","shell.execute_reply.started":"2021-09-10T18:38:14.79734Z","shell.execute_reply":"2021-09-10T18:38:14.802442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(tf.keras.utils.Sequence):\n    def __init__(self, df, directory, \n                 batch_size=32, \n                 random_state=1127802825, \n                 shuffle=True, target=True, ext='.npy'):\n        np.random.seed(random_state)\n        \n        self.directory = directory\n        self.df = df\n        self.shuffle = shuffle\n        self.target = target\n        self.batch_size = batch_size\n        self.ext = ext\n        \n        self.on_epoch_end()\n    \n    def __len__(self):\n        return np.ceil(self.df.shape[0] / self.batch_size).astype(int)\n    \n    def __getitem__(self, idx):\n        start_idx = idx * self.batch_size\n        batch = self.df[start_idx: start_idx + self.batch_size]\n        \n        signals = []\n\n        for fname in batch.id:\n            path = os.path.join(self.directory, fname + self.ext)\n            data = np.load(path)\n            signals.append(data)\n        \n        signals = np.stack(signals).astype('float32')\n        signals = tf.pad(signals, tf.constant([[0, 0], [2, 3,], [0, 0]]), \"SYMMETRIC\")\n        signals = tf.tile(tf.expand_dims(signals, axis=-1), multiples=[1,1,1,3])\n        \n        if self.target:\n            return signals, batch.target.values\n        else:\n            return signals\n    \n    def on_epoch_end(self):\n        if self.shuffle:\n            self.df = self.df.sample(frac=1).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:14.805938Z","iopub.execute_input":"2021-09-10T18:38:14.806447Z","iopub.status.idle":"2021-09-10T18:38:14.822983Z","shell.execute_reply.started":"2021-09-10T18:38:14.806401Z","shell.execute_reply":"2021-09-10T18:38:14.821336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vit_model = vit.vit_b16(\n        image_size = (32, 128),\n        activation = 'softmax',\n        pretrained = True,\n        include_top = False,\n        pretrained_top = False,\n        classes = 2)","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:14.825033Z","iopub.execute_input":"2021-09-10T18:38:14.825597Z","iopub.status.idle":"2021-09-10T18:38:25.30705Z","shell.execute_reply.started":"2021-09-10T18:38:14.825538Z","shell.execute_reply":"2021-09-10T18:38:25.305715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model():\n    inputs = layers.Input(shape=(32, 128, 3))\n\n    x = vit_model(inputs)\n    x = tf.keras.layers.BatchNormalization()(x)\n    x = layers.Dense(128, activation = tfa.activations.gelu)(x)\n    x = layers.Dense(1, activation=\"sigmoid\", name=\"sigmoid\")(x)\n\n    model = tf.keras.Model(inputs=inputs, outputs=x)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:46:28.519928Z","iopub.execute_input":"2021-09-10T18:46:28.520326Z","iopub.status.idle":"2021-09-10T18:46:28.529732Z","shell.execute_reply.started":"2021-09-10T18:46:28.520292Z","shell.execute_reply":"2021-09-10T18:46:28.528325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/g2net-gravitational-wave-detection/training_labels.csv')\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:25.329272Z","iopub.execute_input":"2021-09-10T18:38:25.329906Z","iopub.status.idle":"2021-09-10T18:38:25.875694Z","shell.execute_reply.started":"2021-09-10T18:38:25.329771Z","shell.execute_reply":"2021-09-10T18:38:25.874551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cv = StratifiedKFold(n_splits=N_SPLITS, random_state=1127802825, shuffle=True)\ncv_splits = cv.split(X=train, y=train['target'].values)\nfor _fold, (train_idx, valid_idx) in enumerate(cv_splits):\n    if _fold == FOLD:\n        break\n\ntrain_df = train.iloc[train_idx, :]\nvalid_df = train.iloc[valid_idx, :]","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:25.877573Z","iopub.execute_input":"2021-09-10T18:38:25.877972Z","iopub.status.idle":"2021-09-10T18:38:25.990208Z","shell.execute_reply.started":"2021-09-10T18:38:25.877941Z","shell.execute_reply":"2021-09-10T18:38:25.989088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dset = CustomDataset(\n    train_df, '../input/g2net-n-mels-128-train-images', batch_size=64)\n\nvalid_dset = CustomDataset(\n    valid_df, '../input/g2net-n-mels-128-train-images', batch_size=64, shuffle=False)\n\nsample = next(iter(train_dset))\nfor item in sample:\n    print(item.shape)","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:25.991979Z","iopub.execute_input":"2021-09-10T18:38:25.992537Z","iopub.status.idle":"2021-09-10T18:38:26.634445Z","shell.execute_reply.started":"2021-09-10T18:38:25.992487Z","shell.execute_reply":"2021-09-10T18:38:26.633189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_model()\nmodel.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4), \n              loss=\"binary_crossentropy\", \n              metrics=[tf.keras.metrics.AUC()])\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:46:34.818478Z","iopub.execute_input":"2021-09-10T18:46:34.81897Z","iopub.status.idle":"2021-09-10T18:46:36.875785Z","shell.execute_reply.started":"2021-09-10T18:46:34.818939Z","shell.execute_reply":"2021-09-10T18:46:36.874375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt = tf.keras.callbacks.ModelCheckpoint(\n    \"model_weights.h5\", save_best_only=True, save_weights_only=True,\n)\n\ntrain_history = model.fit(\n    train_dset, \n    epochs=8,\n    validation_data=valid_dset,\n    callbacks=[ckpt],\n    verbose=1\n)","metadata":{"execution":{"iopub.status.busy":"2021-09-10T18:38:38.838813Z","iopub.execute_input":"2021-09-10T18:38:38.839319Z","iopub.status.idle":"2021-09-10T18:45:58.908536Z","shell.execute_reply.started":"2021-09-10T18:38:38.839287Z","shell.execute_reply":"2021-09-10T18:45:58.905432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_weights('model_weights.h5')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv('../input/g2net-gravitational-wave-detection/sample_submission.csv')\n\ntest_dset = CustomDataset(\n    sub, \"../input/g2net-n-mels-128-test-images\", batch_size=64, target=False, shuffle=False)\n\ny_pred = model.predict(test_dset, verbose=1)\nsub['target'] = y_pred\nsub.to_csv(f'vit_sub_{FOLD}.csv', index=False)","metadata":{},"execution_count":null,"outputs":[]}]}