{"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":"# Few-shot learning\n\n![Few shot cover](https://i.ibb.co/W5JqJqR/few-shot-cover.png)\n[image credit](https://medium.com/sap-machine-learning-research/deep-few-shot-learning-a1caa289f18)\n\nFew-shot learning is the problem of making predictions based on a limited number of samples. Few-shot learning is different from classical supervised learning. The goal of few-shot learning is not to let the model recognize the images in the training set and then generalize to the test set. Instead, the goal is to learn. “Learn to learn” sounds hard to understand. You can think of it in this way.\n\n## Few-shot classification task\n\nThere are a lot of techniques of few-shot learning. One of the most popular - using support set. Let’s define an N-way-K-Shot classification problem. In that setup, the support set contains samples of N-classes. So, during forward pass such model could classify only 1 class of N. There are a lot of papers and benchmarks tackling that problem - [papers with code page](https://paperswithcode.com/task/few-shot-image-classification).\n\nIt is natural that an increasing number of classes (**N**) will decrease the overall accuracy of our model. So, that approach leads to pure results while working with a huge amount of classes (like [MiniImage](https://paperswithcode.com/sota/few-shot-image-classification-on-mini-1)).\n\n# This work \n\nIn this work, we are going to build a model that will classify images along with **1000 classes** where only **10 images per class** are available (including validation!).\n\nSince we have a limited amount of data per class, traditional classification will lead to poor results. Today we will try to turn the image classification task into metric learning.\n\n## What is metric learning\n\n![Distance metric img](https://www.researchgate.net/profile/Wenyu-Liu-7/publication/221361643/figure/fig1/AS:643202524667904@1530362840212/Adaptive-distance-metric-learning-on-synthetic-data-The-left-figure-indicates-the.png)\n[image credit](http://contrib.scikit-learn.org/metric-learn/introduction.html)\n\nMany approaches in machine learning require a measure of distance between data points. Traditionally, practitioners would choose a standard distance metric (Euclidean, City-Block, Cosine, etc.) using a priori knowledge of the domain. However, it is often difficult to design metrics that are well-suited to the particular data and task of interest.\n\nDistance metric learning (or simply, metric learning) aims at automatically constructing task-specific distance metrics from (weakly) supervised data, in a machine learning manner.\n\nUsually, a deep neural network is taken as a feature extractor. Then standard similarity functions are used (e.g. L2 distance, COS distance, etc).\n\n#### Predict if images are similar\n\nHere is a diagram of the model which can check if two given samples are similar or not:\n\n![Similairty estimation](https://i.ibb.co/C7XtwnB/similarity-1.png)\n\nIn the diagram above we have two modules - a trainable feature extractor and a trainable similarity comparator. Feature extractor builds descriptor of the input image in such a way, that similar images have the similar descriptors (in terms of similarity comparator). Images of different objects would have descriptors whose distance is much bigger.\n\nSo, to check does the input images belongs to the same class we just extract descriptors from both samples, calculate distance and compare that distance with come predefined threshold (calculated based on the dataset). This is a very basic approach. Modern papers use advanced techniques (for example, custom threshold per each class[[1](https://arxiv.org/pdf/1810.11160v1.pdf)])\n\n#### Predict image class\n\nNow, we can move toward the classification task. For that, we just replace the similarity measurement block with the simple K-Nearest Neighbors classifier. According to the assumption before, descriptors of the same class are located close to each other. So, KNN is a good choice to choose the right prediction based on the dataset.\n\n![Few-shot classification](https://i.ibb.co/ZzW3YTM/knn-1.png)\n\n## Dataset description\n\n![Dataset cover](https://i.ibb.co/w4sNJ8j/cover-1.jpg)\n\nHere we are going to use [FSS-1000 dataset](https://www.kaggle.com/meowmeowmeowmeowmeow/fss1000-a-1000-class-fewshot-segmentation). This dataset consists of 1000 classes of various objects (like a tennis racket, polar bear, pizza, pickup, and so on). Dataset was naturally created for a few-shot segmentation task. Here we are going to turn it into a classification task.","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nimport tensorflow_addons as tfa\nimport tensorflow_hub as hub\nimport pandas as pd\nimport os\nimport cv2\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport numpy as np\nimport itertools as it\nfrom sklearn.neighbors import KNeighborsClassifier\nimport seaborn as sns\nfrom skimage import filters\nfrom sklearn.neighbors import KNeighborsClassifier\nfrom sklearn.metrics import accuracy_score, roc_curve, auc\nfrom sklearn.neural_network import MLPClassifier\nfrom sklearn.model_selection import GridSearchCV\nfrom warnings import simplefilter\nfrom sklearn.exceptions import ConvergenceWarning\nfrom tensorflow.keras.metrics import top_k_categorical_accuracy\nfrom collections import namedtuple\nfrom sklearn.model_selection import GridSearchCV, cross_val_score, train_test_split\nimport sklearn\n\nsimplefilter(\"ignore\", category=ConvergenceWarning)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:17:56.350175Z","iopub.execute_input":"2022-02-01T20:17:56.350935Z","iopub.status.idle":"2022-02-01T20:17:56.35908Z","shell.execute_reply.started":"2022-02-01T20:17:56.350887Z","shell.execute_reply":"2022-02-01T20:17:56.358271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# FSS-1000 dataset\n\nFSS-1000 consists of 1000 object classes with pixelwise annotation of ground-truth segmentation. Unique in FSS-1000, our dataset contains a significant number of objects that have never been seen or annotated in previous datasets, such as tiny daily objects, merchandise, cartoon characters, logos, etc.\nOnly one class is presented per one image.\n\nLet's load annotation of the dataset.","metadata":{}},{"cell_type":"code","source":"csv = pd.read_csv('/kaggle/input/fss1000-a-1000-class-fewshot-segmentation/fss-1000.csv', index_col=0)\ncsv.sample(3)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:17:56.360794Z","iopub.execute_input":"2022-02-01T20:17:56.361399Z","iopub.status.idle":"2022-02-01T20:17:56.416403Z","shell.execute_reply.started":"2022-02-01T20:17:56.361361Z","shell.execute_reply":"2022-02-01T20:17:56.415647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let me briefly describe DataFrame appeared above:\n\n - *in_file* - Relative path to the RGB image\n - *out_file* - Relative path to the grayscale image which represents mask of object presented on the RGB image\n - *class* - String representation of the object class which is presented on the RGB image\n - *x_min*, *x_max*, *y_min*, *y_max* - Relative coordinates of bounding box around object mask\n - *width*, *height* - Relative width and height of bounding box\n - *class_id* - Unique ID of the object class\n\nHere we are going to focus on the classification problem, so we will discard mask data. Instead, we will use a precalculated bounding box right out of DataFrame above.\nFirst of all, we will adjust paths to images. It will help us to load images in further steps without any attention to paths.","metadata":{}},{"cell_type":"code","source":"csv['in_file'] = csv['in_file'].map(lambda x: os.path.join('/kaggle/input/fss1000-a-1000-class-fewshot-segmentation/FSS-1000/', x))\ncsv['out_file'] = csv['out_file'].map(lambda x: os.path.join('/kaggle/input/fss1000-a-1000-class-fewshot-segmentation/FSS-1000/', x))\n\n# Assert that all files are available\nassert csv['in_file'].map(os.path.isfile).all()\nassert csv['out_file'].map(os.path.isfile).all()\n\ncsv.sample(3)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:17:56.418044Z","iopub.execute_input":"2022-02-01T20:17:56.418481Z","iopub.status.idle":"2022-02-01T20:18:18.442894Z","shell.execute_reply.started":"2022-02-01T20:17:56.418437Z","shell.execute_reply":"2022-02-01T20:18:18.442208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"csv.describe()","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:18:18.444241Z","iopub.execute_input":"2022-02-01T20:18:18.444704Z","iopub.status.idle":"2022-02-01T20:18:18.477821Z","shell.execute_reply.started":"2022-02-01T20:18:18.444667Z","shell.execute_reply":"2022-02-01T20:18:18.477196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Minimalistic exploration of distribution proves that images are uniformly distributed over 1000 classes.\n\nMost of the time we will work with class IDs. It would be very helpful to have some util that will easly convert class ID to string representation of the class name.","metadata":{}},{"cell_type":"markdown","source":"# Basic data exploration","metadata":{}},{"cell_type":"markdown","source":"## Class balance exploration","metadata":{}},{"cell_type":"code","source":"class_grouped = csv.groupby('class_id')\n\nsamples_per_class = class_grouped.count().in_file.to_numpy()\nclasses_id = class_grouped.count().index\n\nassert np.all(samples_per_class == 10) # FSS-1000 ensure us that there are exectly 10 samples per class","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:18:18.479702Z","iopub.execute_input":"2022-02-01T20:18:18.479949Z","iopub.status.idle":"2022-02-01T20:18:18.496441Z","shell.execute_reply.started":"2022-02-01T20:18:18.479915Z","shell.execute_reply":"2022-02-01T20:18:18.495837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Raw image exploration\n\nHere is a visualization of raw images. Since we are working with classification tasks, we are interested in bounding box areas only.","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=2, figsize=(24, 8))\n\nfor _ax, (idx, row) in zip(ax.ravel(), csv.sample(10).iterrows()):\n    img = cv2.imread(row.in_file)\n    bb = row.x_min, row.x_max, row.y_min, row.y_max\n    bb = [x*224 for x in bb]\n    rect = patches.Rectangle((bb[0], bb[2]), bb[1]-bb[0], bb[3]-bb[2], linewidth=2, edgecolor='r', facecolor='none')\n    \n    _ax.imshow(img[..., ::-1])\n    _ax.add_patch(rect)\n    \n    _ax.set_title(row['class'])\n    \n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:18:18.497586Z","iopub.execute_input":"2022-02-01T20:18:18.498033Z","iopub.status.idle":"2022-02-01T20:18:19.535101Z","shell.execute_reply.started":"2022-02-01T20:18:18.497999Z","shell.execute_reply":"2022-02-01T20:18:19.530722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cropped objects visualization","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=2, figsize=(24, 8))\n\nfor _ax, (idx, row) in zip(ax.ravel(), csv.sample(10).iterrows()):\n    img = cv2.imread(row.in_file)\n    bb = row.x_min, row.x_max, row.y_min, row.y_max\n    bb = [int(x*224) for x in bb]\n    \n    cropped_img = img[bb[2]:bb[3], bb[0]:bb[1], :]\n    \n    _ax.imshow(cropped_img[..., ::-1])\n    \n    _ax.set_title(row['class'])\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:18:19.536579Z","iopub.execute_input":"2022-02-01T20:18:19.537068Z","iopub.status.idle":"2022-02-01T20:18:20.427985Z","shell.execute_reply.started":"2022-02-01T20:18:19.537029Z","shell.execute_reply":"2022-02-01T20:18:20.42739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Cropped images have different size. We need to resize it in order to form batches. What size is the best tradeoff between quality and size?","metadata":{}},{"cell_type":"code","source":"#widths = csv.apply(lambda x: x.width*224, axis=1)\n#heights = csv.apply(lambda x: x.height*224, axis=1)\n\n#sns.jointplot(widths, heights, kind='kde')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:48.710726Z","iopub.execute_input":"2022-02-01T20:27:48.71158Z","iopub.status.idle":"2022-02-01T20:27:48.71709Z","shell.execute_reply.started":"2022-02-01T20:27:48.711532Z","shell.execute_reply":"2022-02-01T20:27:48.71629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Looks like most cropped images have a size of 175x175.","metadata":{}},{"cell_type":"code","source":"def class_id_to_str(cls_id):\n    return csv[csv.class_id == cls_id].iloc[0]['class']","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:50.35107Z","iopub.execute_input":"2022-02-01T20:27:50.351322Z","iopub.status.idle":"2022-02-01T20:27:50.355072Z","shell.execute_reply.started":"2022-02-01T20:27:50.351294Z","shell.execute_reply":"2022-02-01T20:27:50.354387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train & validation & test split\n\nNow it's time to split the whole dataset into the train and validation splits. Taking into account an extremely small amount of sample per class, a random split over the whole CSV is not the best approach.\n\nLet's split samples randomly per class. In that approach, we will definitely obtain uniformly balanced train and validation splits.","metadata":{}},{"cell_type":"code","source":"csv_train = []\ncsv_val = []\ncsv_test = []\n\nclasses = np.unique(csv.class_id)\nnp.random.shuffle(classes)\n\nfor class_id in classes[:800]:\n    csv_sliced = csv[csv.class_id == class_id].sample(frac=1)  # Slice DataFrame for one class only. shuffle it\n    csv_train.append(csv_sliced.iloc[:7])  # Take first 7 samples to train\n    csv_val.append(csv_sliced.iloc[7:])  #  Take some to validation\n    \nfor class_id in classes[800:]:\n    csv_sliced = csv[csv.class_id == class_id].sample(frac=1)\n    csv_test.append(csv_sliced)\n    \n# Concatenate per class splits\ncsv_train = pd.concat(csv_train, axis=0).sample(frac=1)\ncsv_val = pd.concat(csv_val, axis=0).sample(frac=1)\ncsv_test = pd.concat(csv_test, axis=0).sample(frac=1)\n\n# Lightway check that we have no data leakage\nfor train_file in csv_train.in_file:\n    for val_file in csv_val.in_file:\n        assert not train_file == val_file\n        \nfor train_file in csv_train.in_file:\n    for test_file in csv_test.in_file:\n        assert not train_file == test_file\n        \nfor val_file in csv_val.in_file:\n    for test_file in csv_test.in_file:\n        assert not val_file == test_file","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:52.258355Z","iopub.execute_input":"2022-02-01T20:27:52.258895Z","iopub.status.idle":"2022-02-01T20:27:58.572343Z","shell.execute_reply.started":"2022-02-01T20:27:52.258855Z","shell.execute_reply":"2022-02-01T20:27:58.571599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preapre dataset loader\n\nAt that point, we have two DataFrame - train and validation. Since we are going to work with DL models we need to have an efficient data pipeline. Loading data on the fly allows us to avoid bottlenecks and speed up the overall training routine.\n\nTensorFlow efficient data pipeline seems to be the best solution (at least for TF :)).\nTensorFlow data pipelines work like a built-in python generator. We are going to work with images. Obviously, it is not a solution to load the whole dataset into the memory. So, we just load paths to images. Then, map function, that loaded image, crop it. After that, we may map the addition data preprocessing step (normalization, augmentation, etc). In the end, we pack samples into batches. **Note,** no image has been loaded jet.\nThe image will be loaded when you start iterating over the dataset. It is very efficient, that TensorFlow handles all routines about multiprocessing, paralleling. You don't pay attention to it at all. GPU trains your model while the CPU loads the next batches to it. No bottlenecks there.\n\nFollow that link for a detailed tutorial about [TensorFlow data pipelines](https://www.tensorflow.org/guide/data).\n\nLet's define some basic utils to help us work with DataFrames defined above.","metadata":{}},{"cell_type":"code","source":"# Extracts bounding box from DataFrame row\n# row: DataFrame row\ndef bb_from_row(row):\n    return row.x_min, row.x_max, row.y_min, row.y_max\n\n# Load imabe from path and crop it according to boundig box\n# file: string Tensor ,path to the image\n# bb: relative boundig box coordinates - minx, maxx, miny, maxy\n# label: class label of sample\ndef crop_image(file, bb, label):\n    img_data = tf.io.read_file(file)\n    img = tf.io.decode_jpeg(img_data, channels=3)\n    \n    bb *= 224\n    bb = tf.cast(bb, tf.int32)\n    \n#     img = img[bb[2]:bb[3], bb[0]:bb[1], :]\n    img = tf.image.resize(img, [224, 224])\n    img /= 255\n    \n    return img, label","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:28:02.875297Z","iopub.execute_input":"2022-02-01T20:28:02.875549Z","iopub.status.idle":"2022-02-01T20:28:02.882423Z","shell.execute_reply.started":"2022-02-01T20:28:02.87552Z","shell.execute_reply":"2022-02-01T20:28:02.881496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Generator that yields image path, bounding box and class label\n# _csv: DataFrame to yield samples from\ndef file_bb_lbl_gen(_csv: pd.DataFrame):\n    def generator():\n        for idx, row in _csv.iterrows():\n            file = row.in_file\n\n            bb = bb_from_row(row)\n            label = row.class_id\n\n            yield file, bb, label\n    return generator\n\n\n# Makes TensorFlow datasets from a DataFrame\n# _csv: DataFrame to make dataset from\ndef make_img_lbl_ds(_csv):\n    ds = tf.data.Dataset.from_generator(file_bb_lbl_gen(_csv), output_types=(tf.string, tf.float32, tf.int32))\n    # Now ds yields tuples of (file_path, bounding_box, label)\n    ds = ds.map(crop_image) # Maps image loading and cropping to the dataset\n    # Now ds yields tuples of (image, label)\n    \n    return ds\n\n\n# Perofrm augmentation to the sample\ndef aug(img, label):\n    x = tf.image.random_flip_left_right(img)\n    return x, label","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:28:04.615613Z","iopub.execute_input":"2022-02-01T20:28:04.61617Z","iopub.status.idle":"2022-02-01T20:28:04.622863Z","shell.execute_reply.started":"2022-02-01T20:28:04.616133Z","shell.execute_reply":"2022-02-01T20:28:04.621958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"You may ask why there are redundant arguments in preprocessing utils (e.g. label)? The answer is simple. Since we form samples (which consists of image and label). A function that will be mapped to such dataset should have signature which reflects dataset sample - image and label in our case.\n\nNow, let's create dataset objects","metadata":{}},{"cell_type":"code","source":"# We keep original datasets without batches. Thet's why names are started from undescroe\n_train_dataset, _val_dataset, _test_dataset = make_img_lbl_ds(csv_train), make_img_lbl_ds(csv_val), make_img_lbl_ds(csv_test)\n\n# Build your input pipelines. Here we perform shufling, batching and prefetching\ntrain_dataset = _train_dataset.map(aug).shuffle(1024).batch(256).prefetch(2)\ntrain_dataset_no_aug = _train_dataset.batch(256).prefetch(2)\n\nval_dataset = _val_dataset.batch(256).prefetch(4)","metadata":{"id":"iXvByj6wcT7d","execution":{"iopub.status.busy":"2022-02-01T20:28:07.082235Z","iopub.execute_input":"2022-02-01T20:28:07.082656Z","iopub.status.idle":"2022-02-01T20:28:09.524674Z","shell.execute_reply.started":"2022-02-01T20:28:07.082605Z","shell.execute_reply":"2022-02-01T20:28:09.523966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize samples from dataset pipeline","metadata":{}},{"cell_type":"markdown","source":"We need carefully inspect data at each step. Let's unsure that data from TensorFlow pipelines looks good.","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=2, figsize=(24, 8))\n\nfor _ax, (img, label) in zip(ax.ravel(), _train_dataset.shuffle(444).take(10)):\n    _ax.imshow(img)\n    _ax.set_title(class_id_to_str(label.numpy()))\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:28:12.214657Z","iopub.execute_input":"2022-02-01T20:28:12.215211Z","iopub.status.idle":"2022-02-01T20:28:15.733321Z","shell.execute_reply.started":"2022-02-01T20:28:12.215164Z","shell.execute_reply":"2022-02-01T20:28:15.732692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build model\n\n![TF Hub](https://miro.medium.com/max/688/0*B8VDCnh8qBnwUuM4)\n\n### Tensorflow Hub\nIt is a good idea to start with a pretrained model on the imagenet. We could try to export such a model from TensorFlow Hub. TensorFlow Hub provides us with various pretrained models which may be reused in one line of code. Such a model is imported as a Keras layer. Most of the models are trainable, which means you are able to fine-tune imported weights. Despite the opportunity to train the imported layer, we cannot access internal blocks. In other words, the imported model is monolithic, and no external logic may be attached to it except using input and output. It is the biggest drawback of TensorFlow Hub. For example, you could reuse the classification model, but it cannot be reused for segmentation tasks thus since we cannot cut off the last max-pooling layer. However, despite drawbacks, it is very convenient to use TensorFlow Hub for most of the daily research tasks.\n\n### Model architecture\n\n![Embading model](./img/embading-model.png)\n\nWe choose EfficientNet-B2 as a feature extractor. One fully connected layer is used to form embeddings. Note, that number of neurons in those layers controls the size of embeddings. Finally, sample-wise L2 normalization is used in order to control the embedding scale.\n","metadata":{}},{"cell_type":"code","source":"model = tf.keras.Sequential([\n    hub.KerasLayer(\"https://tfhub.dev/google/imagenet/efficientnet_v2_imagenet21k_b2/feature_vector/2\", trainable=False),\n    tf.keras.layers.Flatten(),\n    tf.keras.layers.Dense(256, activation=None),\n    tf.keras.layers.Lambda(lambda x: tf.math.l2_normalize(x, axis=1)) # L2 normalize embeddings\n\n])","metadata":{"id":"djpoAvfWNyL5","execution":{"iopub.status.busy":"2022-02-01T20:28:21.406472Z","iopub.execute_input":"2022-02-01T20:28:21.407036Z","iopub.status.idle":"2022-02-01T20:28:29.388516Z","shell.execute_reply.started":"2022-02-01T20:28:21.406996Z","shell.execute_reply":"2022-02-01T20:28:29.387695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train and Evaluate\n\nWe need a special loss function in order to force the model to learn meaningful embeddings. \nThe goal of the triplet loss[[2]](https://arxiv.org/abs/1503.03832) is to make sure that:\n* Two examples with the same label have their embeddings close together in the embedding space\n* Two examples with different labels have their embeddings far away.\n\nSo, triplet loss takes three embeddings - *A*, *P* and *N*\n* A - random embedding sample\n* P - embedding that represents the same class as an *A*\n* N - embedding that represents other class\n\n\n![Triplet loss](https://omoindrot.github.io/assets/triplet_loss/triplet_loss.png) [image credit](https://omoindrot.github.io/triplet-loss)\n    \nTriplet loss formula:\n    \n$$\\Large\\lambda = max(d(A, P) - d(A, N) + m, 0)$$\n    \nMargin *m* here is used in order to control desired inter-class embedding distances.\n    \nUnfortunately, random triplet mining (i.e. random (A, P, N)) leads to a poor result. Deep models tend to overfit training data (especially if training on small datasets). Using triplet loss allows us to increase the number of samples exponentially (by incorporating three samples there is a lot of different combination). Despite a large number of different triplets, there is still a chance to overfit the model. That is because of model optimizes all distances between samples despite having very good distances between easy samples. In order to prevent such behavior, a triplet mining policy is implemented. \n   \nBased on the definition of the loss, there are three categories of triplets:\n\n* Easy triplets - triplets which have a loss of 0, because $d(A, P) + m < d(A, N)$\n* Hard triplets - triplets where the negative is closer to the anchor than the positive, i.e. $d(A, N) < d(A, P)$\n* Semi-hard triplets -  triplets where the negative is not closer to the anchor than the positive, but which still have positive loss: $d(A, P) < d(A, N) < d(A, P) + m$\n\n<img style=\"text-align: centr\" src=\"https://omoindrot.github.io/assets/triplet_loss/triplets.png\" width=500>\n    \n[image credit](https://omoindrot.github.io/triplet-loss)\n    \nTake a look at [Triplet Loss and Online Triplet Mining in TensorFlow](https://omoindrot.github.io/triplet-loss) article if you are interested in more details.\n    \n[TensorFlow addons](https://www.tensorflow.org/addons) provides additional utils for TensorFlow. So, it is a good option to use triplet loss from that package. It provides us with different triplet mining policies.","metadata":{"id":"HYE-BxhOzFQp"}},{"cell_type":"code","source":"# Compile the model\nmodel.compile(\n    optimizer=tf.keras.optimizers.Adam(0.001),\n    loss=tfa.losses.TripletSemiHardLoss())","metadata":{"id":"NxfYhtiSzHf-","execution":{"iopub.status.busy":"2022-02-01T20:28:29.391189Z","iopub.execute_input":"2022-02-01T20:28:29.391741Z","iopub.status.idle":"2022-02-01T20:28:29.407928Z","shell.execute_reply.started":"2022-02-01T20:28:29.391699Z","shell.execute_reply":"2022-02-01T20:28:29.407144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train model\n\nNow, it's time to train our model. Since using, a pretrained backbone model it is good practice to freeze the backbone fist to prevent fresh gradients broke the pretrained weights.\nSo, our training procedure consists of two stages:\n * Train only FC layers\n * Unfreeze backbone and fine-tune the whole model","metadata":{}},{"cell_type":"code","source":"# Train the network\nsave_cb = tf.keras.callbacks.ModelCheckpoint('./best_val', save_best_only=True, monitor='val_loss', save_weights_only=True)\nhistory1 = model.fit(train_dataset, epochs=20, validation_data=val_dataset, verbose=2, callbacks=[save_cb]) \n\nmodel.load_weights('./best_val')\nmodel.layers[0].trainable = True\n\nhistory2 = model.fit(train_dataset, epochs=20, validation_data=val_dataset, verbose=2, callbacks=[save_cb]) ","metadata":{"id":"TGBYNGxgVDrj","execution":{"iopub.status.busy":"2022-02-01T20:28:37.146425Z","iopub.execute_input":"2022-02-01T20:28:37.146704Z","iopub.status.idle":"2022-02-01T20:29:45.575859Z","shell.execute_reply.started":"2022-02-01T20:28:37.146672Z","shell.execute_reply":"2022-02-01T20:29:45.573188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training & validation loss\n\nLet's explore training and validation losses","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize=(24, 4))\n\nax.plot(np.arange(1, 21, 1), history1.history['loss'], label='Frozen backbone (train)', c='tab:blue', linestyle='--')\nax.plot(np.arange(21, 41, 1), history2.history['loss'], c='tab:red', label='Fine tune backbone (train)', linestyle='--')\n\nax.plot(np.arange(1, 21, 1), history1.history['val_loss'], label='Frozen backbone (validation)', c='tab:blue')\nax.plot(np.arange(21, 41, 1), history2.history['val_loss'], c='tab:red', label='Fine tune backbone (validation)')\n\nax.grid()\nax.legend()\nax.set_title('Training loss graph')\nax.set\n\nax.set_ylabel('Loss', fontsize=16)\nax.set_xlabel('Epoch', fontsize=16);","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.233083Z","iopub.status.idle":"2022-02-01T20:27:35.233625Z","shell.execute_reply.started":"2022-02-01T20:27:35.233387Z","shell.execute_reply":"2022-02-01T20:27:35.233411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluate model\n\nSince we have saves the best model checkpoint it makes sence to perform evaluation using that checkpoints","metadata":{}},{"cell_type":"code","source":"model.load_weights('./best_val')\n\ntrain_loss = model.evaluate(train_dataset_no_aug)\nval_loss = model.evaluate(val_dataset)\n\nprint(f'Loss | Train: {train_loss:0.4f} \\t Validation: {val_loss:0.4f}')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.234703Z","iopub.status.idle":"2022-02-01T20:27:35.235246Z","shell.execute_reply.started":"2022-02-01T20:27:35.235016Z","shell.execute_reply":"2022-02-01T20:27:35.235042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Explore embeddings\n\nHaving a single loss value doesn't tell us much about model performance. It would be useful to estimate distances between inter/intra-class embeddings. Let's do that.\nThe function above calculates embeddings for each batch of images. Afterward, it calculates distances for each pair from the given batch. We can't calculate distances between each sample pair due to computational resources. But sampling uniformly distributed samples give us a good estimation of distance distribution.\nOnce we have a distance between inter/intra-class embeddings it is worth finding the best distance threshold which separates similar and opposite embeddings. The best option is to use Otsu thresholding.","metadata":{}},{"cell_type":"code","source":"def get_embeddings_and_labels(dataset):\n    embaddings = model.predict(dataset)\n    labels = np.concatenate([y.numpy() for x, y in dataset])\n    assert embaddings.shape[0] == labels.shape[0]\n    return embaddings, labels\n\nval_embeddings, val_labels = get_embeddings_and_labels(val_dataset)\ntrain_embeddings, train_labels = get_embeddings_and_labels(train_dataset_no_aug)\n\n\ntest_dataset = _test_dataset.batch(256).prefetch(4)\ntest_embeddings, test_labels = get_embeddings_and_labels(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.236306Z","iopub.status.idle":"2022-02-01T20:27:35.236869Z","shell.execute_reply.started":"2022-02-01T20:27:35.236632Z","shell.execute_reply":"2022-02-01T20:27:35.23667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm\n\ndef embading_distances(embeddings, labels, samples):\n    dist = tf.keras.metrics.MSE\n    \n    all_classes = np.unique(labels)\n\n    pos_dists = []\n    for i in tqdm(range(samples)):\n#         idx = int(np.random.uniform(0, 1000))\n        idx = np.random.choice(all_classes)\n        possible_idxs = np.where(labels == idx)[0]\n        choose_idx = np.random.choice(possible_idxs, size=2)\n\n        d = dist(embeddings[choose_idx[0]], embeddings[choose_idx[1]])\n        pos_dists.append(d)\n\n    neg_dists = []\n    for i in tqdm(range(samples)):\n#         idx = int(np.random.uniform(0, 1000))\n        idx = np.random.choice(all_classes)\n        other_idxs = np.where(labels != idx)[0]\n        choose_idx = np.random.choice(other_idxs, size=1)\n\n        d = dist(embeddings[idx], embeddings[choose_idx[0]])\n        neg_dists.append(d)\n        \n    pos_dists, neg_dists = sklearn.utils.shuffle(pos_dists, neg_dists)\n    return pos_dists, neg_dists\n\n\ntrain_pos_dists, train_neg_dists = embading_distances(train_embeddings, train_labels, samples=int(2*10e3))\nval_pos_dists, val_neg_dists = embading_distances(val_embeddings, val_labels, samples=int(2*10e3))\ntest_pos_dists, test_neg_dists = embading_distances(test_embeddings, test_labels, samples=int(10e3))","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.237945Z","iopub.status.idle":"2022-02-01T20:27:35.238495Z","shell.execute_reply.started":"2022-02-01T20:27:35.238257Z","shell.execute_reply":"2022-02-01T20:27:35.238281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_thresh = filters.threshold_otsu(np.concatenate([val_pos_dists, val_neg_dists]))\ntrain_thresh = filters.threshold_otsu(np.concatenate([train_pos_dists, train_neg_dists]))\n\nprint(f'Threshold for train set: {train_thresh:0.5f} Threshold for test set: {val_thresh:0.5f} \\t Validation threshold {(np.abs(train_thresh - val_thresh) / train_thresh)*100:0.2f}%')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.239564Z","iopub.status.idle":"2022-02-01T20:27:35.240117Z","shell.execute_reply.started":"2022-02-01T20:27:35.239888Z","shell.execute_reply":"2022-02-01T20:27:35.239911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize embedding distances\n\nThe best way to visualize data distribution - use a histogram. Let's compare histograms of train and validation distance distribution. We also compare thresholds from training and validation sets. It indicates that our model did not face overfit since thresholds are quite similar.","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize=(24, 4))\n\nsns.distplot(train_pos_dists, label='Positive distances (train)')\nsns.distplot(train_neg_dists, label='Negative distances (train)')\n\nax.vlines(train_thresh, 0, 1100, color='r', linestyle='--', label='Threshold')\nax.set_title('Validation distance distribution')\n\nax.set_xlabel('L2 distance', fontsize=16)\nax.legend();\n\nplt.show()\n\n\nfig, ax = plt.subplots(figsize=(24, 4))\n\nsns.distplot(val_pos_dists, label='Positive distances (validation)')\nsns.distplot(val_neg_dists, label='Negative distances (validation)')\n\nax.vlines(train_thresh, 0, 2000, color='r', linestyle='--', label='Threshold')\nax.vlines(val_thresh, 0, 2000, color='g', linestyle='--', label='Perfect threshold for validation set')\n\nax.set_xlabel('L2 distance', fontsize=16)\nax.legend()\nax.set_title('Validation distance distribution')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.241354Z","iopub.status.idle":"2022-02-01T20:27:35.241927Z","shell.execute_reply.started":"2022-02-01T20:27:35.241698Z","shell.execute_reply":"2022-02-01T20:27:35.241721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Binary accuracy\n\n## Otsu thresholding\n\nHaving the thresholds described above we can predict if the given embedding belongs to the same class or not. Let's calculate scores for such a task","metadata":{}},{"cell_type":"code","source":"def threshold_classifier(thresh):\n    def classify(x):\n        y = np.where(x < thresh, 1, 0)\n        return y\n    return classify\n\n\nBinClassifierDesc = namedtuple('BinClassifierDesc', ['acc', 'fpr', 'tpr', 'auc'])\n\nthresh_clf = threshold_classifier(train_thresh)\n\ntrain_dists = np.concatenate([train_pos_dists, train_neg_dists])\nval_dists = np.concatenate([val_pos_dists, val_neg_dists])\ntest_dists = np.concatenate([test_pos_dists, test_neg_dists])\n\ntrain_dist_gt_labels = np.concatenate([np.full_like(train_pos_dists, 1), np.full_like(train_neg_dists, 0)])\nval_dist_gt_labels = np.concatenate([np.full_like(val_pos_dists, 1), np.full_like(val_neg_dists, 0)])\ntest_dist_gt_labels = np.concatenate([np.full_like(test_pos_dists, 1), np.full_like(test_neg_dists, 0)])\n\ndef thresh_classifier_eval(dists, lbls):\n    y_pred = thresh_clf(dists)\n    acc = accuracy_score(lbls, y_pred)\n    return acc\n    \nthresh_train_score = thresh_classifier_eval(train_dists, train_dist_gt_labels)\nthresh_val_score = thresh_classifier_eval(val_dists, val_dist_gt_labels)\nthresh_test_score = thresh_classifier_eval(test_dists, test_dist_gt_labels)\n\nfig, ax = plt.subplots(figsize=(12, 4))\n\nsns.barplot([thresh_train_score, thresh_val_score, thresh_test_score], ['Train', 'Validation', 'Test'])\nax.set_xlabel('Accuracy score', fontsize=16)\nax.set_ylabel('Subset', fontsize=16)\nfig.suptitle(f'Threshold classifier score | Train: {thresh_train_score:0.3f} Validation: {thresh_val_score:0.3f} Test: {thresh_test_score:0.3f}', fontsize=16)\nax.grid();","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.243003Z","iopub.status.idle":"2022-02-01T20:27:35.243562Z","shell.execute_reply.started":"2022-02-01T20:27:35.243317Z","shell.execute_reply":"2022-02-01T20:27:35.24334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Distance classifier","metadata":{}},{"cell_type":"code","source":"thresh_clf_finder = GridSearchCV(MLPClassifier(), param_grid={\n    'hidden_layer_sizes': [(2,), (4,), (10,), (12,)],\n    'activation': ['logistic', 'tanh', 'relu'],\n    'solver': ['lbfgs', 'sgd', 'adam'],\n    'learning_rate': ['constant', 'invscaling', 'adaptive']\n}, n_jobs=-1)\n\n_train_dists = np.expand_dims(train_dists, axis=-1)\n\nthresh_clf_finder.fit(_train_dists, train_dist_gt_labels)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.244631Z","iopub.status.idle":"2022-02-01T20:27:35.245246Z","shell.execute_reply.started":"2022-02-01T20:27:35.244955Z","shell.execute_reply":"2022-02-01T20:27:35.245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_val_dists = np.expand_dims(val_dists, axis=-1)\n_test_dists = np.expand_dims(test_dists, axis=-1)\n\nBinClassifierDesc = namedtuple('BinClassifierDesc', ['acc', 'fpr', 'tpr', 'auc'])\n\ndef eval_bin_classifier(dists, lbls):\n    _dists = np.expand_dims(dists, axis=-1)\n    y_pred = thresh_clf_finder.best_estimator_.predict(_dists)\n    fpr, tpr, _ = roc_curve(lbls, y_pred)\n    auc_score = auc(fpr, tpr)\n    score = thresh_clf_finder.best_estimator_.score(_dists, lbls)\n    return BinClassifierDesc(score, fpr, tpr, auc_score)\n\ntrain_eval_bin_classifier = eval_bin_classifier(train_dists, train_dist_gt_labels)\nval_eval_bin_classifier = eval_bin_classifier(val_dists, val_dist_gt_labels)\ntest_eval_bin_classifier = eval_bin_classifier(test_dists, test_dist_gt_labels)\n\n\nfig, ax = plt.subplots(figsize=(8, 8))\n\nax.plot(train_eval_bin_classifier.fpr, train_eval_bin_classifier.tpr, label=f'D-Classifier ROC (train) | AUC {train_eval_bin_classifier.auc: 0.3f}')\nax.plot(val_eval_bin_classifier.fpr, val_eval_bin_classifier.tpr, label=f'D-Classifier ROC (validation) | AUC {val_eval_bin_classifier.auc: 0.3f}')\nax.plot(test_eval_bin_classifier.fpr, test_eval_bin_classifier.tpr, label=f'D-Classifier ROC (test) | AUC {test_eval_bin_classifier.auc: 0.3f}')\nax.legend()\nfig.suptitle(f'Distance classifier score | Train: {train_eval_bin_classifier.acc:0.3f} '\n             f'Validation: {val_eval_bin_classifier.acc:0.3f} '\n             f'Test: {test_eval_bin_classifier.acc:0.3f}', fontsize=16);","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.246474Z","iopub.status.idle":"2022-02-01T20:27:35.247041Z","shell.execute_reply.started":"2022-02-01T20:27:35.246806Z","shell.execute_reply":"2022-02-01T20:27:35.24683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Compare distance classifiers","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize=(24, 5), ncols=3)\n\nsns.barplot([thresh_train_score, train_eval_bin_classifier.acc], ['Threshold', 'MLP'], ax=ax[0], palette=['tab:purple', 'tab:green'])\nsns.barplot([thresh_val_score, val_eval_bin_classifier.acc], ['Threshold', 'MLP'], ax=ax[1], palette=['tab:blue', 'tab:orange'])\nsns.barplot([thresh_test_score, test_eval_bin_classifier.acc], ['Threshold', 'MLP'], ax=ax[2], palette=['tab:grey', 'tab:red'])\n\nax[0].vlines(np.min([thresh_train_score, train_eval_bin_classifier.acc]), -0.5, 1.5, linestyle='--', color='r')\nax[1].vlines(np.min([thresh_val_score, val_eval_bin_classifier.acc]), -0.5, 1.5, linestyle='--', color='r')\nax[2].vlines(np.min([thresh_test_score, test_eval_bin_classifier.acc]), -0.5, 1.5, linestyle='--', color='r')\n\nax[0].set_title(f'Train set | Otsu {thresh_train_score:0.3f} MLP {train_eval_bin_classifier.acc:0.3f}', fontsize=15)\nax[1].set_title(f'Validation set | Otsu {thresh_val_score:0.3f} MLP {val_eval_bin_classifier.acc:0.3f}', fontsize=15)\nax[2].set_title(f'Test set | Otsu {thresh_test_score:0.3f} MLP {test_eval_bin_classifier.acc:0.3f}', fontsize=15)\n\nfor _ax in ax:\n    _ax.set_xlabel('Accuracy score', fontsize=16)\n    _ax.set_ylabel('Binary classifier', fontsize=16)\n    _ax.grid()\n    \nfig.suptitle('Threshold and MLP classifier comparison', fontsize=16);","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.248121Z","iopub.status.idle":"2022-02-01T20:27:35.248686Z","shell.execute_reply.started":"2022-02-01T20:27:35.248427Z","shell.execute_reply":"2022-02-01T20:27:35.248451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1000-class classification","metadata":{}},{"cell_type":"code","source":"knn_finder = GridSearchCV(KNeighborsClassifier(), param_grid={'n_neighbors': [1,2,3,4,5,6,7], 'weights': ['uniform', 'distance']}, n_jobs=-1)\nknn_finder.fit(train_embeddings, train_labels)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.249744Z","iopub.status.idle":"2022-02-01T20:27:35.250281Z","shell.execute_reply.started":"2022-02-01T20:27:35.250052Z","shell.execute_reply":"2022-02-01T20:27:35.250076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"knn_train_score = knn_finder.best_estimator_.score(train_embeddings, train_labels)\nknn_val_score = knn_finder.best_estimator_.score(val_embeddings, val_labels)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.251332Z","iopub.status.idle":"2022-02-01T20:27:35.251895Z","shell.execute_reply.started":"2022-02-01T20:27:35.251665Z","shell.execute_reply":"2022-02-01T20:27:35.25169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"knn_unseen_classes = KNeighborsClassifier(**knn_finder.best_params_)\n\ntest_cv_scores = cross_val_score(knn_unseen_classes, test_embeddings, test_labels, cv=5)\n\nunseen_train_emb, unseen_test_emb, unseen_train_lbl, unseen_test_lbl = tuple(train_test_split(test_embeddings, test_labels))\nknn_unseen_classes.fit(unseen_train_emb, unseen_train_lbl)","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.252952Z","iopub.status.idle":"2022-02-01T20:27:35.253498Z","shell.execute_reply.started":"2022-02-01T20:27:35.253262Z","shell.execute_reply":"2022-02-01T20:27:35.253285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize=(12, 4))\n\nsns.barplot([knn_train_score, knn_val_score, test_cv_scores.mean()], ['Train', 'Validation', 'Test'])\nax.set_xlabel('Accuracy score', fontsize=16)\nax.set_ylabel('Subset', fontsize=16)\nax.grid()\nfig.suptitle(f'Accuracy scores | Train: {knn_train_score:0.3f} | Validation: {knn_val_score:0.3f} | Unseen 200 classes: {test_cv_scores.mean():0.3f}', fontsize=16);","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.25455Z","iopub.status.idle":"2022-02-01T20:27:35.255101Z","shell.execute_reply.started":"2022-02-01T20:27:35.254872Z","shell.execute_reply":"2022-02-01T20:27:35.254896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def top_k_knn_acc(clf, emb, lbl, ks, classes):\n    accs = []\n    for k in tqdm(ks):\n        y_proba = clf.predict_proba(emb)\n        lbl_oh = tf.one_hot(lbl, classes)\n        acc = np.mean(top_k_categorical_accuracy(lbl_oh, y_proba, k=k))\n        accs.append(acc)\n        \n    return accs\n\nks = [1, 2, 3, 4, 5, 10]\ntrain_knn_top_k_acc = top_k_knn_acc(knn_finder.best_estimator_, train_embeddings, train_labels, ks, classes=800)\nval_knn_top_k_acc = top_k_knn_acc(knn_finder.best_estimator_, val_embeddings, val_labels, ks, classes=800)\ntest_knn_top_k_acc = top_k_knn_acc(knn_unseen_classes, unseen_test_emb, unseen_test_lbl, ks, classes=200)\n\n\nfig, ax = plt.subplots(figsize=(12, 4))\n\nax.plot(ks, train_knn_top_k_acc, label='Train')\nax.plot(ks, val_knn_top_k_acc, label='Validation')\nax.plot(ks, test_knn_top_k_acc, label='Test')\n\nax.scatter(ks, train_knn_top_k_acc)\nax.scatter(ks, val_knn_top_k_acc, label='Validation')\nax.scatter(ks, test_knn_top_k_acc, label='Test')\nax.legend()\nax.grid()\nfig.suptitle('Top-K accuracy', fontsize=16);","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.256169Z","iopub.status.idle":"2022-02-01T20:27:35.256733Z","shell.execute_reply.started":"2022-02-01T20:27:35.256476Z","shell.execute_reply":"2022-02-01T20:27:35.256508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Example of predictions","metadata":{}},{"cell_type":"code","source":"def load_sample_from_row(row):\n    img = cv2.imread(row.in_file)[..., ::-1]\n    bb = bb_from_row(row)\n    bb = [int(i*224) for i in bb]\n    img = img[bb[2]:bb[3], bb[0]:bb[1]]\n    img = tf.image.resize(img, (224, 244))\n    img /= 255\n    label = row.class_id\n    return img, label\n\ndef predict_img(emb_extractor, classifier_head, img):\n    img = tf.expand_dims(img, axis=0)\n    embeddings = emb_extractor.predict(img)\n    pred = classifier_head.predict(embeddings)\n    return pred","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.257798Z","iopub.status.idle":"2022-02-01T20:27:35.258331Z","shell.execute_reply.started":"2022-02-01T20:27:35.258104Z","shell.execute_reply":"2022-02-01T20:27:35.258128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prediction on validatinon set\n\n### Random cases","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=3, figsize=(24, 12))\n\nnp.random.seed(0)\nfor _ax, (idx, row) in zip(ax.ravel(), csv_val.sample(15).iterrows()):\n    img, label = load_sample_from_row(row)\n    pred = predict_img(model, knn_finder.best_estimator_, img)\n    \n    _ax.imshow(img)\n    \n    gt_class = row['class']\n    pred_class = class_id_to_str(pred[0])\n    _ax.set_title(f'Gt: {gt_class}\\nPred: {pred_class}', color='g' if gt_class == pred_class else 'r')\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.259376Z","iopub.status.idle":"2022-02-01T20:27:35.259932Z","shell.execute_reply.started":"2022-02-01T20:27:35.259707Z","shell.execute_reply":"2022-02-01T20:27:35.259731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Incorrect cases","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=3, figsize=(24, 12))\n\nnp.random.seed(0)\nfor _ax in ax.ravel():\n    for idx, row in csv_val.sample(frac=1).iterrows():\n        img, label = load_sample_from_row(row)\n        pred = predict_img(model, knn_finder.best_estimator_, img)\n\n\n        gt_class = row['class']\n        pred_class = class_id_to_str(pred[0])\n        \n        if gt_class != pred_class:\n            break\n            \n    _ax.imshow(img)\n    _ax.set_title(f'Gt: {gt_class}\\nPred: {pred_class}', color='g' if gt_class == pred_class else 'r')\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.260987Z","iopub.status.idle":"2022-02-01T20:27:35.261531Z","shell.execute_reply.started":"2022-02-01T20:27:35.261293Z","shell.execute_reply":"2022-02-01T20:27:35.261316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Unseen classes predictions\n\n## Random cases","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=3, figsize=(24, 12))\n\nnp.random.seed(0)\nfor _ax, (idx, row) in zip(ax.ravel(), csv_test.sample(15).iterrows()):\n    img, label = load_sample_from_row(row)\n    pred = predict_img(model, knn_unseen_classes, img)\n    \n    _ax.imshow(img)\n    \n    gt_class = row['class']\n    pred_class = class_id_to_str(pred[0])\n    _ax.set_title(f'Gt: {gt_class}\\nPred: {pred_class}', color='g' if gt_class == pred_class else 'r')\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.262606Z","iopub.status.idle":"2022-02-01T20:27:35.263163Z","shell.execute_reply.started":"2022-02-01T20:27:35.262937Z","shell.execute_reply":"2022-02-01T20:27:35.26296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Incorrect cases","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(ncols=5, nrows=3, figsize=(24, 12))\n\nnp.random.seed(0)\nfor _ax in ax.ravel():\n    for idx, row in csv_test.sample(frac=1).iterrows():\n        img, label = load_sample_from_row(row)\n        pred = predict_img(model, knn_unseen_classes, img)\n\n\n        gt_class = row['class']\n        pred_class = class_id_to_str(pred[0])\n        \n        if gt_class != pred_class:\n            break\n            \n    _ax.imshow(img)\n    _ax.set_title(f'Gt: {gt_class}\\nPred: {pred_class}', color='g' if gt_class == pred_class else 'r')\n    _ax.axis('off')","metadata":{"execution":{"iopub.status.busy":"2022-02-01T20:27:35.264234Z","iopub.status.idle":"2022-02-01T20:27:35.264799Z","shell.execute_reply.started":"2022-02-01T20:27:35.264556Z","shell.execute_reply":"2022-02-01T20:27:35.264579Z"},"trusted":true},"execution_count":null,"outputs":[]}]}