{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":18237,"databundleVersionId":1053191,"sourceType":"competition"},{"sourceId":7224777,"sourceType":"datasetVersion","datasetId":4182256},{"sourceId":7226070,"sourceType":"datasetVersion","datasetId":4183141}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\nimport numpy as np # linear algebra\nimport 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\nimport os\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\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":"2023-12-17T22:23:15.213364Z","iopub.execute_input":"2023-12-17T22:23:15.213751Z","iopub.status.idle":"2023-12-17T22:23:15.560384Z","shell.execute_reply.started":"2023-12-17T22:23:15.213714Z","shell.execute_reply":"2023-12-17T22:23:15.559618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install faiss-gpu","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:23:16.033697Z","iopub.execute_input":"2023-12-17T22:23:16.034161Z","iopub.status.idle":"2023-12-17T22:23:31.355980Z","shell.execute_reply.started":"2023-12-17T22:23:16.034134Z","shell.execute_reply":"2023-12-17T22:23:31.354882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport json\nimport PIL\nimport faiss\nimport random\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\n\nfrom tqdm import tqdm\nfrom PIL import Image\n\nfrom torchvision.models import resnet50\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import MultiLabelBinarizer\nfrom sklearn.metrics import accuracy_score","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:23:31.358169Z","iopub.execute_input":"2023-12-17T22:23:31.358536Z","iopub.status.idle":"2023-12-17T22:23:35.474601Z","shell.execute_reply.started":"2023-12-17T22:23:31.358502Z","shell.execute_reply":"2023-12-17T22:23:35.473815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = {\n    'batch_size': 64,\n    'learning_rate': 1e-4,\n    'num_epochs': 10,\n    'device': \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    'num_workers': 4,\n    \"num_classes\": 46\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:32:15.914997Z","iopub.execute_input":"2023-12-17T23:32:15.915357Z","iopub.status.idle":"2023-12-17T23:32:15.920829Z","shell.execute_reply.started":"2023-12-17T23:32:15.915322Z","shell.execute_reply":"2023-12-17T23:32:15.919901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_desc_path = \"/kaggle/input/imaterialist-fashion-2020-fgvc7/label_descriptions.json\"\nannotations_path = \"/kaggle/input/imaterialist-fashion-2020-fgvc7/train.csv\"\ntrain_images_path = \"/kaggle/input/imaterialist-fashion-2020-fgvc7/train\"\ntest_images_path = \"/kaggle/input/imaterialist-fashion-2020-fgvc7/test\"","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:11:06.017412Z","iopub.execute_input":"2023-12-18T00:11:06.018299Z","iopub.status.idle":"2023-12-18T00:11:06.022809Z","shell.execute_reply.started":"2023-12-18T00:11:06.018262Z","shell.execute_reply":"2023-12-18T00:11:06.021782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_json(json_path):\n    with open(json_path) as file:\n        data = json.load(file)\n        \n    return data","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:23:35.524655Z","iopub.execute_input":"2023-12-17T22:23:35.524950Z","iopub.status.idle":"2023-12-17T22:23:35.534100Z","shell.execute_reply.started":"2023-12-17T22:23:35.524926Z","shell.execute_reply":"2023-12-17T22:23:35.533243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"annotations_df = pd.read_csv(annotations_path)\nlabels = load_json(label_desc_path)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:23:35.535076Z","iopub.execute_input":"2023-12-17T22:23:35.535331Z","iopub.status.idle":"2023-12-17T22:24:04.076526Z","shell.execute_reply.started":"2023-12-17T22:23:35.535309Z","shell.execute_reply":"2023-12-17T22:24:04.075695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aggregated_classes = annotations_df.groupby('ImageId')[\"ClassId\"].agg(list).reset_index()","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:04.077994Z","iopub.execute_input":"2023-12-17T22:24:04.078354Z","iopub.status.idle":"2023-12-17T22:24:05.402036Z","shell.execute_reply.started":"2023-12-17T22:24:04.078309Z","shell.execute_reply":"2023-12-17T22:24:05.401219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mlb = MultiLabelBinarizer()\none_hot_classes = mlb.fit_transform(aggregated_classes[\"ClassId\"]).tolist()\n\naggregated_classes['EncodedClasses'] = one_hot_classes","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:05.403153Z","iopub.execute_input":"2023-12-17T22:24:05.403492Z","iopub.status.idle":"2023-12-17T22:24:05.755647Z","shell.execute_reply.started":"2023-12-17T22:24:05.403460Z","shell.execute_reply":"2023-12-17T22:24:05.754877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aggregated_classes","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:05.756886Z","iopub.execute_input":"2023-12-17T22:24:05.757239Z","iopub.status.idle":"2023-12-17T22:24:05.783940Z","shell.execute_reply.started":"2023-12-17T22:24:05.757207Z","shell.execute_reply":"2023-12-17T22:24:05.783071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, val_df = train_test_split(aggregated_classes, test_size=0.1, random_state=42)\n\ntrain_df.reset_index(inplace=True)\nval_df.reset_index(inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:05.787292Z","iopub.execute_input":"2023-12-17T22:24:05.787624Z","iopub.status.idle":"2023-12-17T22:24:05.801712Z","shell.execute_reply.started":"2023-12-17T22:24:05.787601Z","shell.execute_reply":"2023-12-17T22:24:05.800828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_distribution(df, descr):\n    label_counts = df[\"ClassId\"].explode().value_counts()\n\n    classes = label_counts.index\n    counts = label_counts.values\n\n    plt.figure(figsize=(15, 5))\n    sns.barplot(x=classes, y=counts)\n    plt.title(f\"Distribution of classes for {descr}\")\n    plt.xticks(rotation=90)\n    plt.show()\n    \n    \ndef plot_image(image, labels, denormalize=False):\n    if denormalize:\n        image = denorm(image)\n\n    image = image.permute(1, 2, 0)\n    \n    plt.imshow(image)\n    plt.title(f'Image labels: {labels}')\n    plt.axis('off')\n    plt.show()\n    \n    \ndef calculate_classification_metrics(model, dataloader):\n    model.eval()\n    \n    correct_preds = 0\n    total_samples = 0\n    \n    true_positives = 0\n    false_positives = 0\n    false_negatives = 0\n    \n    with torch.no_grad():\n        for inputs, labels in dataloader:\n            inputs = inputs.to(config[\"device\"])\n            labels = labels.to(config[\"device\"])\n\n            outputs = model(inputs)\n\n            preds = torch.sigmoid(outputs)\n\n            preds = preds.detach().cpu() > 0.5\n            labels = labels.detach().cpu()\n\n            correct_preds += torch.sum(preds == labels).item()\n            total_samples += labels.numel()\n            \n            true_positives += torch.sum(preds & labels).item()\n            false_positives += torch.sum(preds & ~labels).item()\n            false_negatives += torch.sum(~preds & labels).item()\n            \n    precision = true_positives / (true_positives + false_positives)\n    recall = true_positives / (true_positives + false_negatives)\n    accuracy = correct_preds / total_samples\n        \n    return accuracy, precision, recall","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:05.802629Z","iopub.execute_input":"2023-12-17T22:24:05.802905Z","iopub.status.idle":"2023-12-17T22:24:05.814314Z","shell.execute_reply.started":"2023-12-17T22:24:05.802879Z","shell.execute_reply":"2023-12-17T22:24:05.813436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_distribution(train_df, \"Train data\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:05.815559Z","iopub.execute_input":"2023-12-17T22:24:05.816022Z","iopub.status.idle":"2023-12-17T22:24:06.462568Z","shell.execute_reply.started":"2023-12-17T22:24:05.815990Z","shell.execute_reply":"2023-12-17T22:24:06.461685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_distribution(val_df, \"Val data\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:06.463772Z","iopub.execute_input":"2023-12-17T22:24:06.464094Z","iopub.status.idle":"2023-12-17T22:24:07.042625Z","shell.execute_reply.started":"2023-12-17T22:24:06.464067Z","shell.execute_reply":"2023-12-17T22:24:07.041782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MutlilabelClassificationDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n        \n    def __getitem__(self, idx):\n        entry = self.df.iloc[idx]\n        image_id = entry[\"ImageId\"]\n        image_path = os.path.join(train_images_path, f\"{image_id}.jpg\")\n        class_ids_encoded = entry[\"EncodedClasses\"]\n        \n        image = PIL.Image.open(image_path).convert('RGB')\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, torch.tensor(class_ids_encoded)\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:07.043616Z","iopub.execute_input":"2023-12-17T22:24:07.043890Z","iopub.status.idle":"2023-12-17T22:24:07.050915Z","shell.execute_reply.started":"2023-12-17T22:24:07.043867Z","shell.execute_reply":"2023-12-17T22:24:07.049989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = [0.485, 0.456, 0.406]\nstd = [0.229, 0.224, 0.225]\n\nnorm = transforms.Normalize(mean=mean, std=std)\ndenorm = transforms.Normalize(mean=[-m/s for m, s in zip(mean, std)], std=[1/s for s in std])\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.ColorJitter(brightness = 0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    norm\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    norm\n])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:11:43.658477Z","iopub.execute_input":"2023-12-18T00:11:43.659216Z","iopub.status.idle":"2023-12-18T00:11:43.666465Z","shell.execute_reply.started":"2023-12-18T00:11:43.659161Z","shell.execute_reply":"2023-12-18T00:11:43.665521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = MutlilabelClassificationDataset(train_df[:7000], transform=train_transform)\nval_dataset = MutlilabelClassificationDataset(val_df[:700], transform=val_transform)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T19:37:13.391469Z","iopub.execute_input":"2023-12-17T19:37:13.392209Z","iopub.status.idle":"2023-12-17T19:37:13.397232Z","shell.execute_reply.started":"2023-12-17T19:37:13.392170Z","shell.execute_reply":"2023-12-17T19:37:13.396126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])\nval_loader = DataLoader(val_dataset, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T12:57:00.334344Z","iopub.execute_input":"2023-12-17T12:57:00.334986Z","iopub.status.idle":"2023-12-17T12:57:00.340335Z","shell.execute_reply.started":"2023-12-17T12:57:00.334953Z","shell.execute_reply":"2023-12-17T12:57:00.339229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultilabelClassificationRESNET(nn.Module):\n    def __init__(self, num_classes, pretrained=False):\n        super(MultilabelClassificationRESNET, self).__init__()\n        \n        self.resnet = resnet50(pretrained=pretrained)\n        \n        self.in_features = self.resnet.fc.in_features\n        \n        self.resnet = torch.nn.Sequential(*(list(self.resnet.children())[:-1]))\n        \n        self.fc = nn.Sequential(\n            nn.Linear(self.in_features, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, num_classes),\n        )\n        \n    def forward(self, x):\n        x = self.resnet(x)\n        \n        x = x.view(x.size(0), -1)\n        \n        return self.fc(x)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:07.067779Z","iopub.execute_input":"2023-12-17T22:24:07.068317Z","iopub.status.idle":"2023-12-17T22:24:07.077610Z","shell.execute_reply.started":"2023-12-17T22:24:07.068285Z","shell.execute_reply":"2023-12-17T22:24:07.076798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training model from scratch","metadata":{}},{"cell_type":"code","source":"model = MultilabelClassificationRESNET(config[\"num_classes\"]).to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T13:17:49.724848Z","iopub.execute_input":"2023-12-17T13:17:49.725798Z","iopub.status.idle":"2023-12-17T13:17:50.145257Z","shell.execute_reply.started":"2023-12-17T13:17:49.725765Z","shell.execute_reply":"2023-12-17T13:17:50.144430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(), lr=config[\"learning_rate\"])\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:28:17.306355Z","iopub.execute_input":"2023-12-17T14:28:17.306756Z","iopub.status.idle":"2023-12-17T14:28:17.317393Z","shell.execute_reply.started":"2023-12-17T14:28:17.306722Z","shell.execute_reply":"2023-12-17T14:28:17.316590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, train_loader, val_loader):\n    train_loss = []\n    val_loss = []\n\n    for epoch in range(1, config[\"num_epochs\"] + 1):\n        model.train()\n        total_train_loss = 0\n\n        current_lr1 = optimizer.param_groups[0]['lr']\n\n        for (inputs, targets) in tqdm(train_loader, leave=False, desc=f'Epoch {epoch}, lr 1: {current_lr1}'):\n            inputs = inputs.to(config[\"device\"])\n            targets = targets.to(config[\"device\"])\n\n            outputs = model(inputs)\n            \n            loss = criterion(outputs, targets.float())\n\n            loss.backward()\n            \n            optimizer.step()\n            optimizer.zero_grad()\n\n            total_train_loss += loss.item()\n\n        avg_train_loss = total_train_loss / len(train_loader)\n        train_loss.append(avg_train_loss)\n\n        print(f\"Average train loss at {epoch} epoch: {avg_train_loss}\")\n\n        total_val_loss = 0\n\n        model.eval()\n        for (inputs, targets) in tqdm(val_loader, leave=False, desc=f'Validation - Epoch {epoch}'):\n            inputs = inputs.to(config[\"device\"])\n            targets = targets.to(config[\"device\"])\n\n            with torch.no_grad():\n                outputs = model(inputs)\n\n                loss = criterion(outputs, targets.float())\n\n            total_val_loss += loss.item()\n\n        avg_val_loss = total_val_loss / len(val_loader)\n        val_loss.append(avg_val_loss)\n        \n        accuracy, precision, recall = calculate_classification_metrics(model, val_loader)\n        \n        scheduler.step()\n\n        print(f\"Average test loss at {epoch} epoch: {avg_val_loss}\")\n        print(f\"Average acc: {accuracy}, precision: {precision}, recall: {recall}\")\n        \n    return train_loss, val_loss","metadata":{"execution":{"iopub.status.busy":"2023-12-17T13:17:51.598127Z","iopub.execute_input":"2023-12-17T13:17:51.598872Z","iopub.status.idle":"2023-12-17T13:17:51.610060Z","shell.execute_reply.started":"2023-12-17T13:17:51.598832Z","shell.execute_reply":"2023-12-17T13:17:51.609077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss, val_loss = train(model, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T13:17:52.119651Z","iopub.execute_input":"2023-12-17T13:17:52.120406Z","iopub.status.idle":"2023-12-17T14:19:49.966334Z","shell.execute_reply.started":"2023-12-17T13:17:52.120376Z","shell.execute_reply":"2023-12-17T14:19:49.965089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"model_from_scratch.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:21:23.374884Z","iopub.execute_input":"2023-12-17T14:21:23.375862Z","iopub.status.idle":"2023-12-17T14:21:23.641321Z","shell.execute_reply.started":"2023-12-17T14:21:23.375825Z","shell.execute_reply":"2023-12-17T14:21:23.640467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_losses(train_loss, validation_loss):\n    epochs = range(1, len(train_loss) + 1)\n\n    plt.figure(figsize=(10, 6))\n    plt.plot(epochs, train_loss, label='Train Loss', marker='o')\n    plt.plot(epochs, validation_loss, label='Validation Loss', marker='o')\n\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss')\n    plt.legend()\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-17T22:24:07.078585Z","iopub.execute_input":"2023-12-17T22:24:07.078897Z","iopub.status.idle":"2023-12-17T22:24:07.088434Z","shell.execute_reply.started":"2023-12-17T22:24:07.078869Z","shell.execute_reply":"2023-12-17T22:24:07.087642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_losses(train_loss, val_loss)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:21:34.089047Z","iopub.execute_input":"2023-12-17T14:21:34.089956Z","iopub.status.idle":"2023-12-17T14:21:34.406152Z","shell.execute_reply.started":"2023-12-17T14:21:34.089917Z","shell.execute_reply":"2023-12-17T14:21:34.404589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(\"cpu\")\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:21:40.162031Z","iopub.execute_input":"2023-12-17T14:21:40.162817Z","iopub.status.idle":"2023-12-17T14:21:40.453719Z","shell.execute_reply.started":"2023-12-17T14:21:40.162784Z","shell.execute_reply":"2023-12-17T14:21:40.452705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finetunning model","metadata":{}},{"cell_type":"code","source":"model = MultilabelClassificationRESNET(config[\"num_classes\"], pretrained=True).to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:21:42.149447Z","iopub.execute_input":"2023-12-17T14:21:42.150326Z","iopub.status.idle":"2023-12-17T14:21:44.590232Z","shell.execute_reply.started":"2023-12-17T14:21:42.150292Z","shell.execute_reply":"2023-12-17T14:21:44.589098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for param in model.resnet.parameters():\n    param.requires_grad = False\n\nfor param in model.fc.parameters():\n    param.requires_grad = True\n\nfor param in model.resnet[-2].parameters():\n    param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:28:02.408754Z","iopub.execute_input":"2023-12-17T14:28:02.409438Z","iopub.status.idle":"2023-12-17T14:28:02.416132Z","shell.execute_reply.started":"2023-12-17T14:28:02.409404Z","shell.execute_reply":"2023-12-17T14:28:02.415094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss, val_loss = train(model, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T14:28:22.805116Z","iopub.execute_input":"2023-12-17T14:28:22.805455Z","iopub.status.idle":"2023-12-17T15:27:56.923035Z","shell.execute_reply.started":"2023-12-17T14:28:22.805429Z","shell.execute_reply":"2023-12-17T15:27:56.921785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"finetuned_model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T15:27:56.925310Z","iopub.execute_input":"2023-12-17T15:27:56.925852Z","iopub.status.idle":"2023-12-17T15:27:57.098161Z","shell.execute_reply.started":"2023-12-17T15:27:56.925813Z","shell.execute_reply":"2023-12-17T15:27:57.097321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_losses(train_loss, val_loss)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T15:28:00.814717Z","iopub.execute_input":"2023-12-17T15:28:00.815089Z","iopub.status.idle":"2023-12-17T15:28:01.135874Z","shell.execute_reply.started":"2023-12-17T15:28:00.815058Z","shell.execute_reply":"2023-12-17T15:28:01.134886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metric learning","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(\"/kaggle/input/finetunned-model/finetuned_model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:44:27.868285Z","iopub.execute_input":"2023-12-17T18:44:27.869380Z","iopub.status.idle":"2023-12-17T18:44:28.858179Z","shell.execute_reply.started":"2023-12-17T18:44:27.869342Z","shell.execute_reply":"2023-12-17T18:44:28.857177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = MultilabelClassificationRESNET(config[\"num_classes\"], pretrained=True).to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:45:42.973568Z","iopub.execute_input":"2023-12-17T18:45:42.974483Z","iopub.status.idle":"2023-12-17T18:45:43.992389Z","shell.execute_reply.started":"2023-12-17T18:45:42.974444Z","shell.execute_reply":"2023-12-17T18:45:43.991537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(checkpoint)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:46:28.627482Z","iopub.execute_input":"2023-12-17T18:46:28.627865Z","iopub.status.idle":"2023-12-17T18:46:28.647219Z","shell.execute_reply.started":"2023-12-17T18:46:28.627832Z","shell.execute_reply":"2023-12-17T18:46:28.646437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SiameseDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __getitem__(self, index):\n        anchor = self.df.iloc[index]\n        \n        label = torch.tensor(anchor[\"EncodedClasses\"])\n            \n        candidates = self.df[self.df.index != index]\n        candidates = candidates.sample(frac=1)\n        \n        for idx, candidate in candidates.iterrows():\n            intersection = set(anchor['ClassId']) & set(candidate['ClassId'])\n            \n            if len(intersection) > 0:\n                positive_index = idx\n                break\n                \n        for idx, candidate in candidates.iterrows():\n            intersection = set(anchor['ClassId']) & set(candidate['ClassId'])\n            \n            if len(intersection) == 0:\n                negative_index = idx\n                break\n                \n        positive = self.df.iloc[positive_index]\n        negative = self.df.iloc[negative_index]\n\n        anchor_image = self.load_image(anchor[\"ImageId\"])\n        positive_image = self.load_image(positive[\"ImageId\"])\n        negative_image = self.load_image(negative[\"ImageId\"])\n            \n        if self.transform:\n            anchor_image = self.transform(anchor_image)\n            positive_image = self.transform(positive_image)\n            negative_image = self.transform(negative_image)\n\n        return anchor_image, positive_image, negative_image, label\n    \n    def load_image(self, image_id):\n        image_path = os.path.join(train_images_path, f\"{image_id}.jpg\")\n        image = PIL.Image.open(image_path).convert('RGB')\n        \n        return image\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:53:14.622871Z","iopub.execute_input":"2023-12-17T23:53:14.623666Z","iopub.status.idle":"2023-12-17T23:53:14.634539Z","shell.execute_reply.started":"2023-12-17T23:53:14.623633Z","shell.execute_reply":"2023-12-17T23:53:14.633584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SiameseNetwork(nn.Module):\n    def __init__(self, num_classes, pretrained=False):\n        super(SiameseNetwork, self).__init__()\n        \n        self.resnet = resnet50(pretrained=pretrained)\n        \n        self.in_features = self.resnet.fc.in_features\n        \n        self.resnet = nn.Sequential(*list(self.resnet.children())[:-1])\n        \n        self.fc = nn.Sequential(\n            nn.Linear(self.in_features, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, num_classes),\n        )\n        \n    def forward_once(self, x):\n        output = self.resnet(x)\n        \n        output = output.view(output.size(0), -1)\n        \n        return self.fc(output)\n        \n    def forward(self, anchor, pos, neg):\n        anchor = self.forward_once(anchor)\n        pos = self.forward_once(pos)\n        neg = self.forward_once(neg)\n        \n        return anchor, pos, neg","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:04.312920Z","iopub.execute_input":"2023-12-17T18:49:04.313335Z","iopub.status.idle":"2023-12-17T18:49:04.321284Z","shell.execute_reply.started":"2023-12-17T18:49:04.313305Z","shell.execute_reply":"2023-12-17T18:49:04.320218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset_example = SiameseDataset(train_df[:5000])\nval_dataset_example = SiameseDataset(val_df[:500], transform=val_transform)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:53:18.736501Z","iopub.execute_input":"2023-12-17T23:53:18.737436Z","iopub.status.idle":"2023-12-17T23:53:18.744887Z","shell.execute_reply.started":"2023-12-17T23:53:18.737401Z","shell.execute_reply":"2023-12-17T23:53:18.743896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader_example = DataLoader(train_dataset_example, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])\nval_loader_example = DataLoader(val_dataset_example, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:07.663429Z","iopub.execute_input":"2023-12-17T18:49:07.663786Z","iopub.status.idle":"2023-12-17T18:49:07.669213Z","shell.execute_reply.started":"2023-12-17T18:49:07.663756Z","shell.execute_reply":"2023-12-17T18:49:07.668195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"anchor, positive, negative, labels = train_dataset_example[5]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:03:24.778569Z","iopub.execute_input":"2023-12-18T00:03:24.778954Z","iopub.status.idle":"2023-12-18T00:03:24.856854Z","shell.execute_reply.started":"2023-12-18T00:03:24.778926Z","shell.execute_reply":"2023-12-18T00:03:24.855888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image(anchor, \"Anchor\", denormalize=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image(positive, \"Positive\", denormalize=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T15:43:33.705961Z","iopub.execute_input":"2023-12-17T15:43:33.706720Z","iopub.status.idle":"2023-12-17T15:43:33.939524Z","shell.execute_reply.started":"2023-12-17T15:43:33.706688Z","shell.execute_reply":"2023-12-17T15:43:33.938559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image(negative, \"Negative\", denormalize=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T15:43:35.840148Z","iopub.execute_input":"2023-12-17T15:43:35.840514Z","iopub.status.idle":"2023-12-17T15:43:36.021170Z","shell.execute_reply.started":"2023-12-17T15:43:35.840483Z","shell.execute_reply":"2023-12-17T15:43:36.020023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"siamese_model = SiameseNetwork(config[\"num_classes\"]).to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:14.304071Z","iopub.execute_input":"2023-12-17T18:49:14.304828Z","iopub.status.idle":"2023-12-17T18:49:14.318080Z","shell.execute_reply.started":"2023-12-17T18:49:14.304791Z","shell.execute_reply":"2023-12-17T18:49:14.317021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for param in siamese_model.resnet.parameters():\n    param.requires_grad = False\n\nfor param in siamese_model.fc.parameters():\n    param.requires_grad = True\n\nfor param in siamese_model.resnet[-2].parameters():\n    param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:26.925397Z","iopub.execute_input":"2023-12-17T18:49:26.925772Z","iopub.status.idle":"2023-12-17T18:49:26.932447Z","shell.execute_reply.started":"2023-12-17T18:49:26.925741Z","shell.execute_reply":"2023-12-17T18:49:26.931383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"triplet_loss_criterion = nn.TripletMarginLoss(margin=1.0, p=2)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:29.087085Z","iopub.execute_input":"2023-12-17T18:49:29.087809Z","iopub.status.idle":"2023-12-17T18:49:29.092097Z","shell.execute_reply.started":"2023-12-17T18:49:29.087771Z","shell.execute_reply":"2023-12-17T18:49:29.091124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.AdamW(siamese_model.parameters(), lr=config[\"learning_rate\"])\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:49:30.378933Z","iopub.execute_input":"2023-12-17T18:49:30.379885Z","iopub.status.idle":"2023-12-17T18:49:30.385971Z","shell.execute_reply.started":"2023-12-17T18:49:30.379848Z","shell.execute_reply":"2023-12-17T18:49:30.385007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss = []\nval_loss = []\n\nfor epoch in range(1, config[\"num_epochs\"] + 1):\n    siamese_model.train()\n    total_train_loss = 0\n\n    current_lr1 = optimizer.param_groups[0]['lr']\n\n    for (anchor, positive, negative, targets) in tqdm(train_loader_example, leave=False, desc=f'Epoch {epoch}, lr 1: {current_lr1}'):\n        anchor, positive, negative = anchor.to(config[\"device\"]), positive.to(config[\"device\"]), negative.to(config[\"device\"])\n        targets = targets.to(config[\"device\"])\n\n        anchor_embedding, positive_embedding, negative_embedding = siamese_model(anchor, positive, negative)\n\n        loss = triplet_loss_criterion(anchor_embedding, positive_embedding, negative_embedding)\n\n        loss.backward()\n\n        optimizer.step()\n        optimizer.zero_grad()\n\n        total_train_loss += loss.item()\n\n    avg_train_loss = total_train_loss / len(train_loader_example)\n    train_loss.append(avg_train_loss)\n\n    print(f\"Average train loss at {epoch} epoch: {avg_train_loss}\")\n\n    total_val_loss = 0\n\n    siamese_model.eval()\n    \n    for (anchor, positive, negative, targets) in tqdm(val_loader_example, leave=False, desc=f'Validation - Epoch {epoch}'):\n        anchor, positive, negative = anchor.to(config[\"device\"]), positive.to(config[\"device\"]), negative.to(config[\"device\"])\n        targets = targets.to(config[\"device\"])\n\n        with torch.no_grad():\n            anchor_embedding, positive_embedding, negative_embedding = siamese_model(anchor, positive, negative)\n\n            loss = triplet_loss_criterion(anchor_embedding, positive_embedding, negative_embedding)\n\n        total_val_loss += loss.item()\n\n    avg_val_loss = total_val_loss / len(val_loader_example)\n    val_loss.append(avg_val_loss)\n\n    scheduler.step()\n\n    print(f\"Average test loss at {epoch} epoch: {avg_val_loss}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(siamese_model.state_dict(), \"siamese_model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:16:08.059456Z","iopub.execute_input":"2023-12-17T18:16:08.059876Z","iopub.status.idle":"2023-12-17T18:16:08.234539Z","shell.execute_reply.started":"2023-12-17T18:16:08.059836Z","shell.execute_reply":"2023-12-17T18:16:08.233672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_losses(train_loss, val_loss)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T18:16:46.403733Z","iopub.execute_input":"2023-12-17T18:16:46.404422Z","iopub.status.idle":"2023-12-17T18:16:46.699129Z","shell.execute_reply.started":"2023-12-17T18:16:46.404387Z","shell.execute_reply":"2023-12-17T18:16:46.698220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Self-supervised SimCLR","metadata":{}},{"cell_type":"code","source":"class SimCLRAugmentor:\n    def __init__(self, img_size, s=1):\n        color_jitter = transforms.ColorJitter(0.8 * s, 0.8 * s, 0.8 * s, 0.2 * s)\n        blur = transforms.GaussianBlur((3, 3), (0.1, 2.0))\n        \n        self.transform = transforms.Compose([\n            transforms.RandomResizedCrop(size=img_size),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomApply([color_jitter], p=0.8),\n            transforms.RandomApply([blur], p=0.5),\n            transforms.RandomGrayscale(p=0.2),\n            transforms.ToTensor(),\n            norm\n        ])\n        \n    def __call__(self, x):\n        return self.transform(x), self.transform(x)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:54:47.636992Z","iopub.execute_input":"2023-12-17T23:54:47.637799Z","iopub.status.idle":"2023-12-17T23:54:47.644821Z","shell.execute_reply.started":"2023-12-17T23:54:47.637752Z","shell.execute_reply":"2023-12-17T23:54:47.643813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CLRtransform = transforms.Compose([transforms.Resize((224, 224)), transforms.ToTensor()])\n\ntrain_dataset = MutlilabelClassificationDataset(train_df[:4096], transform=CLRtransform)\nval_dataset = MutlilabelClassificationDataset(val_df[:512], transform=CLRtransform)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:38.093309Z","iopub.execute_input":"2023-12-17T23:20:38.094213Z","iopub.status.idle":"2023-12-17T23:20:38.100019Z","shell.execute_reply.started":"2023-12-17T23:20:38.094178Z","shell.execute_reply":"2023-12-17T23:20:38.099036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])\nval_loader = DataLoader(val_dataset, batch_size=config[\"batch_size\"], shuffle=True, num_workers=config[\"num_workers\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:40.292538Z","iopub.execute_input":"2023-12-17T23:20:40.292923Z","iopub.status.idle":"2023-12-17T23:20:40.297992Z","shell.execute_reply.started":"2023-12-17T23:20:40.292893Z","shell.execute_reply":"2023-12-17T23:20:40.296998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SimCLRModel(nn.Module):\n    def __init__(self, projection_dim=128):\n        super(SimCLRModel, self).__init__()\n\n        self.resnet = resnet50(pretrained=True)\n        \n        self.in_features = self.resnet.fc.in_features\n        \n        self.resnet = nn.Sequential(*list(self.resnet.children())[:-1])\n        \n        self.projector = nn.Sequential(\n            nn.Linear(2048, projection_dim),\n            nn.ReLU(),\n            nn.Linear(projection_dim, projection_dim),\n            nn.ReLU()\n        )\n\n    def forward(self, x):\n        h = self.resnet(x)\n        \n        z = h.view(h.size(0), -1)\n\n        z = self.projector(z)\n\n        return h, z","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:45.477196Z","iopub.execute_input":"2023-12-17T23:20:45.477919Z","iopub.status.idle":"2023-12-17T23:20:45.484974Z","shell.execute_reply.started":"2023-12-17T23:20:45.477886Z","shell.execute_reply":"2023-12-17T23:20:45.484136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrastiveLoss(nn.Module):\n    def __init__(self, batch_size, temperature=0.5):\n        super().__init__()\n        self.batch_size = batch_size\n        self.temperature = temperature\n        self.mask = (~torch.eye(batch_size * 2, batch_size * 2, dtype=bool)).float()\n\n    def calc_similarity_batch(self, a, b):\n        representations = torch.cat([a, b], dim=0)\n        return F.cosine_similarity(representations.unsqueeze(1), representations.unsqueeze(0), dim=2)\n\n    def forward(self, proj_1, proj_2):\n        batch_size = proj_1.shape[0]\n        z_i = F.normalize(proj_1, p=2, dim=1)\n        z_j = F.normalize(proj_2, p=2, dim=1)\n\n        similarity_matrix = self.calc_similarity_batch(z_i, z_j)\n\n        sim_ij = torch.diag(similarity_matrix, batch_size)\n        sim_ji = torch.diag(similarity_matrix, -batch_size)\n\n        positives = torch.cat([sim_ij, sim_ji], dim=0)\n\n        nominator = torch.exp(positives / self.temperature)\n\n        denominator = self.mask.to(config[\"device\"]) * torch.exp(similarity_matrix / self.temperature)\n\n        all_losses = -torch.log(nominator / torch.sum(denominator, dim=1))\n        loss = torch.sum(all_losses) / (2 * self.batch_size)\n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:46.163369Z","iopub.execute_input":"2023-12-17T23:20:46.163687Z","iopub.status.idle":"2023-12-17T23:20:46.173249Z","shell.execute_reply.started":"2023-12-17T23:20:46.163661Z","shell.execute_reply":"2023-12-17T23:20:46.172328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = SimCLRModel().to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:47.174374Z","iopub.execute_input":"2023-12-17T23:20:47.175075Z","iopub.status.idle":"2023-12-17T23:20:47.685398Z","shell.execute_reply.started":"2023-12-17T23:20:47.175037Z","shell.execute_reply":"2023-12-17T23:20:47.684348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = ContrastiveLoss(config[\"batch_size\"])\noptimizer = optim.AdamW(model.parameters(), lr=config[\"learning_rate\"])\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=4, gamma=0.5)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:50.077930Z","iopub.execute_input":"2023-12-17T23:20:50.078300Z","iopub.status.idle":"2023-12-17T23:20:50.085887Z","shell.execute_reply.started":"2023-12-17T23:20:50.078270Z","shell.execute_reply":"2023-12-17T23:20:50.084907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"augmentor = SimCLRAugmentor(224)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:54:51.378351Z","iopub.execute_input":"2023-12-17T23:54:51.378686Z","iopub.status.idle":"2023-12-17T23:54:51.383183Z","shell.execute_reply.started":"2023-12-17T23:54:51.378661Z","shell.execute_reply":"2023-12-17T23:54:51.382232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_images(images, classes, denormalize=False):\n    num_images = len(images)\n    num_cols = 3\n    num_rows = (num_images + num_cols - 1) // num_cols\n    \n    plt.figure(figsize=(15, 3 * num_rows))\n\n    for i, image in enumerate(images):        \n        if denormalize:\n            image = denorm(image)\n            \n        image = image.permute(1, 2, 0)   \n\n        plt.subplot(num_rows, num_cols, i + 1)\n\n        plt.imshow(image)\n        plt.title(f'Label: {classes[i]}')\n        plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:20:35.061665Z","iopub.execute_input":"2023-12-18T00:20:35.062574Z","iopub.status.idle":"2023-12-18T00:20:35.069089Z","shell.execute_reply.started":"2023-12-18T00:20:35.062533Z","shell.execute_reply":"2023-12-18T00:20:35.068105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:02:26.145580Z","iopub.execute_input":"2023-12-18T00:02:26.146137Z","iopub.status.idle":"2023-12-18T00:02:26.150638Z","shell.execute_reply.started":"2023-12-18T00:02:26.146106Z","shell.execute_reply":"2023-12-18T00:02:26.149641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss = []\nval_loss = []\n\nfor epoch in range(1, config[\"num_epochs\"] + 1):\n    model.train()\n    total_train_loss = 0\n\n    current_lr1 = optimizer.param_groups[0]['lr']\n\n    for (inputs, targets) in tqdm(train_loader, leave=False, desc=f'Epoch {epoch}, lr 1: {current_lr1}'):\n        x1, x2 = augmentor(inputs)\n        \n        x1, x2 = x1.to(config[\"device\"]), x2.to(config[\"device\"])\n\n        h1, z1 = model(x1)\n        h2, z2 = model(x2)\n        \n        loss = criterion(z1, z2)\n\n        loss.backward()\n\n        optimizer.step()\n        optimizer.zero_grad()\n\n        total_train_loss += loss.item()\n\n    avg_train_loss = total_train_loss / len(train_loader)\n    train_loss.append(avg_train_loss)\n\n    print(f\"Average train loss at {epoch} epoch: {avg_train_loss}\")\n\n    total_val_loss = 0\n\n    model.eval()\n    \n    for (inputs, targets) in tqdm(val_loader, leave=False, desc=f'Validation - Epoch {epoch}'):\n        inputs = inputs.to(config[\"device\"])\n        targets = targets.to(config[\"device\"])\n\n        with torch.no_grad():\n            h1, z1 = model(inputs)\n            h2, z2 = model(inputs)\n        \n            loss = criterion(z1, z2)\n\n        total_val_loss += loss.item()\n\n    avg_val_loss = total_val_loss / len(val_loader)\n    val_loss.append(avg_val_loss)\n\n    scheduler.step()\n\n    print(f\"Average test loss at {epoch} epoch: {avg_val_loss}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:53.647459Z","iopub.execute_input":"2023-12-17T23:20:53.648274Z","iopub.status.idle":"2023-12-17T23:32:07.044954Z","shell.execute_reply.started":"2023-12-17T23:20:53.648240Z","shell.execute_reply":"2023-12-17T23:32:07.043681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_losses(train_loss, val_loss)","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:19:02.238131Z","iopub.execute_input":"2023-12-17T23:19:02.238918Z","iopub.status.idle":"2023-12-17T23:19:02.537788Z","shell.execute_reply.started":"2023-12-17T23:19:02.238883Z","shell.execute_reply":"2023-12-17T23:19:02.536912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"self_supervised.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-12-17T23:20:02.892965Z","iopub.execute_input":"2023-12-17T23:20:02.893783Z","iopub.status.idle":"2023-12-17T23:20:03.021489Z","shell.execute_reply.started":"2023-12-17T23:20:02.893752Z","shell.execute_reply":"2023-12-17T23:20:03.020672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Trying FAISS backend","metadata":{}},{"cell_type":"code","source":"def extract_features(data_loader, model):\n    all_features = []\n    \n    model.eval()\n\n    with torch.no_grad():\n        for (inputs, targets) in tqdm(data_loader, leave=False):\n            features = model(inputs.to(config[\"device\"]))\n            features = features.detach().cpu()\n            features = features.view(features.size(0), -1).numpy()\n            all_features.append(features)\n        \n    return np.concatenate(all_features, axis=0)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:04.553120Z","iopub.execute_input":"2023-12-18T00:15:04.554223Z","iopub.status.idle":"2023-12-18T00:15:04.561778Z","shell.execute_reply.started":"2023-12-18T00:15:04.554180Z","shell.execute_reply":"2023-12-18T00:15:04.560816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = resnet50(pretrained=True).to(config[\"device\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:07.619448Z","iopub.execute_input":"2023-12-18T00:15:07.620194Z","iopub.status.idle":"2023-12-18T00:15:08.091584Z","shell.execute_reply.started":"2023-12-18T00:15:07.620161Z","shell.execute_reply":"2023-12-18T00:15:08.090768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.nn.Sequential(*(list(model.children())[:-1]))","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:09.941940Z","iopub.execute_input":"2023-12-18T00:15:09.942341Z","iopub.status.idle":"2023-12-18T00:15:09.947628Z","shell.execute_reply.started":"2023-12-18T00:15:09.942310Z","shell.execute_reply":"2023-12-18T00:15:09.946576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_df = train_df[:2000]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:11.251162Z","iopub.execute_input":"2023-12-18T00:15:11.251901Z","iopub.status.idle":"2023-12-18T00:15:11.256571Z","shell.execute_reply.started":"2023-12-18T00:15:11.251865Z","shell.execute_reply":"2023-12-18T00:15:11.255546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_dataset = MutlilabelClassificationDataset(example_df, transform=val_transform)\nexample_loader = DataLoader(example_dataset, batch_size=config[\"batch_size\"], num_workers=config[\"num_workers\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:12.729114Z","iopub.execute_input":"2023-12-18T00:15:12.729598Z","iopub.status.idle":"2023-12-18T00:15:12.735623Z","shell.execute_reply.started":"2023-12-18T00:15:12.729555Z","shell.execute_reply":"2023-12-18T00:15:12.734571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = extract_features(example_loader, model)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:15:13.789380Z","iopub.execute_input":"2023-12-18T00:15:13.789712Z","iopub.status.idle":"2023-12-18T00:16:08.088706Z","shell.execute_reply.started":"2023-12-18T00:15:13.789685Z","shell.execute_reply":"2023-12-18T00:16:08.087567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"index = faiss.IndexFlatL2(features.shape[1])\nindex.add(features)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:16:59.522357Z","iopub.execute_input":"2023-12-18T00:16:59.523165Z","iopub.status.idle":"2023-12-18T00:16:59.540275Z","shell.execute_reply.started":"2023-12-18T00:16:59.523129Z","shell.execute_reply":"2023-12-18T00:16:59.539511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example = train_df.iloc[2001]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:17:18.243314Z","iopub.execute_input":"2023-12-18T00:17:18.243664Z","iopub.status.idle":"2023-12-18T00:17:18.248625Z","shell.execute_reply.started":"2023-12-18T00:17:18.243636Z","shell.execute_reply":"2023-12-18T00:17:18.247722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_image_path = os.path.join(train_images_path, f\"{example['ImageId']}.jpg\")\nexample_image = PIL.Image.open(example_image_path).convert('RGB')\n\nexample_image = val_transform(example_image)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:28:36.505329Z","iopub.execute_input":"2023-12-18T00:28:36.505700Z","iopub.status.idle":"2023-12-18T00:28:36.522038Z","shell.execute_reply.started":"2023-12-18T00:28:36.505669Z","shell.execute_reply":"2023-12-18T00:28:36.520971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_image(example_image, example[\"ClassId\"], denormalize=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:18:21.066890Z","iopub.execute_input":"2023-12-18T00:18:21.067256Z","iopub.status.idle":"2023-12-18T00:18:21.268572Z","shell.execute_reply.started":"2023-12-18T00:18:21.067227Z","shell.execute_reply":"2023-12-18T00:18:21.267537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example_image = example_image.unsqueeze(0).to(config[\"device\"])\n\nmodel.eval()\n\nwith torch.no_grad():\n    query_features = model(example_image)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:18:22.947596Z","iopub.execute_input":"2023-12-18T00:18:22.947977Z","iopub.status.idle":"2023-12-18T00:18:22.997961Z","shell.execute_reply.started":"2023-12-18T00:18:22.947945Z","shell.execute_reply":"2023-12-18T00:18:22.997091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"query_features = query_features.detach().cpu()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:18:56.306797Z","iopub.execute_input":"2023-12-18T00:18:56.307180Z","iopub.status.idle":"2023-12-18T00:18:56.312219Z","shell.execute_reply.started":"2023-12-18T00:18:56.307151Z","shell.execute_reply":"2023-12-18T00:18:56.311060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"query_features = query_features.view(query_features.size(0), -1).numpy()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:18:57.329630Z","iopub.execute_input":"2023-12-18T00:18:57.330015Z","iopub.status.idle":"2023-12-18T00:18:57.334926Z","shell.execute_reply.started":"2023-12-18T00:18:57.329985Z","shell.execute_reply":"2023-12-18T00:18:57.333892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"k = 5\n\ndistances, indices = index.search(query_features, k)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:29:01.712855Z","iopub.execute_input":"2023-12-18T00:29:01.713232Z","iopub.status.idle":"2023-12-18T00:29:01.722650Z","shell.execute_reply.started":"2023-12-18T00:29:01.713203Z","shell.execute_reply":"2023-12-18T00:29:01.721675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"items = []\nlabels = []\n\nfor i in indices[0]:\n    items.append(example_dataset[i][0])\n    labels.append(example_dataset[i][1])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:29:04.021662Z","iopub.execute_input":"2023-12-18T00:29:04.022305Z","iopub.status.idle":"2023-12-18T00:29:04.250276Z","shell.execute_reply.started":"2023-12-18T00:29:04.022270Z","shell.execute_reply":"2023-12-18T00:29:04.249260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"matches = [torch.any(t == torch.tensor(example[\"EncodedClasses\"])).item() for t in labels]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:29:05.229108Z","iopub.execute_input":"2023-12-18T00:29:05.229788Z","iopub.status.idle":"2023-12-18T00:29:05.235249Z","shell.execute_reply.started":"2023-12-18T00:29:05.229756Z","shell.execute_reply":"2023-12-18T00:29:05.234303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.patches import Rectangle\n\ndef plot_images(images, classes, matches=None, denormalize=False):\n    num_images = len(images)\n    num_cols = 3\n    num_rows = (num_images + num_cols - 1) // num_cols\n    \n    plt.figure(figsize=(15, 3 * num_rows))\n\n    for i, (image, match) in enumerate(zip(images, matches)):\n        if denormalize:\n            image = denorm(image)\n        \n        image = image.permute(1, 2, 0)\n\n        plt.subplot(num_rows, num_cols, i + 1)\n\n        plt.imshow(image)\n        plt.title(f'Label: {classes[i]}')\n        plt.axis('off')\n        \n        # Add a frame with different colors for matches and non-matches\n        frame_color = 'green' if match else 'red'\n        rect = Rectangle((0, 0), 1, 1, linewidth=4, edgecolor=frame_color, facecolor='none', transform=plt.gca().transAxes)\n        plt.gca().add_patch(rect)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:33:41.353304Z","iopub.execute_input":"2023-12-18T00:33:41.353685Z","iopub.status.idle":"2023-12-18T00:33:41.362095Z","shell.execute_reply.started":"2023-12-18T00:33:41.353656Z","shell.execute_reply":"2023-12-18T00:33:41.361136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images([example_image] + items, [\"Base Image\", \"1\", \"2\", \"3\", \"4\", \"5\"], matches=[True] + matches, denormalize=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T00:33:59.964711Z","iopub.execute_input":"2023-12-18T00:33:59.965081Z","iopub.status.idle":"2023-12-18T00:34:00.675500Z","shell.execute_reply.started":"2023-12-18T00:33:59.965053Z","shell.execute_reply":"2023-12-18T00:34:00.674623Z"},"trusted":true},"execution_count":null,"outputs":[]}]}