{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":5048,"databundleVersionId":868335,"sourceType":"competition"}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import libraries","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport pickle\nimport numpy as np\nimport seaborn as sns\nimport secrets\nimport cv2\nfrom PIL import Image\nfrom PIL import ImageFile\nfrom sklearn.datasets import load_files\nimport matplotlib.pyplot as plt\nfrom keras.layers import Conv2D, MaxPooling2D, GlobalAveragePooling2D\nfrom keras.layers import Dropout, Flatten, Dense\nfrom keras.models import Sequential\nfrom keras.callbacks import ModelCheckpoint\nfrom keras.utils import to_categorical\nfrom sklearn.metrics import confusion_matrix\nfrom keras.preprocessing import image                  \nfrom tqdm import tqdm\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nimport random\n\n\n\n\nimport seaborn as sns\nfrom sklearn.metrics import accuracy_score,precision_score,recall_score,f1_score","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:36:13.668806Z","iopub.execute_input":"2024-04-01T00:36:13.669486Z","iopub.status.idle":"2024-04-01T00:36:13.676558Z","shell.execute_reply.started":"2024-04-01T00:36:13.669453Z","shell.execute_reply":"2024-04-01T00:36:13.675667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-01T00:36:16.655565Z","iopub.execute_input":"2024-04-01T00:36:16.656218Z","iopub.status.idle":"2024-04-01T00:36:16.660689Z","shell.execute_reply.started":"2024-04-01T00:36:16.656189Z","shell.execute_reply":"2024-04-01T00:36:16.659757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Visualization","metadata":{}},{"cell_type":"code","source":"train_dir = \"/kaggle/input/state-farm-distracted-driver-detection/imgs/train/\"","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:36:18.247798Z","iopub.execute_input":"2024-04-01T00:36:18.248598Z","iopub.status.idle":"2024-04-01T00:36:18.252594Z","shell.execute_reply.started":"2024-04-01T00:36:18.248568Z","shell.execute_reply":"2024-04-01T00:36:18.251793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = ['c0', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7', 'c8', 'c9']\n\n#randomly selecting one image per class\n\nimg_c0_path = train_dir + \"c0/\" + random.choice(os.listdir(train_dir + \"c0/\"))\nimg_c1_path = train_dir + \"c1/\" + random.choice(os.listdir(train_dir + \"c1/\"))\nimg_c2_path = train_dir + \"c2/\" + random.choice(os.listdir(train_dir + \"c2/\"))\nimg_c3_path = train_dir + \"c3/\" + random.choice(os.listdir(train_dir + \"c3/\"))\nimg_c4_path = train_dir + \"c4/\" + random.choice(os.listdir(train_dir + \"c4/\"))\nimg_c5_path = train_dir + \"c5/\" + random.choice(os.listdir(train_dir + \"c5/\"))\nimg_c6_path = train_dir + \"c6/\" + random.choice(os.listdir(train_dir + \"c6/\"))\nimg_c7_path = train_dir + \"c7/\" + random.choice(os.listdir(train_dir + \"c7/\"))\nimg_c8_path = train_dir + \"c8/\" + random.choice(os.listdir(train_dir + \"c8/\"))\nimg_c9_path = train_dir + \"c9/\" + random.choice(os.listdir(train_dir + \"c9/\"))\n\n# reading the images\nimgs = {}\n\nimgs[\"c0\"] = cv2.cvtColor(cv2.imread(img_c0_path), cv2.COLOR_BGR2RGB)\nimgs[\"c1\"] = cv2.cvtColor(cv2.imread(img_c1_path), cv2.COLOR_BGR2RGB)\nimgs[\"c2\"] = cv2.cvtColor(cv2.imread(img_c2_path), cv2.COLOR_BGR2RGB)\nimgs[\"c3\"] = cv2.cvtColor(cv2.imread(img_c3_path), cv2.COLOR_BGR2RGB)\nimgs[\"c4\"] = cv2.cvtColor(cv2.imread(img_c4_path), cv2.COLOR_BGR2RGB)\nimgs[\"c5\"] = cv2.cvtColor(cv2.imread(img_c5_path), cv2.COLOR_BGR2RGB)\nimgs[\"c6\"] = cv2.cvtColor(cv2.imread(img_c6_path), cv2.COLOR_BGR2RGB)\nimgs[\"c7\"] = cv2.cvtColor(cv2.imread(img_c7_path), cv2.COLOR_BGR2RGB)\nimgs[\"c8\"] = cv2.cvtColor(cv2.imread(img_c8_path), cv2.COLOR_BGR2RGB)\nimgs[\"c9\"] = cv2.cvtColor(cv2.imread(img_c9_path), cv2.COLOR_BGR2RGB)\n\n# displaying the images\ndisplay_rows = 5\ndisplay_columns = 5\nfig = plt.figure(figsize=(20, 20))\n\nfor label, index in zip(labels, range(10)):\n    fig.add_subplot(display_rows, display_columns, index + 1)\n    plt.imshow(imgs[label])\n    plt.axis('off')\n    plt.title(label)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:36:18.887399Z","iopub.execute_input":"2024-04-01T00:36:18.887777Z","iopub.status.idle":"2024-04-01T00:36:22.246507Z","shell.execute_reply.started":"2024-04-01T00:36:18.887746Z","shell.execute_reply":"2024-04-01T00:36:22.24556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data distribution","metadata":{}},{"cell_type":"code","source":"samples = {}\nfor label in labels:\n    samples[label] = [len(os.listdir(train_dir + label + \"/\"))]\n\nsamples_df = pd.DataFrame.from_dict(samples)\nsns.barplot(data = samples_df)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:36:41.336465Z","iopub.execute_input":"2024-04-01T00:36:41.337356Z","iopub.status.idle":"2024-04-01T00:36:41.629566Z","shell.execute_reply.started":"2024-04-01T00:36:41.337325Z","shell.execute_reply":"2024-04-01T00:36:41.628683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Hyperparameters","metadata":{}},{"cell_type":"code","source":"batch_size = 32\nbase_learning_rate = 0.001\nbase_momentum=0.9\ninitial_epochs = 30\nfine_tune_epochs = 40\nfine_tune_at = 4\ndrop_out_rate = 0.2\ninitial_patience = 3\nfine_tune_patience = 2","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:40:24.001055Z","iopub.execute_input":"2024-04-01T00:40:24.001674Z","iopub.status.idle":"2024-04-01T00:40:24.00635Z","shell.execute_reply.started":"2024-04-01T00:40:24.001645Z","shell.execute_reply":"2024-04-01T00:40:24.005447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data loading and augmentation","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\nIMG_SIZE = (224, 224)\nIMG_SHAPE = IMG_SIZE + (3,)\nBASE_MODEL_LAYER = 1\nbatch_size = 50\n\npreprocess_input = tf.keras.applications.resnet50.preprocess_input\n\ndata_augmentation = ImageDataGenerator(\n    rotation_range=20,\n    width_shift_range=0.2,\n    height_shift_range=0.2,\n    zoom_range=0.2,\n    horizontal_flip=True,\n    fill_mode='nearest',\n    preprocessing_function=preprocess_input, # I have to make it such that some of the augmented images are the focused versions. Not everyone. Now it would be probabilistically everyone as I am preprocessing it like this with probability = 0.5\n    validation_split=0.2\n)\n\ntraining_dataset = data_augmentation.flow_from_directory(train_dir,\n                                                        subset=\"training\",\n                                                        target_size=IMG_SIZE,\n                                                        batch_size=batch_size,\n                                                        class_mode='sparse')\nvalidation_dataset = data_augmentation.flow_from_directory(train_dir,\n                                                        subset=\"validation\",\n                                                        target_size=IMG_SIZE,\n                                                        batch_size=batch_size,\n                                                        class_mode='sparse')","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:40:24.504872Z","iopub.execute_input":"2024-04-01T00:40:24.505562Z","iopub.status.idle":"2024-04-01T00:40:35.305083Z","shell.execute_reply.started":"2024-04-01T00:40:24.505534Z","shell.execute_reply":"2024-04-01T00:40:35.304334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ResNet 50 Transfer Learning","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras.applications import ResNet50\n\nconv_base = ResNet50(input_shape=IMG_SHAPE,include_top=False,weights='imagenet')","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:40:35.306565Z","iopub.execute_input":"2024-04-01T00:40:35.306873Z","iopub.status.idle":"2024-04-01T00:40:36.385796Z","shell.execute_reply.started":"2024-04-01T00:40:35.306848Z","shell.execute_reply":"2024-04-01T00:40:36.384925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conv_base.summary()","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:40:36.386972Z","iopub.execute_input":"2024-04-01T00:40:36.387266Z","iopub.status.idle":"2024-04-01T00:40:36.613498Z","shell.execute_reply.started":"2024-04-01T00:40:36.387242Z","shell.execute_reply":"2024-04-01T00:40:36.612658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras import models\nfrom tensorflow.keras import layers\nimport tensorflow as tf\n\nmodel = models.Sequential()\nmodel.add(conv_base)\nmodel.add(layers.Flatten())\n# model.add(layers.Dense(512, activation='relu'))\n# model.add(layers.Dropout(0.2))\nmodel.add(layers.Dense(256, activation='relu'))\nmodel.add(layers.Dropout(0.2))\nmodel.add(layers.Dense(10, activation='softmax'))","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:40:36.615544Z","iopub.execute_input":"2024-04-01T00:40:36.61588Z","iopub.status.idle":"2024-04-01T00:40:36.62748Z","shell.execute_reply.started":"2024-04-01T00:40:36.615844Z","shell.execute_reply":"2024-04-01T00:40:36.626606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.build(input_shape=(None, 224, 224, 3))","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:00.821218Z","iopub.execute_input":"2024-04-01T00:41:00.821822Z","iopub.status.idle":"2024-04-01T00:41:00.831462Z","shell.execute_reply.started":"2024-04-01T00:41:00.821792Z","shell.execute_reply":"2024-04-01T00:41:00.830645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:03.578563Z","iopub.execute_input":"2024-04-01T00:41:03.578924Z","iopub.status.idle":"2024-04-01T00:41:03.607787Z","shell.execute_reply.started":"2024-04-01T00:41:03.578899Z","shell.execute_reply":"2024-04-01T00:41:03.606915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('This is the number of trainable weights '\n      'before freezing the conv base:', len(model.trainable_weights))","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:17.191195Z","iopub.execute_input":"2024-04-01T00:41:17.19154Z","iopub.status.idle":"2024-04-01T00:41:17.197482Z","shell.execute_reply.started":"2024-04-01T00:41:17.191513Z","shell.execute_reply":"2024-04-01T00:41:17.196504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conv_base.trainable = False","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:18.590691Z","iopub.execute_input":"2024-04-01T00:41:18.591397Z","iopub.status.idle":"2024-04-01T00:41:18.601167Z","shell.execute_reply.started":"2024-04-01T00:41:18.59136Z","shell.execute_reply":"2024-04-01T00:41:18.600047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('This is the number of trainable weights '\n      'after freezing the conv base:', len(model.trainable_weights))","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:19.608147Z","iopub.execute_input":"2024-04-01T00:41:19.608991Z","iopub.status.idle":"2024-04-01T00:41:19.614441Z","shell.execute_reply.started":"2024-04-01T00:41:19.608959Z","shell.execute_reply":"2024-04-01T00:41:19.613554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.summary()","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:21.426442Z","iopub.execute_input":"2024-04-01T00:41:21.427202Z","iopub.status.idle":"2024-04-01T00:41:21.45603Z","shell.execute_reply.started":"2024-04-01T00:41:21.427172Z","shell.execute_reply":"2024-04-01T00:41:21.455115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.utils import plot_model\nplot_model(model)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:22.608544Z","iopub.execute_input":"2024-04-01T00:41:22.609391Z","iopub.status.idle":"2024-04-01T00:41:23.175456Z","shell.execute_reply.started":"2024-04-01T00:41:22.609358Z","shell.execute_reply":"2024-04-01T00:41:23.174543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras import optimizers\nfrom tensorflow.keras.losses import SparseCategoricalCrossentropy\n\nmy_optimizer=tf.keras.optimizers.SGD(learning_rate=base_learning_rate,\n                                            momentum=base_momentum)\n\n\nmodel.compile(optimizer=my_optimizer,\n                loss=SparseCategoricalCrossentropy(),\n                metrics=['accuracy'])","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:32.299526Z","iopub.execute_input":"2024-04-01T00:41:32.300295Z","iopub.status.idle":"2024-04-01T00:41:32.317214Z","shell.execute_reply.started":"2024-04-01T00:41:32.300265Z","shell.execute_reply":"2024-04-01T00:41:32.316302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model.input_shape)\nprint(model.output_shape)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:35.455935Z","iopub.execute_input":"2024-04-01T00:41:35.456283Z","iopub.status.idle":"2024-04-01T00:41:35.461654Z","shell.execute_reply.started":"2024-04-01T00:41:35.456255Z","shell.execute_reply":"2024-04-01T00:41:35.46067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_path = \"training_1/transferLearning.weights.h5\"\ncheckpoint_dir = os.path.dirname(checkpoint_path)\n\n# create checkpoint callback\ncp_callback = tf.keras.callbacks.ModelCheckpoint(checkpoint_path,\n                                                save_weights_only=True,\n                                                verbose=1)\nes = tf.keras.callbacks.EarlyStopping(monitor='val_loss', mode='min', verbose=1, patience=initial_patience, restore_best_weights=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:39.965353Z","iopub.execute_input":"2024-04-01T00:41:39.965766Z","iopub.status.idle":"2024-04-01T00:41:39.971659Z","shell.execute_reply.started":"2024-04-01T00:41:39.965735Z","shell.execute_reply":"2024-04-01T00:41:39.970627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nhistory = model.fit(\n      training_dataset,\n      epochs=40,\n      validation_data=validation_dataset,callbacks = [cp_callback,es])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fine Tunning","metadata":{}},{"cell_type":"code","source":"checkpoint_path2 = \"training_1/fineTuning.weights.h5\"\ncheckpoint_dir = os.path.dirname(checkpoint_path)\n\nconv_base.trainable = True\nfor layer in conv_base.layers[:-15]:\n    layer.trainable=False\n\n# for layer in conv_base.layers:\n#     print(layer.trainable)    ","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:47.80527Z","iopub.execute_input":"2024-04-01T00:41:47.8056Z","iopub.status.idle":"2024-04-01T00:41:47.818657Z","shell.execute_reply.started":"2024-04-01T00:41:47.805574Z","shell.execute_reply":"2024-04-01T00:41:47.817719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_weights(\"/kaggle/working/training_1/transferLearning.weights.h5\")","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:41:50.753259Z","iopub.execute_input":"2024-04-01T00:41:50.753669Z","iopub.status.idle":"2024-04-01T00:41:50.958632Z","shell.execute_reply.started":"2024-04-01T00:41:50.753614Z","shell.execute_reply":"2024-04-01T00:41:50.957437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fine_tuning_optimizer =tf.keras.optimizers.Adam(learning_rate=0.00001)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cp_callback2 = tf.keras.callbacks.ModelCheckpoint(checkpoint_path2,\n                                                save_weights_only=True,\n                                                verbose=1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.compile(optimizer=fine_tuning_optimizer,\n                loss=SparseCategoricalCrossentropy(),\n                metrics=['accuracy'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n      training_dataset,\n      epochs=60,\n      validation_data=validation_dataset,callbacks = [cp_callback2])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test","metadata":{}},{"cell_type":"code","source":"model.load_weights(\"/kaggle/working/training_1/fineTuning.weights.h5\")","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:42:03.775886Z","iopub.execute_input":"2024-04-01T00:42:03.776542Z","iopub.status.idle":"2024-04-01T00:42:04.39102Z","shell.execute_reply.started":"2024-04-01T00:42:03.776509Z","shell.execute_reply":"2024-04-01T00:42:04.390124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Analysis","metadata":{}},{"cell_type":"code","source":"test_path = '/kaggle/input/state-farm-distracted-driver-detection/imgs/test'\n\ntest_datagen = ImageDataGenerator(\n    preprocessing_function=preprocess_input\n)\n\ntest_dataset = test_datagen.flow_from_directory(\n    directory='/kaggle/input/state-farm-distracted-driver-detection/imgs/',\n    target_size=IMG_SIZE,\n    batch_size=32,\n    classes=['test'],\n    class_mode=None,\n    shuffle=False\n)\n\n# Get filenames and predictions\nfile_names = sorted(os.listdir(test_path))\nlen(file_names)\n\npredictions = model.predict(test_dataset, verbose=1)\n\nfile_names, predictions = zip(*sorted(zip(file_names, predictions.tolist())))\n\n# Vectorized concatenation of file names with predictions\narr = np.column_stack([file_names, predictions])\n\n# Create a DataFrame with the desired columns and save to CSV\nsubmission_df = pd.DataFrame(arr, columns=['img', 'c0', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7', 'c8', 'c9'])\nsubmission_df.to_csv('./submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:42:15.199661Z","iopub.execute_input":"2024-04-01T00:42:15.200012Z","iopub.status.idle":"2024-04-01T00:52:31.640298Z","shell.execute_reply.started":"2024-04-01T00:42:15.199984Z","shell.execute_reply":"2024-04-01T00:52:31.638745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ytest = validation_dataset.classes.tolist()\n# print(ytest)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T00:59:21.926773Z","iopub.execute_input":"2024-04-01T00:59:21.927145Z","iopub.status.idle":"2024-04-01T00:59:21.931585Z","shell.execute_reply.started":"2024-04-01T00:59:21.927117Z","shell.execute_reply":"2024-04-01T00:59:21.930663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = model.predict(validation_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T01:00:03.818625Z","iopub.execute_input":"2024-04-01T01:00:03.819002Z","iopub.status.idle":"2024-04-01T01:00:59.150089Z","shell.execute_reply.started":"2024-04-01T01:00:03.818972Z","shell.execute_reply":"2024-04-01T01:00:59.149166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Convert predictions to class labels (assuming one-hot encoding)\npredicted_labels = np.argmax(predictions, axis=1)\n# print(predicted_labels)\ntrue_labels = ytest\n# print(ytest)\n\n# Generate confusion matrix\ncm = confusion_matrix(true_labels, predicted_labels)\n\n# Plot heatmap\nplt.figure(figsize=(20, 20))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['Class 0', 'Class 1', 'Class 2'], yticklabels=['Class 0', 'Class 1', 'Class 2'])\nplt.xlabel('Predicted Labels')\nplt.ylabel('True Labels')\nplt.title('Confusion Matrix')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-01T01:08:03.106742Z","iopub.execute_input":"2024-04-01T01:08:03.107098Z","iopub.status.idle":"2024-04-01T01:08:03.721871Z","shell.execute_reply.started":"2024-04-01T01:08:03.107071Z","shell.execute_reply":"2024-04-01T01:08:03.720861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predictions = model.predict(validation_dataset, verbose=1)\n# ytest = validation_dataset.classes.tolist()\n\naccuracy = accuracy_score(ytest,predicted_labels)\nprint('Accuracy: %f' % accuracy)\n# precision tp / (tp + fp)\nprecision = precision_score(ytest, predicted_labels,average='weighted')\nprint('Precision: %f' % precision)\n# recall: tp / (tp + fn)\nrecall = recall_score(ytest,predicted_labels,average='weighted')\nprint('Recall: %f' % recall)\n# f1: 2 tp / (2 tp + fp + fn)\nf1 = f1_score(ytest,predicted_labels,average='weighted')\nprint('F1 score: %f' % f1)","metadata":{"execution":{"iopub.status.busy":"2024-04-01T01:08:31.37117Z","iopub.execute_input":"2024-04-01T01:08:31.372015Z","iopub.status.idle":"2024-04-01T01:08:31.405449Z","shell.execute_reply.started":"2024-04-01T01:08:31.371986Z","shell.execute_reply":"2024-04-01T01:08:31.40449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}