{"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":"code","source":"import tensorflow as tf\nprint(f'tensorflow version: {tf.__version__}')\nprint(f'tensorflow keras version: {tf.keras.__version__}')","metadata":{"execution":{"iopub.status.busy":"2023-05-17T07:57:57.518800Z","iopub.execute_input":"2023-05-17T07:57:57.519229Z","iopub.status.idle":"2023-05-17T07:57:57.525147Z","shell.execute_reply.started":"2023-05-17T07:57:57.519189Z","shell.execute_reply":"2023-05-17T07:57:57.524005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect hardware, return appropriate distribution strategy\ntry:\n    TPU = tf.distribute.cluster_resolver.TPUClusterResolver()  # TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.\n    print('Running on TPU ', TPU.master())\nexcept ValueError:\n    print('Running on GPU')\n    TPU = None\n\nif TPU:\n    tf.config.experimental_connect_to_cluster(TPU)\n    tf.tpu.experimental.initialize_tpu_system(TPU)\n    strategy = tf.distribute.experimental.TPUStrategy(TPU)\nelse:\n    strategy = tf.distribute.get_strategy() # default distribution strategy in Tensorflow. Works on CPU and single GPU.\n\nREPLICAS = strategy.num_replicas_in_sync\nprint(f'REPLICAS: {REPLICAS}')\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-17T07:58:01.218988Z","iopub.execute_input":"2023-05-17T07:58:01.219486Z","iopub.status.idle":"2023-05-17T07:58:10.969800Z","shell.execute_reply.started":"2023-05-17T07:58:01.219431Z","shell.execute_reply":"2023-05-17T07:58:10.968832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python -m pip install --upgrade pip\n!pip install vit_keras\n!pip install tensorflow_addons \n# !pip install tensorflow-gpu==2.6.0","metadata":{"execution":{"iopub.status.busy":"2023-05-17T07:58:18.432910Z","iopub.execute_input":"2023-05-17T07:58:18.433309Z","iopub.status.idle":"2023-05-17T07:58:35.749350Z","shell.execute_reply.started":"2023-05-17T07:58:18.433279Z","shell.execute_reply":"2023-05-17T07:58:35.748090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install keras_applications","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:00:02.053900Z","iopub.execute_input":"2023-05-17T08:00:02.054306Z","iopub.status.idle":"2023-05-17T08:00:06.615798Z","shell.execute_reply.started":"2023-05-17T08:00:02.054274Z","shell.execute_reply":"2023-05-17T08:00:06.614617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install albumentations ","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:01:44.043723Z","iopub.execute_input":"2023-05-17T08:01:44.044695Z","iopub.status.idle":"2023-05-17T08:01:58.654869Z","shell.execute_reply.started":"2023-05-17T08:01:44.044662Z","shell.execute_reply":"2023-05-17T08:01:58.653555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import keras_applications\nfrom vit_keras import vit, utils\nfrom tensorflow import keras\nimport tensorflow as tf\n\nimport os, glob\nimport random\nfrom sklearn.model_selection import train_test_split\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport multiprocessing\nfrom copy import deepcopy\nfrom sklearn.metrics import precision_recall_curve, auc\nimport tensorflow.keras as keras\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.callbacks import Callback\nfrom tensorflow.keras.applications.densenet import DenseNet121\nfrom tensorflow.keras.layers import Dense, Flatten\nfrom tensorflow.keras.models import Model, load_model\nfrom tensorflow.keras.utils import Sequence\nfrom albumentations import Compose, VerticalFlip, HorizontalFlip, Rotate, GridDistortion,CenterCrop\nimport matplotlib.pyplot as plt\nfrom IPython.display import Image\nfrom tqdm import tqdm_notebook as tqdm\nfrom numpy.random import seed\nseed(10)\ntf.random.set_seed(10)\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:05.937186Z","iopub.execute_input":"2023-05-17T08:02:05.937928Z","iopub.status.idle":"2023-05-17T08:02:05.952496Z","shell.execute_reply.started":"2023-05-17T08:02:05.937893Z","shell.execute_reply":"2023-05-17T08:02:05.951369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_imgs_folde = '../input/understanding_cloud_organization/test_images/'\ntrain_imgs_folder  = '../input/understanding_cloud_organization/train_images/'\nnum_cores = multiprocessing.cpu_count()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:35.105010Z","iopub.execute_input":"2023-05-17T08:02:35.105432Z","iopub.status.idle":"2023-05-17T08:02:35.110645Z","shell.execute_reply.started":"2023-05-17T08:02:35.105399Z","shell.execute_reply":"2023-05-17T08:02:35.109557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/understanding_cloud_organization/train.csv')\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:36.880005Z","iopub.execute_input":"2023-05-17T08:02:36.880883Z","iopub.status.idle":"2023-05-17T08:02:41.393587Z","shell.execute_reply.started":"2023-05-17T08:02:36.880851Z","shell.execute_reply":"2023-05-17T08:02:41.392046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = train_df[~train_df['EncodedPixels'].isnull()]\ntrain_df['Image'] = train_df['Image_Label'].map(lambda x: x.split('_')[0])\ntrain_df['Class'] = train_df['Image_Label'].map(lambda x: x.split('_')[1])\nclasses = train_df['Class'].unique()\ntrain_df = train_df.groupby('Image')['Class'].agg(set).reset_index()\nfor class_name in classes:\n    train_df[class_name] = train_df['Class'].map(lambda x: 1 if class_name in x else 0)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:43.129628Z","iopub.execute_input":"2023-05-17T08:02:43.130640Z","iopub.status.idle":"2023-05-17T08:02:43.383556Z","shell.execute_reply.started":"2023-05-17T08:02:43.130601Z","shell.execute_reply":"2023-05-17T08:02:43.382227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dictionary for fast access to ohe vectors\nimg_2_ohe_vector = {img:vec for img, vec in zip(train_df['Image'], train_df.iloc[:, 2:].values)}","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:45.703483Z","iopub.execute_input":"2023-05-17T08:02:45.703914Z","iopub.status.idle":"2023-05-17T08:02:45.715057Z","shell.execute_reply.started":"2023-05-17T08:02:45.703882Z","shell.execute_reply":"2023-05-17T08:02:45.713692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_imgs, val_imgs = train_test_split(train_df['Image'].values, \n                                        test_size=0.1, \n                                        stratify=train_df['Class'].map(lambda x: str(sorted(list(x)))), # sorting present classes in lexicographical order, just to be sure\n                                        random_state=43)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:47.264781Z","iopub.execute_input":"2023-05-17T08:02:47.266046Z","iopub.status.idle":"2023-05-17T08:02:47.295066Z","shell.execute_reply.started":"2023-05-17T08:02:47.265995Z","shell.execute_reply":"2023-05-17T08:02:47.293830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataGenenerator(Sequence):\n    def __init__(self, images_list=None, folder_imgs=train_imgs_folder, \n                 batch_size=32, shuffle=True, augmentation=None,\n                 resized_height=224, resized_width=224, num_channels=3):\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n        self.augmentation = augmentation\n        if images_list is None:\n            self.images_list = os.listdir(folder_imgs)\n        else:\n            self.images_list = deepcopy(images_list)\n        self.folder_imgs = folder_imgs\n        self.len = len(self.images_list) // self.batch_size\n        self.resized_height = resized_height\n        self.resized_width = resized_width\n        self.num_channels = num_channels\n        self.num_classes = 4\n        self.is_test = not 'train' in folder_imgs\n        if not shuffle and not self.is_test:\n            self.labels = [img_2_ohe_vector[img] for img in self.images_list[:self.len*self.batch_size]]\n\n    def __len__(self):\n        return self.len\n    \n    def on_epoch_start(self):\n        if self.shuffle:\n            random.shuffle(self.images_list)\n\n    def __getitem__(self, idx):\n        current_batch = self.images_list[idx * self.batch_size: (idx + 1) * self.batch_size]\n        X = np.empty((self.batch_size, self.resized_height, self.resized_width, self.num_channels))\n        y = np.empty((self.batch_size, self.num_classes))\n\n        for i, image_name in enumerate(current_batch):\n            path = os.path.join(self.folder_imgs, image_name)\n            img = cv2.resize(cv2.imread(path), (self.resized_height, self.resized_width)).astype(np.float32)\n            if not self.augmentation is None:\n                augmented = self.augmentation(image=img)\n                img = augmented['image']\n            X[i, :, :, :] = img/255.0\n            if not self.is_test:\n                y[i, :] = img_2_ohe_vector[image_name]\n        return X, y\n\n    def get_labels(self):\n        if self.shuffle:\n            images_current = self.images_list[:self.len*self.batch_size]\n            labels = [img_2_ohe_vector[img] for img in images_current]\n        else:\n            labels = self.labels\n        return np.array(labels)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:02:49.309930Z","iopub.execute_input":"2023-05-17T08:02:49.310338Z","iopub.status.idle":"2023-05-17T08:02:49.328303Z","shell.execute_reply.started":"2023-05-17T08:02:49.310306Z","shell.execute_reply":"2023-05-17T08:02:49.326891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"albumentations_train = Compose([\n    VerticalFlip(), HorizontalFlip(), Rotate(limit=30), GridDistortion()\n], p=1)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:03:04.364424Z","iopub.execute_input":"2023-05-17T08:03:04.364844Z","iopub.status.idle":"2023-05-17T08:03:04.370951Z","shell.execute_reply.started":"2023-05-17T08:03:04.364813Z","shell.execute_reply":"2023-05-17T08:03:04.369906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_generator_train = DataGenenerator(train_imgs, augmentation=albumentations_train)\ndata_generator_train_eval = DataGenenerator(train_imgs, shuffle=False)\ndata_generator_val = DataGenenerator(val_imgs, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:03:12.798779Z","iopub.execute_input":"2023-05-17T08:03:12.799237Z","iopub.status.idle":"2023-05-17T08:03:12.818074Z","shell.execute_reply.started":"2023-05-17T08:03:12.799202Z","shell.execute_reply":"2023-05-17T08:03:12.816963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PrAucCallback(Callback):\n    def __init__(self, data_generator, num_workers=num_cores, \n                 early_stopping_patience=5, \n                 plateau_patience=3, reduction_rate=0.5,\n                 stage='train', checkpoints_path='checkpoints/'):\n        super(Callback, self).__init__()\n        self.data_generator = data_generator\n        self.num_workers = num_workers\n        self.class_names = ['Fish', 'Flower', 'Sugar', 'Gravel']\n        self.history = [[] for _ in range(len(self.class_names) + 1)] # to store per each class and also mean PR AUC\n        self.early_stopping_patience = early_stopping_patience\n        self.plateau_patience = plateau_patience\n        self.reduction_rate = reduction_rate\n        self.stage = stage\n        self.best_pr_auc = -float('inf')\n        if not os.path.exists(checkpoints_path):\n            os.makedirs(checkpoints_path)\n        self.checkpoints_path = checkpoints_path\n        \n    def compute_pr_auc(self, y_true, y_pred):\n        pr_auc_mean = 0\n        print(f\"\\n{'#'*30}\\n\")\n        for class_i in range(len(self.class_names)):\n            precision, recall, _ = precision_recall_curve(y_true[:, class_i], y_pred[:, class_i])\n            pr_auc = auc(recall, precision)\n            pr_auc_mean += pr_auc/len(self.class_names)\n            print(f\"PR AUC {self.class_names[class_i]}, {self.stage}: {pr_auc:.3f}\\n\")\n            self.history[class_i].append(pr_auc)        \n        print(f\"\\n{'#'*20}\\n PR AUC mean, {self.stage}: {pr_auc_mean:.3f}\\n{'#'*20}\\n\")\n        self.history[-1].append(pr_auc_mean)\n        return pr_auc_mean\n              \n    def is_patience_lost(self, patience):\n        if len(self.history[-1]) > patience:\n            best_performance = max(self.history[-1][-(patience + 1):-1])\n            return best_performance == self.history[-1][-(patience + 1)] and best_performance >= self.history[-1][-1]    \n              \n    def early_stopping_check(self, pr_auc_mean):\n        if self.is_patience_lost(self.early_stopping_patience):\n            self.model.stop_training = True    \n              \n    def model_checkpoint(self, pr_auc_mean, epoch):\n        if pr_auc_mean > self.best_pr_auc:\n            # remove previous checkpoints to save space\n            for checkpoint in glob.glob(os.path.join(self.checkpoints_path, 'classifier_epoch_*')):\n                os.remove(checkpoint)\n        self.best_pr_auc = pr_auc_mean\n        self.model.save(os.path.join(self.checkpoints_path, f'classifier_epoch_{epoch}_val_pr_auc_{pr_auc_mean}.h5'))              \n        print(f\"\\n{'#'*20}\\nSaved new checkpoint\\n{'#'*20}\\n\")\n              \n    def reduce_lr_on_plateau(self):\n        if self.is_patience_lost(self.plateau_patience):\n            new_lr = float(keras.backend.get_value(self.model.optimizer.lr)) * self.reduction_rate\n            keras.backend.set_value(self.model.optimizer.lr, new_lr)\n            print(f\"\\n{'#'*20}\\nReduced learning rate to {new_lr}.\\n{'#'*20}\\n\")\n        \n    def on_epoch_end(self, epoch, logs={}):\n        y_pred = self.model.predict_generator(self.data_generator, workers=self.num_workers)\n        y_true = self.data_generator.get_labels()\n        # estimate AUC under precision recall curve for each class\n        pr_auc_mean = self.compute_pr_auc(y_true, y_pred)\n              \n        if self.stage == 'val':\n            # early stop after early_stopping_patience=4 epochs of no improvement in mean PR AUC\n            self.early_stopping_check(pr_auc_mean)\n\n            # save a model with the best PR AUC in validation\n            self.model_checkpoint(pr_auc_mean, epoch)\n\n            # reduce learning rate on PR AUC plateau\n            self.reduce_lr_on_plateau()            \n        \n    def get_pr_auc_history(self):\n        return self.history","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:03:20.283696Z","iopub.execute_input":"2023-05-17T08:03:20.284109Z","iopub.status.idle":"2023-05-17T08:03:20.307472Z","shell.execute_reply.started":"2023-05-17T08:03:20.284077Z","shell.execute_reply":"2023-05-17T08:03:20.306451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metric_callback = PrAucCallback(data_generator_train_eval)\nval_callback = PrAucCallback(data_generator_val, stage='val')\n","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:03:28.761996Z","iopub.execute_input":"2023-05-17T08:03:28.762443Z","iopub.status.idle":"2023-05-17T08:03:28.767770Z","shell.execute_reply.started":"2023-05-17T08:03:28.762408Z","shell.execute_reply":"2023-05-17T08:03:28.766723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n    base_model = vit.vit_b16(\n    image_size=224,\n    activation='sigmoid',\n    pretrained=False,\n    include_top=False,\n    pretrained_top=False,\n    classes=4\n)\n    x = base_model.output\n    y_pred = Dense(4, activation='sigmoid')(x)\n    return Model(inputs=base_model.input, outputs=y_pred)\n\nmodel = get_model()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:03:44.474301Z","iopub.execute_input":"2023-05-17T08:03:44.474743Z","iopub.status.idle":"2023-05-17T08:03:47.892748Z","shell.execute_reply.started":"2023-05-17T08:03:44.474710Z","shell.execute_reply":"2023-05-17T08:03:47.891457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for base_layer in model.layers[:-1]:\n    base_layer.trainable = False\n    \nmodel.compile(optimizer=Adam(lr=0.01), loss='binary_crossentropy',metrics=keras.metrics.BinaryAccuracy())\nhistory_0 = model.fit_generator(generator=data_generator_train,\n                              validation_data=data_generator_val,\n                              epochs=10,\n                              callbacks=[train_metric_callback, val_callback],\n                              workers=num_cores,\n                              verbose=1 \n                             )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for base_layer in model.layers[:-1]:\n    base_layer.trainable = True\n    \nmodel.compile(optimizer=Adam(lr=0.005), loss='binary_crossentropy',metrics=keras.metrics.BinaryAccuracy())\nhistory_1 = model.fit_generator(generator=data_generator_train,\n                              validation_data=data_generator_val,\n                              epochs=15,\n                              callbacks=[train_metric_callback, val_callback],\n                              workers=num_cores,\n                              verbose=1 \n                             )","metadata":{"execution":{"iopub.status.busy":"2023-05-17T08:04:01.387825Z","iopub.execute_input":"2023-05-17T08:04:01.388264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_0","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history_1","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_with_dots(ax, np_array):\n    ax.scatter(list(range(1, len(np_array) + 1)), np_array, s=50)\n    ax.plot(list(range(1, len(np_array) + 1)), np_array)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pr_auc_history_train = train_metric_callback.get_pr_auc_history()\npr_auc_history_val = val_callback.get_pr_auc_history()\n\nplt.figure(figsize=(10, 7))\nplt.plot(list(range(1, len(pr_auc_history_train[-1]) + 1)), pr_auc_history_train[-1])\nplt.plot(list(range(1, len(pr_auc_history_val[-1]) + 1)), pr_auc_history_val[-1])\nplt.scatter(list(range(1, len(pr_auc_history_train[-1]) + 1)), pr_auc_history_train[-1])\nplt.scatter(list(range(1, len(pr_auc_history_val[-1]) + 1)), pr_auc_history_val[-1])\nplt.xlabel('Epoch', fontsize=15)\nplt.ylabel('Mean PR AUC', fontsize=15)\nplt.legend(['Train', 'Val'])\nplt.title('Training and Validation PR AUC', fontsize=20)\nplt.savefig('pr_auc_hist.png')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10, 7))\n# plot_with_dots(plt, history_0.history['loss']+history_1.history['loss'])\n# plot_with_dots(plt, history_0.history['val_loss']+history_1.history['val_loss'])\nplt.plot(list(range(1, len(history_0.history['loss']+history_1.history['loss']) + 1)), history_0.history['loss']+history_1.history['loss'])\nplt.plot(list(range(1, len(history_0.history['val_loss']+history_1.history['val_loss']) + 1)), history_0.history['val_loss']+history_1.history['val_loss'])\nplt.scatter(list(range(1, len(history_0.history['loss']+history_1.history['loss']) + 1)), history_0.history['loss']+history_1.history['loss'])\nplt.scatter(list(range(1, len(history_0.history['val_loss']+history_1.history['val_loss']) + 1)), history_0.history['val_loss']+history_1.history['val_loss'])\nplt.xlabel('Epoch', fontsize=15)\nplt.ylabel('Binary Crossentropy', fontsize=15)\nplt.legend(['Train', 'Val'])\nplt.title('Training and Validation Loss', fontsize=20)\nplt.savefig('loss_hist.png')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10, 7))\n# plot_with_dots(plt, history_0.history['loss']+history_1.history['loss'])\n# plot_with_dots(plt, history_0.history['val_loss']+history_1.history['val_loss'])\nplt.plot(list(range(1, len(history_0.history['binary_accuracy']+history_1.history['binary_accuracy']) + 1)), history_0.history['binary_accuracy']+history_1.history['binary_accuracy'])\nplt.plot(list(range(1, len(history_0.history['val_binary_accuracy']+history_1.history['val_binary_accuracy']) + 1)), history_0.history['val_binary_accuracy']+history_1.history['val_binary_accuracy'])\nplt.scatter(list(range(1, len(history_0.history['binary_accuracy']+history_1.history['binary_accuracy']) + 1)), history_0.history['binary_accuracy']+history_1.history['binary_accuracy'])\nplt.scatter(list(range(1, len(history_0.history['val_binary_accuracy']+history_1.history['val_binary_accuracy']) + 1)), history_0.history['val_binary_accuracy']+history_1.history['val_binary_accuracy'])\nplt.xlabel('Epoch', fontsize=15)\nplt.ylabel('accuracy', fontsize=15)\nplt.legend(['Train', 'Val'])\nplt.title('Training and Validation accuracy', fontsize=20)\nplt.savefig('accuracy_hist.png')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = ['Fish', 'Flower', 'Sugar', 'Gravel']\ndef get_threshold_for_recall(y_true, y_pred, class_i, recall_threshold=0.95, precision_threshold=0.94, plot=False):\n    precision, recall, thresholds = precision_recall_curve(y_true[:, class_i], y_pred[:, class_i])\n    i = len(thresholds) - 1\n    best_recall_threshold = None\n    while best_recall_threshold is None:\n        next_threshold = thresholds[i]\n        next_recall = recall[i]\n        if next_recall >= recall_threshold:\n            best_recall_threshold = next_threshold\n        i -= 1\n        \n\n    best_precision_threshold = [thres for prec, thres in zip(precision, thresholds) if prec >= precision_threshold][0]\n    \n    if plot:\n        plt.figure(figsize=(10, 7))\n        plt.step(recall, precision, color='r', alpha=0.3, where='post')\n        plt.fill_between(recall, precision, alpha=0.3, color='r')\n        plt.axhline(y=precision[i + 1])\n        recall_for_prec_thres = [rec for rec, thres in zip(recall, thresholds) \n                                 if thres == best_precision_threshold][0]\n        plt.axvline(x=recall_for_prec_thres, color='g')\n        plt.xlabel('Recall')\n        plt.ylabel('Precision')\n        plt.ylim([0.0, 1.05])\n        plt.xlim([0.0, 1.0])\n        plt.legend(['PR curve', \n                    'PR area',\n                    f'Precision {precision[i + 1]: .2f} corresponding to selected recall threshold',\n                    f'Recall {recall_for_prec_thres: .2f} corresponding to selected precision threshold'])\n        plt.title(f'Precision-Recall curve for Class {class_names[class_i]}')\n    return best_recall_threshold, best_precision_threshold\n\ny_pred = model.predict_generator(data_generator_val, workers=num_cores)\ny_true = data_generator_val.get_labels()\nrecall_thresholds = dict()\nprecision_thresholds = dict()\nfor i, class_name in tqdm(enumerate(class_names)):\n    recall_thresholds[class_name], precision_thresholds[class_name] = get_threshold_for_recall(y_true, y_pred, i, plot=True)","metadata":{},"execution_count":null,"outputs":[]}]}