{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":164364362,"sourceType":"kernelVersion"}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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 os\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport time\nimport math","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-29T13:15:56.930656Z","iopub.execute_input":"2024-02-29T13:15:56.931050Z","iopub.status.idle":"2024-02-29T13:15:57.884523Z","shell.execute_reply.started":"2024-02-29T13:15:56.931000Z","shell.execute_reply":"2024-02-29T13:15:57.883627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom keras import backend as K\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import Model","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:15:57.886109Z","iopub.execute_input":"2024-02-29T13:15:57.886496Z","iopub.status.idle":"2024-02-29T13:16:11.237438Z","shell.execute_reply.started":"2024-02-29T13:15:57.886471Z","shell.execute_reply":"2024-02-29T13:16:11.236471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#output of other notebook\ntrain_df_all = pd.read_csv(\"/kaggle/input/hms-eda-confidence-of-diagnosis/hms_train_eda.csv\")\n\ntest_df  = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\")\nsample_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.238773Z","iopub.execute_input":"2024-02-29T13:16:11.239443Z","iopub.status.idle":"2024-02-29T13:16:11.562653Z","shell.execute_reply.started":"2024-02-29T13:16:11.239414Z","shell.execute_reply":"2024-02-29T13:16:11.561700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_spectro_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\"\ntrain_spectro_files = os.listdir(train_spectro_dir)\nlen(train_spectro_files)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.565233Z","iopub.execute_input":"2024-02-29T13:16:11.566059Z","iopub.status.idle":"2024-02-29T13:16:11.721794Z","shell.execute_reply.started":"2024-02-29T13:16:11.566022Z","shell.execute_reply":"2024-02-29T13:16:11.720707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.sqrt(10000)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.723029Z","iopub.execute_input":"2024-02-29T13:16:11.723393Z","iopub.status.idle":"2024-02-29T13:16:11.729993Z","shell.execute_reply.started":"2024-02-29T13:16:11.723367Z","shell.execute_reply":"2024-02-29T13:16:11.729053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Split data for train and validation\n\nData should not split randomly, because there are multiple columns for a patient. First, unique patients IDs are identified. Then, patient IDs are splitted for train and validation.","metadata":{}},{"cell_type":"code","source":"patient_all = train_df_all[\"patient_id\"].unique()\nnp.random.seed(1)\nnp.random.shuffle(patient_all)\n\nn_patient = patient_all.shape[0]\nn_patient_train = int(n_patient*0.85)\nn_patient_val = n_patient - n_patient_train\nn_patient, n_patient_train, n_patient_val\n\npatient_train = set(patient_all[:n_patient_train])\npatient_val = set(patient_all[n_patient_train:])\n\n\nfilter1 = train_df_all[\"patient_id\"].apply(lambda x: x in patient_train)\ntrain_df = train_df_all.loc[filter1]\nval_df = train_df_all.loc[~filter1]","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.731397Z","iopub.execute_input":"2024-02-29T13:16:11.731764Z","iopub.status.idle":"2024-02-29T13:16:11.807105Z","shell.execute_reply.started":"2024-02-29T13:16:11.731732Z","shell.execute_reply":"2024-02-29T13:16:11.806342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.shape, val_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.808191Z","iopub.execute_input":"2024-02-29T13:16:11.808473Z","iopub.status.idle":"2024-02-29T13:16:11.814485Z","shell.execute_reply.started":"2024-02-29T13:16:11.808449Z","shell.execute_reply":"2024-02-29T13:16:11.813596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Cleaning\n\nData which satisfy following conditions are selected for developing ML model.\n\n- total vote is greater than or equal to 3\n- This vote is 100% confident. (all votes has the same diagnosis)\n- Spectrogram has single Symptom.\n\nThose preprocess are described in a notebook [HMS-EDA_Confidence_Of_Diagnosis](https://www.kaggle.com/code/hidetaketakahashi/hms-eda-confidence-of-diagnosis).","metadata":{}},{"cell_type":"code","source":"filter1 = (train_df[\"total_vote\"] >= 3) & (train_df[\"max_vote_rate\"] >= 1) #& (train_df[\"SpectrogramWithSingleSymtom\"])\n\ntrain_df_clean = train_df.loc[filter1] \ntrain_df_mess = train_df.loc[~filter1] \n\nfilter1 = (val_df[\"total_vote\"] >= 3) & (val_df[\"max_vote_rate\"] >= 1)# & (val_df[\"SpectrogramWithSingleSymtom\"])\n\nval_df_clean = val_df.loc[filter1] \nval_df_mess = val_df.loc[~filter1] ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.815684Z","iopub.execute_input":"2024-02-29T13:16:11.815962Z","iopub.status.idle":"2024-02-29T13:16:11.834546Z","shell.execute_reply.started":"2024-02-29T13:16:11.815940Z","shell.execute_reply":"2024-02-29T13:16:11.833727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_clean.shape, val_df_clean.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.835737Z","iopub.execute_input":"2024-02-29T13:16:11.836001Z","iopub.status.idle":"2024-02-29T13:16:11.841836Z","shell.execute_reply.started":"2024-02-29T13:16:11.835977Z","shell.execute_reply":"2024-02-29T13:16:11.840816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"select = [\"spectrogram_id\", \"expert_consensus\"]\ntrain_df_clean_unique = train_df_clean[select].value_counts().reset_index()      \nval_df_clean_unique = val_df_clean[select].value_counts().reset_index()      ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.845493Z","iopub.execute_input":"2024-02-29T13:16:11.845813Z","iopub.status.idle":"2024-02-29T13:16:11.871946Z","shell.execute_reply.started":"2024-02-29T13:16:11.845789Z","shell.execute_reply":"2024-02-29T13:16:11.871004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_clean_unique.shape, val_df_clean_unique.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.873152Z","iopub.execute_input":"2024-02-29T13:16:11.873486Z","iopub.status.idle":"2024-02-29T13:16:11.879510Z","shell.execute_reply.started":"2024-02-29T13:16:11.873457Z","shell.execute_reply":"2024-02-29T13:16:11.878639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_clean_unique[\"expert_consensus\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.880734Z","iopub.execute_input":"2024-02-29T13:16:11.881509Z","iopub.status.idle":"2024-02-29T13:16:11.892367Z","shell.execute_reply.started":"2024-02-29T13:16:11.881472Z","shell.execute_reply":"2024-02-29T13:16:11.891505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_df","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.909621Z","iopub.execute_input":"2024-02-29T13:16:11.910853Z","iopub.status.idle":"2024-02-29T13:16:11.923329Z","shell.execute_reply.started":"2024-02-29T13:16:11.910820Z","shell.execute_reply":"2024-02-29T13:16:11.922386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_index_dict = {\"Seizure\":0, \"LPD\":1, \"GPD\":2, \"LRDA\":3, \"GRDA\":4, \"Other\":5}\nlabel_index_dict","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.924720Z","iopub.execute_input":"2024-02-29T13:16:11.924994Z","iopub.status.idle":"2024-02-29T13:16:11.933534Z","shell.execute_reply.started":"2024-02-29T13:16:11.924970Z","shell.execute_reply":"2024-02-29T13:16:11.932548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def readSpectrogram(fileID, test = False, dropTime = True):\n    fid = str(fileID)\n    \n    if test:\n        spectro_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/\"\n        data = pd.read_parquet(spectro_dir + fid + \".parquet\")\n        \n    else:\n        spectro_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\"\n        data = pd.read_parquet(spectro_dir + fid + \".parquet\")\n    \n    if dropTime:\n        return data.drop(\"time\", axis = 1).to_numpy()\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.934672Z","iopub.execute_input":"2024-02-29T13:16:11.935006Z","iopub.status.idle":"2024-02-29T13:16:11.941543Z","shell.execute_reply.started":"2024-02-29T13:16:11.934980Z","shell.execute_reply":"2024-02-29T13:16:11.940612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpectrogramLoader:\n    \n    def __init__(self):\n        self.shuffle = True\n        self.scaling = True\n        self.L = 48\n        self.H = 96\n        self.impute_value = 0\n        self.selectChannel = True\n        \n    def imputation(self, data):\n        filter1 = np.isnan(data)\n        \n        if np.mean(filter1) > 0.1:\n            return False\n        data[filter1] = self.impute_value\n        return True\n        \n    def SecToIndex(self, sec):\n        \n        return max(int(int(sec)/2), 0)\n    \n    def cropData300(self, data, sec):\n        \n        idxOffset = self.SecToIndex(sec)\n        data = data[idxOffset:]\n        \n        L1 = min(data.shape[0], 300)\n        \n        if L1 < 300 - 10:\n            return data, False\n        \n        return data, True\n        \n    \n    def cropDataShort(self, data):\n        \n        L1 = data.shape[0]\n        start = 0\n        if L1 > self.L:\n            start = np.random.randint(L1 - self.L)\n\n        end = start + self.L\n        data =  data[start:end]\n        \n        return data\n    \n    def reshapeTo4Channels(self, data):\n        \n        mat = np.zeros((data.shape[0], self.H, 4))\n        \n        starts = [0,100,200,300] \n        for i in range(4):\n            mat[:,:,i] = data[:,starts[i]:(starts[i] + self.H)]\n\n        return mat\n        \n    def shuffleChannel(self, mat):\n        \n        idx_array = np.arange(4)\n        np.random.shuffle(idx_array)\n        \n        return mat[:,:,idx_array]\n    \n    def selectChannelRandom(self, mat):\n        \n        channel = np.random.randint(4)\n        \n        return mat[:,:,channel].reshape((self.L, self.H, 1))\n        \n            \n    def loadSpectrogram(self, spectID, offsetSec):\n        \n        data = readSpectrogram(spectID)\n        \n        data, crop_result = self.cropData300(data, offsetSec)\n        if crop_result is False:\n            return None\n\n        if self.scaling:\n            #data = data/np.max(data)\n            #data = data/np.quantile(data, 0.995)\n            data = data/1000\n            filter1 = data > 1\n            data[filter1] = 1\n                \n        data = self.cropDataShort(data)\n                \n        data = self.reshapeTo4Channels(data)\n\n        impute_result = self.imputation(data)\n        \n        if impute_result is False:\n            return None\n        \n        \n        if self.selectChannel:\n            return self.selectChannelRandom(data)\n            \n        return data\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.942975Z","iopub.execute_input":"2024-02-29T13:16:11.943329Z","iopub.status.idle":"2024-02-29T13:16:11.960458Z","shell.execute_reply.started":"2024-02-29T13:16:11.943305Z","shell.execute_reply":"2024-02-29T13:16:11.959616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df[[\"spectrogram_id\", \"spectrogram_label_offset_seconds\"]].sample()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.961265Z","iopub.execute_input":"2024-02-29T13:16:11.961503Z","iopub.status.idle":"2024-02-29T13:16:11.979305Z","shell.execute_reply.started":"2024-02-29T13:16:11.961482Z","shell.execute_reply":"2024-02-29T13:16:11.978522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SL = SpectrogramLoader()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.980251Z","iopub.execute_input":"2024-02-29T13:16:11.980496Z","iopub.status.idle":"2024-02-29T13:16:11.984577Z","shell.execute_reply.started":"2024-02-29T13:16:11.980475Z","shell.execute_reply":"2024-02-29T13:16:11.983639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(10):\n    spect = SL.loadSpectrogram(927827189, 0)\n    print(np.max(spect))\n","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:11.986018Z","iopub.execute_input":"2024-02-29T13:16:11.986573Z","iopub.status.idle":"2024-02-29T13:16:12.490818Z","shell.execute_reply.started":"2024-02-29T13:16:11.986541Z","shell.execute_reply":"2024-02-29T13:16:12.489747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataBatchCreator:\n    \n    def __init__(self):\n        \n        self.label_index_dict = {\"Seizure\":0, \"LPD\":1, \"GPD\":2, \"LRDA\":3, \"GRDA\":4, \"Other\":5}\n        self.keys = list(self.label_index_dict.keys())\n        self.n_keys = len(self.keys)\n        self.SL = SpectrogramLoader()\n        self.train_df = {}\n        self.val_df = {}\n            \n    def readDF(self, df, train = True):\n        \n        for key in self.keys:\n            filter1 = df[\"expert_consensus\"] == key\n            data = df[[\"spectrogram_id\", \"spectrogram_label_offset_seconds\"]].loc[filter1].to_numpy()\n            \n            if train:\n                self.train_df[key] = data\n            else:\n                self.val_df[key] = data\n    \n    def loadBatch(self, df, n_sample = 2):\n        \n        y_array = np.zeros((self.n_keys*n_sample, self.n_keys), dtype = int)\n        X_mat = np.zeros((self.n_keys*n_sample, self.SL.L, self.SL.H, 1))\n        \n        idx = 0\n        \n        for key in self.keys:\n            n_df = df[key].shape[0]\n            for i in range(n_sample):\n                \n                while True:\n                \n                    sample_idx = np.random.randint(n_df) \n                    spectID = int(df[key][sample_idx, 0])\n                    offsetSec = df[key][sample_idx, 1]\n                    mat = self.SL.loadSpectrogram(spectID, offsetSec)\n                    if mat is not None:\n                        X_mat[idx] =  mat\n                        y_array[idx, self.label_index_dict[key]] = 1\n                        idx += 1\n                        break\n            \n        return X_mat, y_array        ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:12.492135Z","iopub.execute_input":"2024-02-29T13:16:12.492546Z","iopub.status.idle":"2024-02-29T13:16:12.504719Z","shell.execute_reply.started":"2024-02-29T13:16:12.492486Z","shell.execute_reply":"2024-02-29T13:16:12.503826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DBC = DataBatchCreator()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:12.506295Z","iopub.execute_input":"2024-02-29T13:16:12.506642Z","iopub.status.idle":"2024-02-29T13:16:12.517393Z","shell.execute_reply.started":"2024-02-29T13:16:12.506611Z","shell.execute_reply":"2024-02-29T13:16:12.516559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DBC.readDF(train_df_clean)\nDBC.readDF(val_df_clean, False)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:12.518645Z","iopub.execute_input":"2024-02-29T13:16:12.519130Z","iopub.status.idle":"2024-02-29T13:16:12.586244Z","shell.execute_reply.started":"2024-02-29T13:16:12.519099Z","shell.execute_reply":"2024-02-29T13:16:12.585359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X1, y1 = DBC.loadBatch(DBC.train_df, 6)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:12.587477Z","iopub.execute_input":"2024-02-29T13:16:12.587843Z","iopub.status.idle":"2024-02-29T13:16:14.483538Z","shell.execute_reply.started":"2024-02-29T13:16:12.587816Z","shell.execute_reply":"2024-02-29T13:16:14.482712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X1.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:14.484866Z","iopub.execute_input":"2024-02-29T13:16:14.485254Z","iopub.status.idle":"2024-02-29T13:16:14.491420Z","shell.execute_reply.started":"2024-02-29T13:16:14.485223Z","shell.execute_reply":"2024-02-29T13:16:14.490470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ConvBlock(X, channel, ksizes, psizes, pstrides, max_pool = True):\n    \n    X = layers.Conv2D(channel, kernel_size = ksizes, strides = 1, padding = \"same\")(X)        \n    X = layers.BatchNormalization()(X)\n    X = layers.ReLU()(X)\n    if max_pool:\n        X =  layers.MaxPooling2D(pool_size = psizes, strides = pstrides)(X)    \n    \n    return X","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:14.711345Z","iopub.execute_input":"2024-02-29T13:16:14.711733Z","iopub.status.idle":"2024-02-29T13:16:14.717515Z","shell.execute_reply.started":"2024-02-29T13:16:14.711707Z","shell.execute_reply":"2024-02-29T13:16:14.716528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ConvLine(X, n_output = 6):\n    \n#    X = layers.Rescaling(scale = 1./100., offset= 0, )(X)\n\n    X = ConvBlock(X, 64, (3,3), (2,2), (2,2)) # 48 x 96\n    X = ConvBlock(X, 64*2, (3,3), (2,2), (2,2)) #24 x 48\n    X = ConvBlock(X, 64*4, (3,3), (2,2), (2,2)) #12 x 24\n    X = ConvBlock(X, 64*8, (3,3), (2,2), (2,2)) #6 x 12\n    X = ConvBlock(X, 64*8, (3,3), (2,2), (2,2)) #3 x 6\n    \n    return X","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:14.718652Z","iopub.execute_input":"2024-02-29T13:16:14.718917Z","iopub.status.idle":"2024-02-29T13:16:14.730892Z","shell.execute_reply.started":"2024-02-29T13:16:14.718895Z","shell.execute_reply":"2024-02-29T13:16:14.729957Z"},"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":"def createModel(n_output = 6):\n\n    Input =  layers.Input(shape=(48, 96,1))\n    X = ConvLine(Input, n_output)\n    X = layers.Flatten()(X)\n    X = layers.Dropout(0.3)(X)\n    X = layers.Dense(n_output)(X)\n    \n\n    model = Model(inputs = Input, outputs = X)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:14.731989Z","iopub.execute_input":"2024-02-29T13:16:14.732250Z","iopub.status.idle":"2024-02-29T13:16:14.740872Z","shell.execute_reply.started":"2024-02-29T13:16:14.732229Z","shell.execute_reply":"2024-02-29T13:16:14.740074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_Spect = createModel()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:14.741775Z","iopub.execute_input":"2024-02-29T13:16:14.742035Z","iopub.status.idle":"2024-02-29T13:16:15.683938Z","shell.execute_reply.started":"2024-02-29T13:16:14.742013Z","shell.execute_reply":"2024-02-29T13:16:15.682944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_Spect.summary()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.685809Z","iopub.execute_input":"2024-02-29T13:16:15.686241Z","iopub.status.idle":"2024-02-29T13:16:15.752086Z","shell.execute_reply.started":"2024-02-29T13:16:15.686207Z","shell.execute_reply":"2024-02-29T13:16:15.751054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#tf.keras.utils.plot_model(model_Spect)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.753253Z","iopub.execute_input":"2024-02-29T13:16:15.753549Z","iopub.status.idle":"2024-02-29T13:16:15.758058Z","shell.execute_reply.started":"2024-02-29T13:16:15.753524Z","shell.execute_reply":"2024-02-29T13:16:15.757096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BCE = tf.keras.losses.BinaryCrossentropy(from_logits=True)\n\ndef custom_loss(y_true, y_pred):\n\n    filter1 = y_true == 1\n\n    p_loss = BCE(y_true[filter1], y_pred[filter1])\n\n    filter1 = y_true == 0\n    n_loss = BCE(y_true[filter1], y_pred[filter1])\n\n    loss =  p_loss + n_loss*2\n\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.759372Z","iopub.execute_input":"2024-02-29T13:16:15.759709Z","iopub.status.idle":"2024-02-29T13:16:15.767997Z","shell.execute_reply.started":"2024-02-29T13:16:15.759675Z","shell.execute_reply":"2024-02-29T13:16:15.766970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def val_score(X, Y, model):\n    \n    Y_pred = model.predict(X, verbose = 0)\n    loss =  custom_loss(Y, Y_pred)\n    \n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.769169Z","iopub.execute_input":"2024-02-29T13:16:15.769502Z","iopub.status.idle":"2024-02-29T13:16:15.778467Z","shell.execute_reply.started":"2024-02-29T13:16:15.769478Z","shell.execute_reply":"2024-02-29T13:16:15.777520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@tf.function\ndef train_step(X, Y, model):\n        \n    with tf.GradientTape() as tape:\n        Y_pred = model(X)\n        loss = custom_loss(Y, Y_pred)\n        \n        \n    grads = tape.gradient(loss, model.trainable_weights)\n    optimizer.apply_gradients(zip(grads, model.trainable_weights))\n    \n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.779826Z","iopub.execute_input":"2024-02-29T13:16:15.780121Z","iopub.status.idle":"2024-02-29T13:16:15.789785Z","shell.execute_reply.started":"2024-02-29T13:16:15.780098Z","shell.execute_reply":"2024-02-29T13:16:15.788533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_probability(X, model):\n    pred = model.predict(X, verbose = 0)\n    odd = np.exp(pred)\n    prob = odd/(1+odd)\n    \n    return prob","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.791114Z","iopub.execute_input":"2024-02-29T13:16:15.791417Z","iopub.status.idle":"2024-02-29T13:16:15.800817Z","shell.execute_reply.started":"2024-02-29T13:16:15.791381Z","shell.execute_reply":"2024-02-29T13:16:15.799649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def trainLoop(model, history, n_epoch, n_train_repeat = 50, n_batch_val = 300, n_batch_train = 12):\n    \n    \n    for k in range(n_epoch):\n        time1 = time.time()\n        loss_list = []\n        for i in range(n_train_repeat):\n            \n            X_train, y_train = DBC.loadBatch(DBC.train_df, n_batch_train)\n            \n            loss = train_step(X_train, y_train, model)\n            loss_list.append(loss)\n            \n        train_loss = np.mean(loss_list)   \n        \n        X_val, y_val = DBC.loadBatch(DBC.val_df, n_batch_val)\n        \n        val_loss = val_score(X_val, y_val, model)\n        \n        time2 = time.time()\n        time3 = np.round(time2- time1)\n        \n        print(k, \"train loss\", train_loss, \", val loss \", val_loss, \" time[s] = \", time3)\n        \n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        \n        if val_loss < history[\"best_val_loss\"]:\n            history[\"best_val_loss\"] = val_loss\n            model.save_weights(\"HMS_CNN1/ckpt1\")\n            print(\"write model at epoch \", k)\n            \n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.802504Z","iopub.execute_input":"2024-02-29T13:16:15.802924Z","iopub.status.idle":"2024-02-29T13:16:15.819803Z","shell.execute_reply.started":"2024-02-29T13:16:15.802897Z","shell.execute_reply":"2024-02-29T13:16:15.818704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer =  tf.keras.optimizers.Adam(learning_rate=0.0005)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.820839Z","iopub.execute_input":"2024-02-29T13:16:15.821093Z","iopub.status.idle":"2024-02-29T13:16:15.835553Z","shell.execute_reply.started":"2024-02-29T13:16:15.821071Z","shell.execute_reply":"2024-02-29T13:16:15.834469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_history(history):\n    \n    fig, ax = plt.subplots(figsize = (5,4))\n    ax.plot(history[\"train_loss\"], label = \"train loss\")    \n    ax.plot(history[\"val_loss\"], label = \"val loss\")\n    ax.set_title(\"loss\")\n    ax.legend()\n    ax.grid()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.837020Z","iopub.execute_input":"2024-02-29T13:16:15.837882Z","iopub.status.idle":"2024-02-29T13:16:15.842864Z","shell.execute_reply.started":"2024-02-29T13:16:15.837852Z","shell.execute_reply":"2024-02-29T13:16:15.841917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_CNN = {}\nhistory_CNN[\"train_loss\"] = []\nhistory_CNN[\"val_loss\"] = []\nhistory_CNN[\"best_val_loss\"] = 1000000.","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.843974Z","iopub.execute_input":"2024-02-29T13:16:15.844252Z","iopub.status.idle":"2024-02-29T13:16:15.854063Z","shell.execute_reply.started":"2024-02-29T13:16:15.844228Z","shell.execute_reply":"2024-02-29T13:16:15.853044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainLoop(model_Spect, history_CNN, n_epoch = 30, n_train_repeat = 50, n_batch_val = 300, n_batch_train = 12)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:16:15.855400Z","iopub.execute_input":"2024-02-29T13:16:15.855755Z","iopub.status.idle":"2024-02-29T13:17:02.439533Z","shell.execute_reply.started":"2024-02-29T13:16:15.855727Z","shell.execute_reply":"2024-02-29T13:17:02.437980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history_CNN)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T13:17:02.440695Z","iopub.status.idle":"2024-02-29T13:17:02.441145Z","shell.execute_reply.started":"2024-02-29T13:17:02.440925Z","shell.execute_reply":"2024-02-29T13:17:02.440944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}