{"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":"# Importing the required libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nimport tensorflow as tf\nimport os\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pydicom as pyd\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nimport albumentations as A\nimport cv2\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.utils import Sequence\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.callbacks import *","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:34.280851Z","iopub.execute_input":"2021-09-12T06:06:34.281107Z","iopub.status.idle":"2021-09-12T06:06:40.504969Z","shell.execute_reply.started":"2021-09-12T06:06:34.281081Z","shell.execute_reply":"2021-09-12T06:06:40.504088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reading the input files","metadata":{}},{"cell_type":"code","source":"images_path = '../input/rsna-pneumonia-detection-challenge/stage_2_train_images'\ntrain_labels_df = pd.read_csv('../input/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv')\nlabel_meta_data = pd.read_csv('../input/rsna-pneumonia-detection-challenge/stage_2_detailed_class_info.csv')","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.507623Z","iopub.execute_input":"2021-09-12T06:06:40.50784Z","iopub.status.idle":"2021-09-12T06:06:40.606956Z","shell.execute_reply.started":"2021-09-12T06:06:40.507815Z","shell.execute_reply":"2021-09-12T06:06:40.606059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> The dataset is available as part of the RSNA Pneumonia Detection Competition in Kaggle itself. The data has been added to this workspace and we are storing the path of the images as well as converting the CSV's into a Pandas Dataframe","metadata":{}},{"cell_type":"code","source":"train_labels_df.head(10)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.608433Z","iopub.execute_input":"2021-09-12T06:06:40.608697Z","iopub.status.idle":"2021-09-12T06:06:40.636799Z","shell.execute_reply.started":"2021-09-12T06:06:40.608661Z","shell.execute_reply":"2021-09-12T06:06:40.635924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> The stage_2_train_labels CSV contains the following information \n* patientId : For uniquely identifying the patient whose X-ray scan is in the image dataset\n* Coordinates of the bounding boxes : The x_min, y_min, width and height of the rectangular boudning boxes that detect inflammation leading to the diagnosis of pneumonia\n* Taeget Class : the categorical output which tells if the patient has pneumonia or not.\n\nNote that if the target is 0, then then the coordinate columns carry NAN values","metadata":{}},{"cell_type":"code","source":"label_meta_data.head(10)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.638951Z","iopub.execute_input":"2021-09-12T06:06:40.639279Z","iopub.status.idle":"2021-09-12T06:06:40.648806Z","shell.execute_reply.started":"2021-09-12T06:06:40.639246Z","shell.execute_reply":"2021-09-12T06:06:40.647938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> The stage_2_detailed_class_info CSV contains the following information\n* patientId : For uniquely identifying the patient whose X-ray scan is in the image dataset\n* class : detailed class info wherein there are three categories in picture\n\n1.     Normal : No pneumonia (Target = 0)\n2.     Lung Opacity : Pneumonia (Target = 1)\n3.     No Lung Opacity/ Not Normal : Inflammation not leading to Pneumonia ( Target = 0)\n","metadata":{}},{"cell_type":"code","source":"print('Size of Dataset 1: ',train_labels_df.shape)\nprint('Size of Dataset 2: ',label_meta_data.shape)\nprint('Number of Unique X-Rays in Dataset 1 : ',train_labels_df['patientId'].nunique())\nprint('Number of Unique X-Rays in Dataset 2 : ',label_meta_data['patientId'].nunique())\n","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.650505Z","iopub.execute_input":"2021-09-12T06:06:40.65087Z","iopub.status.idle":"2021-09-12T06:06:40.684808Z","shell.execute_reply.started":"2021-09-12T06:06:40.650833Z","shell.execute_reply":"2021-09-12T06:06:40.684012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Both datasets have same number of rows - 30,227 but out of which there are only 26,684 unique patient IDs -> which means some patient scans may have more than one bounding box.","metadata":{}},{"cell_type":"code","source":"train_labels_df.drop_duplicates(inplace=True)\nlabel_meta_data.drop_duplicates(inplace=True)\nprint('Size of Dataset 1: ',train_labels_df.shape)\nprint('Size of Dataset 2: ',label_meta_data.shape)\nprint('Number of Unique X-Rays in Dataset 1 : ',train_labels_df['patientId'].nunique())\nprint('Number of Unique X-Rays in Dataset 2 : ',label_meta_data['patientId'].nunique())","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.686174Z","iopub.execute_input":"2021-09-12T06:06:40.686442Z","iopub.status.idle":"2021-09-12T06:06:40.741735Z","shell.execute_reply.started":"2021-09-12T06:06:40.686411Z","shell.execute_reply":"2021-09-12T06:06:40.740882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> This shows that there are no duplicates in the bounding box dataset","metadata":{}},{"cell_type":"code","source":"patient_info_df = pd.DataFrame(columns=['age','sex','patientId'])\nfor ix,id_ in tqdm(enumerate(label_meta_data['patientId'])):\n    age=pyd.read_file(os.path.join(images_path,\n                                   id_+'.dcm')).PatientAge\n    sex=pyd.read_file(os.path.join(images_path,\n                               id_+'.dcm')).PatientSex\n\n    patient_info_df.loc[ix,'age']=age\n    patient_info_df.loc[ix,'sex']=sex\n    patient_info_df.loc[ix,'patientId'] = id_","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:06:40.743169Z","iopub.execute_input":"2021-09-12T06:06:40.743438Z","iopub.status.idle":"2021-09-12T06:11:42.013084Z","shell.execute_reply.started":"2021-09-12T06:06:40.743405Z","shell.execute_reply":"2021-09-12T06:11:42.012407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Extracting the age and gender of the patient for all the IDs specified in the detailed info CSV into a separate data frame","metadata":{}},{"cell_type":"code","source":"# Merging this patient information with the exiting meta data\npatient_info_df = pd.merge(label_meta_data,patient_info_df,on=\"patientId\")\npatient_info_df","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.014368Z","iopub.execute_input":"2021-09-12T06:11:42.014754Z","iopub.status.idle":"2021-09-12T06:11:42.049668Z","shell.execute_reply.started":"2021-09-12T06:11:42.014719Z","shell.execute_reply":"2021-09-12T06:11:42.048863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"merged_data_info = pd.merge(train_labels_df,patient_info_df,on=\"patientId\")\nmerged_data_info.tail(10)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.051064Z","iopub.execute_input":"2021-09-12T06:11:42.051342Z","iopub.status.idle":"2021-09-12T06:11:42.088422Z","shell.execute_reply.started":"2021-09-12T06:11:42.051309Z","shell.execute_reply":"2021-09-12T06:11:42.087612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Merging both the files into a single dataset to make sure that there are no inconsistent patient IDs across both the datasets","metadata":{}},{"cell_type":"code","source":"print(merged_data_info.shape)\nprint(merged_data_info.isna().sum())","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.091543Z","iopub.execute_input":"2021-09-12T06:11:42.091748Z","iopub.status.idle":"2021-09-12T06:11:42.110761Z","shell.execute_reply.started":"2021-09-12T06:11:42.091724Z","shell.execute_reply":"2021-09-12T06:11:42.109951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> The shape of the merged one is the same as the original bounding box dataset and there are no other null values in the other output columns","metadata":{}},{"cell_type":"code","source":"print('Minimum Age in the dataset:', merged_data_info['age'].min())\nprint('Maximum Age in the dataset', merged_data_info['age'].max())\n\nmerged_data_info.columns = merged_data_info.columns.str.strip()\nmerged_data_info['sex']=merged_data_info['sex'].replace({ 'M' : 0, 'F' : 1  })\nmerged_data_info['label']=merged_data_info['class']\nmerged_data_info['label']=merged_data_info['label'].replace({ 'Normal' : 0, 'Lung Opacity' : 1, 'No Lung Opacity / Not Normal' : 2 })\n","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.112363Z","iopub.execute_input":"2021-09-12T06:11:42.11261Z","iopub.status.idle":"2021-09-12T06:11:42.163809Z","shell.execute_reply.started":"2021-09-12T06:11:42.112578Z","shell.execute_reply":"2021-09-12T06:11:42.163125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> The Age can be considered as an continuous input. After removing the white spaces at the beginning and end of string columns, we encode the categorical columns into integers with the following code. The label will be the new classification output column\n* Male : 0\n* Female : 1\n\n* Normal : 0\n* Lung Opacity : 1\n* No Lung Opacity / Not Normal : 2","metadata":{}},{"cell_type":"code","source":"merged_data_info['age'] = merged_data_info['age'].astype('int64')","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.164903Z","iopub.execute_input":"2021-09-12T06:11:42.165635Z","iopub.status.idle":"2021-09-12T06:11:42.174487Z","shell.execute_reply.started":"2021-09-12T06:11:42.1656Z","shell.execute_reply":"2021-09-12T06:11:42.173792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"merged_data_info.info()","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.175622Z","iopub.execute_input":"2021-09-12T06:11:42.176009Z","iopub.status.idle":"2021-09-12T06:11:42.201Z","shell.execute_reply.started":"2021-09-12T06:11:42.175973Z","shell.execute_reply":"2021-09-12T06:11:42.200127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Other than the class description column and the patient ID colummn all the others are of numerical datatype","metadata":{}},{"cell_type":"markdown","source":"> ","metadata":{}},{"cell_type":"markdown","source":"# Target Distribution","metadata":{}},{"cell_type":"code","source":"label_count=label_meta_data['class'].value_counts()\nexplode = (0.01,0.01,0.01)  \n\nfig1, ax1 = plt.subplots(figsize=(5,5))\nax1.pie(label_count.values, explode=explode, labels=label_count.index, autopct='%1.1f%%',\n        shadow=True, startangle=90)\nax1.axis('equal') \nplt.title('Class Distribution')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.202432Z","iopub.execute_input":"2021-09-12T06:11:42.202814Z","iopub.status.idle":"2021-09-12T06:11:42.337628Z","shell.execute_reply.started":"2021-09-12T06:11:42.202774Z","shell.execute_reply":"2021-09-12T06:11:42.336902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Only 22% of the patients have Lung opacity, which shows that the dataset is very much imbalanced. This is based on the unique patientIDs. Out of these patient X-ray scans each image may have one or more bounding boxes.","metadata":{}},{"cell_type":"code","source":"# lets take a look at our Target Distribution\nlabel_count=merged_data_info['Target'].value_counts()\nexplode = (0.1,0.0)  \n\nfig1, ax1 = plt.subplots(figsize=(5,5))\nax1.pie(label_count.values, explode=explode, labels=['Normal','Pneumonia'], autopct='%1.1f%%',\n        shadow=True, startangle=90)\nax1.axis('equal') \nplt.title('Target Distribution')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.338917Z","iopub.execute_input":"2021-09-12T06:11:42.339172Z","iopub.status.idle":"2021-09-12T06:11:42.432816Z","shell.execute_reply.started":"2021-09-12T06:11:42.339138Z","shell.execute_reply":"2021-09-12T06:11:42.432032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> We need to atleast double the positive datasets artificially to make it a balanced data set","metadata":{}},{"cell_type":"markdown","source":"# Visualization of the Images","metadata":{}},{"cell_type":"code","source":"r=c=3\nfig= plt.figure(figsize=(15,14))\nfor i in range(1,r*c+1):\n    id_= np.random.choice(merged_data_info['patientId'].values)\n    label_0= np.unique(merged_data_info['Target'][merged_data_info['patientId']==id_])\n    label_1= np.unique(merged_data_info['class'][merged_data_info['patientId']==id_])\n    \n    #read xray\n    img=pyd.read_file(os.path.join(images_path,id_+'.dcm')).pixel_array\n    fig.add_subplot(r,c,i)\n    plt.imshow(img,cmap='gray')\n    if label_0==1:\n        plt.title('Pneumonia Infected'+' | '+label_1)\n    else:\n        plt.title('Normal Xray'+' | '+label_1)\n    plt.xticks([])\n    plt.yticks([])","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:42.434162Z","iopub.execute_input":"2021-09-12T06:11:42.434438Z","iopub.status.idle":"2021-09-12T06:11:44.297459Z","shell.execute_reply.started":"2021-09-12T06:11:42.434405Z","shell.execute_reply":"2021-09-12T06:11:44.296765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> These random set of 9 images show various combinations of Target class and the detailed category","metadata":{}},{"cell_type":"markdown","source":"# Visualization of the areas of inflammation","metadata":{}},{"cell_type":"code","source":"id_= np.random.choice(merged_data_info[merged_data_info['Target'] == 1]['patientId'].values)\nclass_=merged_data_info['class'][merged_data_info['patientId']==id_]\n\nplt.figure(figsize=(15,10))\ncurrent_axis = plt.gca()\nimg=pyd.read_file(os.path.join(images_path,id_+'.dcm')).pixel_array\nplt.imshow(img,cmap='bone')\n\n\ncurrent_axis = plt.gca()\nboxes=train_labels_df[['x','y','width','height']][merged_data_info['patientId']==id_].values\n\nfor box in boxes:\n    x=box[0]\n    y=box[1]\n    w=box[2]\n    h=box[3]\n    current_axis.add_patch(plt.Rectangle((x, y), w, h, \n                                         color='red', fill=False, linewidth=3))  \n    ","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:44.298375Z","iopub.execute_input":"2021-09-12T06:11:44.298593Z","iopub.status.idle":"2021-09-12T06:11:44.808917Z","shell.execute_reply.started":"2021-09-12T06:11:44.298566Z","shell.execute_reply":"2021-09-12T06:11:44.808112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(x=merged_data_info['Target'], hue=merged_data_info['sex'])","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:44.810082Z","iopub.execute_input":"2021-09-12T06:11:44.810437Z","iopub.status.idle":"2021-09-12T06:11:45.013365Z","shell.execute_reply.started":"2021-09-12T06:11:44.810398Z","shell.execute_reply":"2021-09-12T06:11:45.012714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Its seems the representation of males is slightly greater than females by almost ~2000 datapoints for both positive as well as negative cases","metadata":{}},{"cell_type":"code","source":"X = merged_data_info[['patientId','age','sex']].values\ny = tf.keras.utils.to_categorical(merged_data_info['label'].values,num_classes=3)\nX_train, X_val, y_train, y_val = train_test_split(X,y,test_size=0.2)\nprint('Samples in Training Data: ', len(X_train))\nprint('Samples in Validation Data: ', len(X_val))","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:01.993919Z","iopub.execute_input":"2021-09-12T06:35:01.994444Z","iopub.status.idle":"2021-09-12T06:35:02.011905Z","shell.execute_reply.started":"2021-09-12T06:35:01.994405Z","shell.execute_reply":"2021-09-12T06:35:02.011241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 32\nHEIGHT = 224\nWIDTH = 224\ndims=(HEIGHT,WIDTH,1)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:45.026989Z","iopub.execute_input":"2021-09-12T06:11:45.027434Z","iopub.status.idle":"2021-09-12T06:11:45.032409Z","shell.execute_reply.started":"2021-09-12T06:11:45.027399Z","shell.execute_reply":"2021-09-12T06:11:45.030995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class KustomGenerator(Sequence):\n    def __init__(self,input_data,batch_size=BATCH_SIZE,dims=(HEIGHT,WIDTH,1),is_train=True):\n        self.input_ids=input_data[0][:,0]\n        self.input_para=input_data[0][:,1:3]\n        self.input_targets=input_data[1]\n        self.batch_size=batch_size\n        self.dims=dims\n        self.is_train=is_train\n        self.on_epoch_end()\n    \n    def on_epoch_end(self):\n        self.indexes=np.arange(len(self.input_ids))\n        if self.is_train:\n            np.random.shuffle(self.indexes)\n    \n    \n    def __len__(self):\n        return int(len(self.input_ids)/self.batch_size)\n    \n    def __getitem__(self,index):\n        \n        indexes=self.indexes[index*self.batch_size:(index+1)*self.batch_size]\n        X_ids = [self.input_ids[i] for i in indexes]\n        Y_=[self.input_targets[i] for i in indexes]\n        X_para = [self.input_para[i] for i in indexes]\n        X=self.__data_generation(X_ids)\n        return [X,X_para],np.array(Y_)\n\n    def __data_generation(self,input_x):\n        tmp_imgs=np.zeros((self.batch_size,*self.dims))\n        \n        for ix,rows in enumerate(input_x):\n            #Read Image\n            img=pyd.read_file(os.path.join(images_path,input_x[ix]+'.dcm')).pixel_array\n            img_shape=img.shape\n            img=cv2.resize(img,(self.dims[0],self.dims[1]))\n            img=np.expand_dims(img,2)\n            \n            \n            # Augmentation\n            if self.is_train:\n                img=self.__augmentation(img)\n                       \n            tmp_imgs[ix]=img.astype('float')/255.\n            \n        return tmp_imgs\n    \n    def __augmentation(self,image):\n        transform = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Transpose(p=0.5),])\n        t=transform(image=image)\n        return t['image']","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:03.917718Z","iopub.execute_input":"2021-09-12T06:35:03.918466Z","iopub.status.idle":"2021-09-12T06:35:03.931985Z","shell.execute_reply.started":"2021-09-12T06:35:03.918425Z","shell.execute_reply":"2021-09-12T06:35:03.931249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Get Generator Object\ntrain_gen=KustomGenerator([X_train,y_train])\nval_gen=KustomGenerator([X_val,y_val],is_train=False)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:06.615921Z","iopub.execute_input":"2021-09-12T06:35:06.6166Z","iopub.status.idle":"2021-09-12T06:35:06.622007Z","shell.execute_reply.started":"2021-09-12T06:35:06.616561Z","shell.execute_reply":"2021-09-12T06:35:06.621156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = [X_train,y_train]\nindexes = [0,1,2,3]\ninput_para= x[0][:,1:3]\nX_para = [input_para[i] for i in indexes]\nX_para","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:31:21.471685Z","iopub.execute_input":"2021-09-12T06:31:21.471952Z","iopub.status.idle":"2021-09-12T06:31:21.480575Z","shell.execute_reply.started":"2021-09-12T06:31:21.471924Z","shell.execute_reply":"2021-09-12T06:31:21.479644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ninputs = layers.Input(shape=(HEIGHT, WIDTH, 1))\ninp=Input(2)\nx_=Dense(512,activation='relu')(inp)\n\n\nx=layers.Conv2D(3,(5,5),1,padding='same')(inputs)\nx=layers.LayerNormalization()(x)\nx=layers.Activation('relu')(x)\n\nbase_feat = EfficientNetB0(include_top=False, weights='imagenet')\nfor layer in base_feat.layers:\n    layer.trainable=True\n\nbase_feat=base_feat(x)    \nbase_feat=layers.GlobalAveragePooling2D()(base_feat)\noutputs=layers.Dense(3,activation='sigmoid')(x)\n\nmodel = tf.keras.Model(inputs, outputs)\n\nmodel.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"acc\"])\nmodel.summary()\n'''","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:11:45.06484Z","iopub.execute_input":"2021-09-12T06:11:45.065101Z","iopub.status.idle":"2021-09-12T06:11:45.074706Z","shell.execute_reply.started":"2021-09-12T06:11:45.06507Z","shell.execute_reply":"2021-09-12T06:11:45.073915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs = layers.Input(shape=(HEIGHT, WIDTH, 1))\nx=layers.Conv2D(3,(5,5),1,padding='same')(inputs)\nx=layers.LayerNormalization()(x)\nx=layers.Activation('relu')(x)\nbase_feat = EfficientNetB0(include_top=False, weights='imagenet')(x)\n#for layer in base_feat.layer:\n#    layer.trainable=True\n#base_feat=base_feat(x)   \nbase_feat=layers.GlobalAveragePooling2D()(base_feat)\nx=layers.Dense(512,activation='relu')(base_feat)\nx = tf.keras.Model(inputs=inputs, outputs=x)\n\ninput_para = layers.Input(shape=(2,))\ny = layers.Dense(1024, activation=\"relu\")(input_para)\ny = layers.Dense(512, activation=\"relu\")(y)\ny = tf.keras.Model(inputs=input_para, outputs=y)\n\n\n# combine the output of the two branches\ncombined = layers.concatenate([x.output, y.output])\n\nz=layers.Dense(1024,activation='relu')(combined)\noutputs=layers.Dense(3,activation='sigmoid')(z)\n\nmodel = tf.keras.Model(inputs=[x.input, y.input], outputs=outputs)\n\nmodel.compile(optimizer=\"adam\", loss=\"binary_crossentropy\", metrics=[\"acc\"])\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:10.136717Z","iopub.execute_input":"2021-09-12T06:35:10.136975Z","iopub.status.idle":"2021-09-12T06:35:12.273764Z","shell.execute_reply.started":"2021-09-12T06:35:10.136946Z","shell.execute_reply":"2021-09-12T06:35:12.27312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Callbacks\nmc = ModelCheckpoint('best_model.h5',monitor='val_loss',mode='min',save_best_only=True)\nrop = ReduceLROnPlateau(monitor='val_loss',mode='min',patience=4,min_lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:12.54251Z","iopub.execute_input":"2021-09-12T06:35:12.54301Z","iopub.status.idle":"2021-09-12T06:35:12.54856Z","shell.execute_reply.started":"2021-09-12T06:35:12.542968Z","shell.execute_reply":"2021-09-12T06:35:12.547842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCH=10\nmodel.fit(train_gen,steps_per_epoch=train_gen.__len__(),\n          epochs=EPOCH,validation_data=val_gen,\n          validation_steps=val_gen.__len__(),callbacks=[mc,rop])","metadata":{"execution":{"iopub.status.busy":"2021-09-12T06:35:15.354664Z","iopub.execute_input":"2021-09-12T06:35:15.355247Z","iopub.status.idle":"2021-09-12T06:35:15.701921Z","shell.execute_reply.started":"2021-09-12T06:35:15.355205Z","shell.execute_reply":"2021-09-12T06:35:15.700697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}