{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":25563,"databundleVersionId":2094376,"sourceType":"competition"},{"sourceId":17498,"sourceType":"modelInstanceVersion","modelInstanceId":14561},{"sourceId":19839,"sourceType":"modelInstanceVersion","modelInstanceId":16457}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Plant Pathology","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"%matplotlib inline\n\nimport pandas as pd\nimport glob\nimport os\nimport matplotlib\nfrom matplotlib import pyplot as plt\nimport matplotlib.image as mpimg\nimport numpy as np\nimport imageio as im\nimport keras\nfrom keras.callbacks import History, EarlyStopping\nfrom sklearn.model_selection import train_test_split\nimport tensorflow as tf\nfrom tensorflow.keras import models\nfrom tensorflow.keras import regularizers\nfrom tensorflow.keras.models import Sequential\nfrom keras.models import load_model\nfrom tensorflow.keras.layers import Conv2D\nfrom tensorflow.keras.layers import MaxPooling2D\nfrom tensorflow.keras.layers import Flatten\nfrom tensorflow.keras.layers import Dense\nfrom tensorflow.keras.layers import Dropout\nfrom tensorflow.keras.layers import BatchNormalization\nfrom tensorflow.keras.preprocessing import image\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.optimizers import Adam, Adamax\n#from tensorflow.keras.utils import image_dataset_from_directory\nfrom tensorflow.keras.callbacks import ModelCheckpoint","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:33.390352Z","iopub.execute_input":"2024-04-24T19:46:33.390707Z","iopub.status.idle":"2024-04-24T19:46:33.405575Z","shell.execute_reply.started":"2024-04-24T19:46:33.390679Z","shell.execute_reply":"2024-04-24T19:46:33.404579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Investigating the Dataset:","metadata":{}},{"cell_type":"code","source":"dataset = pd.read_csv('/kaggle/input/plant-pathology-2021-fgvc8/train.csv')\nprint(dataset)\ndataset_test = pd.read_csv('/kaggle/input/plant-pathology-2021-fgvc8/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:36.092118Z","iopub.execute_input":"2024-04-24T19:46:36.092528Z","iopub.status.idle":"2024-04-24T19:46:36.159751Z","shell.execute_reply.started":"2024-04-24T19:46:36.092500Z","shell.execute_reply":"2024-04-24T19:46:36.158577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"categories = dataset['labels'].value_counts().index\ncounts = dataset['labels'].value_counts().values\nplt.bar(categories, counts, width=0.5)\nplt.xticks(rotation=30, ha='right')\nplt.title('Count of labels in training data')","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:38.140161Z","iopub.execute_input":"2024-04-24T19:46:38.141088Z","iopub.status.idle":"2024-04-24T19:46:38.542074Z","shell.execute_reply.started":"2024-04-24T19:46:38.141053Z","shell.execute_reply":"2024-04-24T19:46:38.541041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The dataset contains links to an associated image base, as well as labels for the types of diseases. The dataset matches that of the alternative dataset provided in the Canvas project description. The alternative dataset contained pre-made labels to create a test/train/validation split, reserving 1000 records each for testing and validation, leading to roughly 90% of the data being used for training, 5% for testing, and 5% for validation. Our model will later utilize an 80/10/10 split for training/testing/validation. ","metadata":{}},{"cell_type":"code","source":"class_labels = [] #Holds unique class labels\ncount = 0\nfor label in dataset.iloc[:,1]:\n    count+=1\n    if label not in class_labels:\n        class_labels.append(label)\nprint(class_labels)\nprint('The dataset contains', len(class_labels), 'labels.')\nprint('The dataset contains', count, 'total images.')","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:40.571690Z","iopub.execute_input":"2024-04-24T19:46:40.572388Z","iopub.status.idle":"2024-04-24T19:46:40.587387Z","shell.execute_reply.started":"2024-04-24T19:46:40.572350Z","shell.execute_reply":"2024-04-24T19:46:40.586176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The dataset contains 18,632 total records, and appears to be spread among 12 different classes. However, some plants are labelled with multiple diseases (such as 'scab frog_eye_leaf_spot'). We ultimately would like to be able to detect such combinations, but must first separate the values so we can determine how many classes the model should be able to detect.","metadata":{}},{"cell_type":"markdown","source":"There are only 6 unique classes in the dataset including healthy samples, but each sample may exhibit multiple diseases at once.","metadata":{}},{"cell_type":"code","source":"unique_split_class_labels = []\nsplit_class_labels = []\nfor label in dataset.iloc[:,1]:\n    split_label = label.split()\n    split_class_labels.append(split_label)\n    for value in split_label:\n        if value not in unique_split_class_labels:\n            unique_split_class_labels.append(value)\nprint(unique_split_class_labels)\nprint('The dataset contains', len(unique_split_class_labels), 'unique class labels.')","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:45.253245Z","iopub.execute_input":"2024-04-24T19:46:45.253871Z","iopub.status.idle":"2024-04-24T19:46:45.279418Z","shell.execute_reply.started":"2024-04-24T19:46:45.253820Z","shell.execute_reply":"2024-04-24T19:46:45.278439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_list = dataset.iloc[:4,0]\nimg_set = []\nfor path in path_list:\n    img_set.append(mpimg.imread(\"/kaggle/input/plant-pathology-2021-fgvc8/train_images/\" + path))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:47.057426Z","iopub.execute_input":"2024-04-24T19:46:47.058037Z","iopub.status.idle":"2024-04-24T19:46:47.650969Z","shell.execute_reply.started":"2024-04-24T19:46:47.057992Z","shell.execute_reply":"2024-04-24T19:46:47.649720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"figure, axarr = plt.subplots(2,2)\naxarr[0,0].imshow(img_set[0])\naxarr[0,1].imshow(img_set[1])\naxarr[1,0].imshow(img_set[2])\naxarr[1,1].imshow(img_set[3])\nplt.show()\nprint(np.shape(img_set[0]))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:49.378772Z","iopub.execute_input":"2024-04-24T19:46:49.379153Z","iopub.status.idle":"2024-04-24T19:46:54.534624Z","shell.execute_reply.started":"2024-04-24T19:46:49.379123Z","shell.execute_reply":"2024-04-24T19:46:54.533600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data Preprocessing (FCN)\nThe images are large RGB images. To prepare the data for use in a CNN, we perform the following preprocessing steps:\n\n1. Create an encoding capable of dealing with multilabelled data. \n2. Create an ImageDataGenerator for the training, validation, and testing sets\n3. Perform data augmentations for the requisite model, such as normalizing the data (divide each value by 255) or rotating/flipping/zooming\n","metadata":{}},{"cell_type":"markdown","source":"#### Build the array","metadata":{}},{"cell_type":"code","source":"#Partition the multilabels for the dataset\ndataset[\"labels\"]=dataset[\"labels\"].apply(lambda x:x.split(\", \"))","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:46:56.748026Z","iopub.execute_input":"2024-04-24T19:46:56.749075Z","iopub.status.idle":"2024-04-24T19:46:57.813382Z","shell.execute_reply.started":"2024-04-24T19:46:56.749039Z","shell.execute_reply":"2024-04-24T19:46:57.812077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create a train/test/valid split\n*80% of data for training\n\n*10% of data for validation\n\n*10% of data for testing","metadata":{}},{"cell_type":"code","source":"#Create a train/valid/test split\ntrain_split = 0.8\ntest_split = 0.5 #ratio of test:validation samples after splitting from train split\n\ntrain_df, valid_test_df = train_test_split(dataset, train_size=train_split, shuffle=True, random_state = 42)\nvalid_df, test_df = train_test_split(valid_test_df, train_size=test_split, shuffle=True, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:47:03.025880Z","iopub.execute_input":"2024-04-24T19:47:03.026520Z","iopub.status.idle":"2024-04-24T19:47:03.040612Z","shell.execute_reply.started":"2024-04-24T19:47:03.026477Z","shell.execute_reply":"2024-04-24T19:47:03.039484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## My CNN Model","metadata":{}},{"cell_type":"code","source":"#Create data generators for training, testing, and validation sets (Dont run this one yet)\ndatagen = ImageDataGenerator(rotation_range=30, fill_mode='nearest',\n                            width_shift_range=0.2, height_shift_range=0.2,\n                            horizontal_flip=True, vertical_flip=True,\n                            zoom_range=0.1) #functions to apply to training set\n\ntest_datagen = ImageDataGenerator() #functions to apply to testing/validation sets\n\ntrain_generator=datagen.flow_from_dataframe(dataframe=train_df,\n                                            directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                            x_col=\"image\",\n                                            y_col=\"labels\",\n                                            batch_size=32,\n                                            seed=42,\n                                            shuffle=True,\n                                            class_mode=\"categorical\",\n                                            classes=unique_split_class_labels,\n                                            target_size=(256,256))\n\nvalid_generator=test_datagen.flow_from_dataframe(dataframe=valid_df,\n                                                 directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                                 x_col=\"image\",\n                                                 y_col=\"labels\",\n                                                 batch_size=32,\n                                                 seed=42,\n                                                 shuffle=True,\n                                                 class_mode=\"categorical\",\n                                                 classes=unique_split_class_labels,\n                                                 target_size=(256,256))\n\ntest_generator=test_datagen.flow_from_dataframe(dataframe=test_df,\n                                                directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                                x_col=\"image\",\n                                                y_col=\"labels\",\n                                                batch_size=1,\n                                                shuffle=False,\n                                                class_mode=\"categorical\",\n                                                classes=unique_split_class_labels,\n                                                target_size=(256,256))\n\n#directory=\"/kaggle/input/plant-pathology-2021-fgvc8/test_images\", dataframe=dataset_test[:]","metadata":{"execution":{"iopub.status.busy":"2024-04-24T22:52:31.844896Z","iopub.execute_input":"2024-04-24T22:52:31.845602Z","iopub.status.idle":"2024-04-24T22:52:40.315863Z","shell.execute_reply.started":"2024-04-24T22:52:31.845571Z","shell.execute_reply":"2024-04-24T22:52:40.314896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Sequential()\n    \nmodel.add(Conv2D(32, (3, 3), padding='same', input_shape = (256, 256, 3), activation = 'relu'))\nmodel.add(Conv2D(32, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.25)) \n\n# Adding a second convolutional block\nmodel.add(Conv2D(64, (3, 3), padding='same', activation = 'relu'))\nmodel.add(Conv2D(64, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.25)) \n\n# Adding a third convolutional block\nmodel.add(Conv2D(64, (3, 3), padding='same', activation = 'relu'))\nmodel.add(Conv2D(64, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.25))\n\n# Adding a fourth convolutional block\nmodel.add(Conv2D(128, (3, 3), padding='same', activation = 'relu'))\nmodel.add(Conv2D(128, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.25))\n\n# Adding a fifth convolutional block\nmodel.add(Conv2D(128, (3, 3), padding='same', activation = 'relu'))\nmodel.add(Conv2D(128, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D(pool_size=(2, 2)))\nmodel.add(Dropout(0.25))\n\n# Step 3 - Flattening\nmodel.add(Flatten())\n\n# Step 4 - Full connection\nmodel.add(Dense(units = 256, activation = 'sigmoid'))\nmodel.add(Dropout(0.1)) \nmodel.add(Dense(units = 6, activation = 'sigmoid'))\n\n# Compile model\nmodel.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy'])\n\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:49:28.541605Z","iopub.execute_input":"2024-04-24T19:49:28.542330Z","iopub.status.idle":"2024-04-24T19:49:29.688746Z","shell.execute_reply.started":"2024-04-24T19:49:28.542299Z","shell.execute_reply":"2024-04-24T19:49:29.687700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Setting up callbacks","metadata":{}},{"cell_type":"markdown","source":"Train the model","metadata":{}},{"cell_type":"code","source":"checkpoint_filepath = 'M1-epoch-{epoch:02d}.weights.h5'\ncheckpoint = keras.callbacks.ModelCheckpoint(\n        filepath=checkpoint_filepath,\n        save_weights_only=True,\n        monitor='val_accuracy',\n        mode='max',\n        save_best_only=True,\n        save_freq=10)\nes = EarlyStopping(monitor='val_accuracy',mode='max', patience=5)\ncallbacks_list = [checkpoint, es]","metadata":{"execution":{"iopub.status.busy":"2024-04-24T19:49:34.117438Z","iopub.execute_input":"2024-04-24T19:49:34.117794Z","iopub.status.idle":"2024-04-24T19:49:34.123433Z","shell.execute_reply.started":"2024-04-24T19:49:34.117766Z","shell.execute_reply":"2024-04-24T19:49:34.122438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nnp.random.seed(0)\nhistory = model.fit(train_generator, validation_data=valid_generator, epochs=10, batch_size=None, callbacks=callbacks_list)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:08:20.971318Z","iopub.execute_input":"2024-04-24T20:08:20.971725Z","iopub.status.idle":"2024-04-24T22:23:04.659526Z","shell.execute_reply.started":"2024-04-24T20:08:20.971692Z","shell.execute_reply":"2024-04-24T22:23:04.658468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# evaluation of the model on Test split\nscores = model.evaluate(test_generator, verbose=0)\nprint('Final CNN accuracy: ', scores[1])","metadata":{"execution":{"iopub.status.busy":"2024-04-24T22:52:51.992236Z","iopub.execute_input":"2024-04-24T22:52:51.992590Z","iopub.status.idle":"2024-04-24T22:54:52.445311Z","shell.execute_reply.started":"2024-04-24T22:52:51.992564Z","shell.execute_reply":"2024-04-24T22:54:52.444283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = model.predict(test_generator)\ny_pred = np.argmax(preds, axis=1)\nprint(y_pred)\n\nprint(preds)\nprint(test_generator.class_indices)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T22:56:17.769821Z","iopub.execute_input":"2024-04-24T22:56:17.770587Z","iopub.status.idle":"2024-04-24T22:57:56.539438Z","shell.execute_reply.started":"2024-04-24T22:56:17.770547Z","shell.execute_reply":"2024-04-24T22:57:56.538262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('model loss')\nplt.ylabel('loss')\nplt.xlabel('epoch')\nplt.legend(['train', 'val'], loc='upper left')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-24T23:21:30.453933Z","iopub.execute_input":"2024-04-24T23:21:30.454348Z","iopub.status.idle":"2024-04-24T23:21:30.765464Z","shell.execute_reply.started":"2024-04-24T23:21:30.454317Z","shell.execute_reply":"2024-04-24T23:21:30.764461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EfficientNet Model","metadata":{}},{"cell_type":"code","source":"#Create data generators for training, testing, and validation sets\n#Note that EfficientNetB3 does not require normalizing the color pixels like hand-made model\ndatagen = ImageDataGenerator(horizontal_flip=True)\n\ntest_datagen = ImageDataGenerator() #functions to apply to testing/validation sets\n\nsubmission_datagen = ImageDataGenerator() #The test samples used to submit to Kaggle\n\ntrain_generator=datagen.flow_from_dataframe(dataframe=train_df,\n                                            directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                            x_col=\"image\",\n                                            y_col=\"labels\",\n                                            batch_size=32,\n                                            seed=42,\n                                            shuffle=True,\n                                            class_mode=\"categorical\",\n                                            classes=unique_split_class_labels,\n                                            target_size=(224,224))\n\nvalid_generator=test_datagen.flow_from_dataframe(dataframe=valid_df,\n                                                 directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                                 x_col=\"image\",\n                                                 y_col=\"labels\",\n                                                 batch_size=32,\n                                                 seed=42,\n                                                 shuffle=True,\n                                                 class_mode=\"categorical\",\n                                                 classes=unique_split_class_labels,\n                                                 target_size=(224,224))\n\ntest_generator=test_datagen.flow_from_dataframe(dataframe=test_df,\n                                                directory=\"/kaggle/input/plant-pathology-2021-fgvc8/train_images\",\n                                                x_col=\"image\",\n                                                y_col=\"labels\",\n                                                batch_size=1,\n                                                shuffle=False,\n                                                class_mode=\"categorical\",\n                                                classes=unique_split_class_labels,\n                                                target_size=(224,224))\n\nsubmission_generator=submission_datagen.flow_from_dataframe(dataframe=dataset_test,\n                                                directory=\"/kaggle/input/plant-pathology-2021-fgvc8/test_images\",\n                                                x_col=\"image\",\n                                                y_col=None,\n                                                batch_size=1,\n                                                shuffle=False,\n                                                class_mode=None,\n                                                #classes=unique_split_class_labels,\n                                                target_size=(224,224))","metadata":{"execution":{"iopub.status.busy":"2024-03-24T20:43:11.309100Z","iopub.execute_input":"2024-03-24T20:43:11.309755Z","iopub.status.idle":"2024-03-24T20:43:17.246009Z","shell.execute_reply.started":"2024-03-24T20:43:11.309721Z","shell.execute_reply":"2024-03-24T20:43:17.245190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Create an efficientnetb3 model\nbase_model = tf.keras.applications.efficientnet.EfficientNetB3(include_top= False, weights= \"imagenet\", input_shape= (224,224,3), pooling= 'max')\n\nmodel_eff = Sequential()\nmodel_eff.add(base_model)\nmodel_eff.add(BatchNormalization(axis= -1, momentum= 0.99, epsilon= 0.001))\nmodel_eff.add(Dense(64, kernel_regularizer= regularizers.l2(l2= 0.016), activity_regularizer= regularizers.l1(0.006),\n                bias_regularizer= regularizers.l1(0.006), activation= 'relu'))\nmodel_eff.add(Dropout(rate= 0.45, seed= 123))\nmodel_eff.add(Dense(units = 6, activation= 'sigmoid'))\n\nmodel_eff.compile(Adamax(learning_rate= 0.001), loss= 'categorical_crossentropy', metrics= ['accuracy'])\n#model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy'])\n\nmodel_eff.summary()","metadata":{"execution":{"iopub.status.busy":"2024-03-24T21:21:24.823462Z","iopub.execute_input":"2024-03-24T21:21:24.824166Z","iopub.status.idle":"2024-03-24T21:21:28.589914Z","shell.execute_reply.started":"2024-03-24T21:21:24.824136Z","shell.execute_reply":"2024-03-24T21:21:28.589058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nnp.random.seed(0)\n# Fit the model\nmodel_eff.fit(train_generator, validation_data=valid_generator, epochs=10, batch_size=None, callbacks=callbacks_list)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T04:26:02.620645Z","iopub.execute_input":"2024-03-21T04:26:02.621039Z","iopub.status.idle":"2024-03-21T06:51:23.218879Z","shell.execute_reply.started":"2024-03-21T04:26:02.621009Z","shell.execute_reply":"2024-03-21T06:51:23.217953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# evaluation of the model on Test split\nscores = model_eff.evaluate(test_generator, verbose=0)\nprint('Final EfficientNet CNN accuracy: ', scores[1])","metadata":{"execution":{"iopub.status.busy":"2024-03-21T06:53:48.834829Z","iopub.execute_input":"2024-03-21T06:53:48.835720Z","iopub.status.idle":"2024-03-21T06:55:48.893104Z","shell.execute_reply.started":"2024-03-21T06:53:48.835666Z","shell.execute_reply":"2024-03-21T06:55:48.892146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_eff.summary()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T06:58:00.097069Z","iopub.execute_input":"2024-03-21T06:58:00.097435Z","iopub.status.idle":"2024-03-21T06:58:00.147349Z","shell.execute_reply.started":"2024-03-21T06:58:00.097407Z","shell.execute_reply":"2024-03-21T06:58:00.146521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Looking at the output","metadata":{}},{"cell_type":"code","source":"preds = model.predict(test_generator)\ny_pred = np.argmax(preds, axis=1)\nprint(y_pred)","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-03-24T20:12:03.699040Z","iopub.execute_input":"2024-03-24T20:12:03.699982Z","iopub.status.idle":"2024-03-24T20:12:03.774490Z","shell.execute_reply.started":"2024-03-24T20:12:03.699948Z","shell.execute_reply":"2024-03-24T20:12:03.773400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(preds)\nprint(test_generator.class_indices)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T07:12:21.894179Z","iopub.execute_input":"2024-03-21T07:12:21.894536Z","iopub.status.idle":"2024-03-21T07:12:21.900185Z","shell.execute_reply.started":"2024-03-21T07:12:21.894507Z","shell.execute_reply":"2024-03-21T07:12:21.899246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"For a sample, the model can output a calculated likelihood of that label existing in the sample. To predict multiple labels on a sample, each entity above a certain threshold can be selected. Here, we (tbc...)","metadata":{}},{"cell_type":"code","source":"# Read the test data\ntest = pd.read_csv('../input/test.csv')\n# Treat the test data in the same way as training data. In this case, pull same columns.\ntest_X = test[predictor_cols]\n# Use the model to make predictions\npredicted_prices = my_model.predict(test_X)\n# We will look at the predicted prices to ensure we have something sensible.\nprint(predicted_prices)","metadata":{},"execution_count":null,"outputs":[]}]}