{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install pytorch-metric-learning\n!pip install faiss-gpu\n!pip install torch_ema","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:00.155572Z","iopub.execute_input":"2022-12-19T18:20:00.156325Z","iopub.status.idle":"2022-12-19T18:20:34.436775Z","shell.execute_reply.started":"2022-12-19T18:20:00.156233Z","shell.execute_reply":"2022-12-19T18:20:34.435559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nimport math\nimport pandas as pd\nimport numpy as np\nimport logging\nfrom IPython.display import FileLink\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision.io import ImageReadMode, read_image\nfrom torchvision import transforms as T \n\nimport pytorch_metric_learning\nimport pytorch_metric_learning.utils.logging_presets as LP\nfrom pytorch_metric_learning.utils import common_functions\nfrom pytorch_metric_learning import losses, miners, samplers, testers, trainers\nfrom pytorch_metric_learning.utils.accuracy_calculator import AccuracyCalculator\nfrom pytorch_metric_learning.utils.inference import InferenceModel\nfrom pytorch_metric_learning.losses.arcface_loss import ArcFaceLoss\nfrom torchvision import models\n\nfrom torch_ema import ExponentialMovingAverage\n\nfor handler in logging.root.handlers[:]:\n    logging.root.removeHandler(handler)\n\nlogging.getLogger().setLevel(logging.INFO)\nlogging.info(\"VERSION %s\" % pytorch_metric_learning.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-19T18:20:34.440418Z","iopub.execute_input":"2022-12-19T18:20:34.440777Z","iopub.status.idle":"2022-12-19T18:20:35.855059Z","shell.execute_reply.started":"2022-12-19T18:20:34.440744Z","shell.execute_reply":"2022-12-19T18:20:35.854100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_CLASSES=5004\nEMBEDDING_SIZE = 512\nN_EPOCH=35\nBATCH_SIZE=128\nVAL_BATCH_SIZE = 1024\nPRINT_STEP = 1\nN_WORKER=2\n\ntorch.backends.cudnn.benchmark = True\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.856425Z","iopub.execute_input":"2022-12-19T18:20:35.857208Z","iopub.status.idle":"2022-12-19T18:20:35.863526Z","shell.execute_reply.started":"2022-12-19T18:20:35.857167Z","shell.execute_reply":"2022-12-19T18:20:35.862271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/humpback-whale-identification/train'\nTEST_DIR = '/kaggle/input/humpback-whale-identification/test'\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.866422Z","iopub.execute_input":"2022-12-19T18:20:35.866798Z","iopub.status.idle":"2022-12-19T18:20:35.875004Z","shell.execute_reply.started":"2022-12-19T18:20:35.866762Z","shell.execute_reply":"2022-12-19T18:20:35.874063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HumpbackDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        image_dir: str,\n        train = True\n    ):\n        self.df = df\n        self.image_dir = image_dir\n        \n        if train:\n            self.image_transform = T.Compose(\n                [\n                    \n                    T.ToPILImage(),\n                    T.RandomHorizontalFlip(p=0.5),\n                    T.RandomRotation(10),\n                    T.ColorJitter(brightness=0.05, contrast=0.05, saturation=0.05, hue=0.05),\n#                     T.AutoAugment(T.AutoAugmentPolicy.IMAGENET),\n                    T.Grayscale(1),\n                    T.Resize((224,224)),\n                    T.ToTensor(),\n                    T.Normalize(mean=[0.5],\n                                 std=[0.5]),\n                    \n                ]\n            )\n        else:\n            self.image_transform = T.Compose(\n                [\n                    T.ToPILImage(),\n                    T.Grayscale(1),\n                    T.Resize((224,224)),\n                    T.ToTensor(),\n                    T.Normalize(mean=[0.5],\n                                 std=[0.5]),\n\n                ]\n            )\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, idx):\n        \n        image_path = os.path.join(self.image_dir, self.df[\"Image\"].iloc[idx])\n        image = read_image(path=image_path)\n        \n        image = self.image_transform(image)\n            \n        label = self.df['Id_int'].iloc[idx]\n\n        return image.type(torch.FloatTensor), torch.tensor(label).type(torch.LongTensor)\n\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.876519Z","iopub.execute_input":"2022-12-19T18:20:35.876878Z","iopub.status.idle":"2022-12-19T18:20:35.888690Z","shell.execute_reply.started":"2022-12-19T18:20:35.876845Z","shell.execute_reply":"2022-12-19T18:20:35.887785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, loss_func, device, train_loader, optimizer, epoch, ema=None, scheduler = None):\n    model.train()\n    pbar = tqdm(enumerate(train_loader))\n    for batch_idx, (data, labels) in pbar:\n        data, labels = data.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        embeddings = model(data)\n        loss = loss_func(embeddings, labels)\n        loss.backward()\n        optimizer.step()\n        \n        if ema:\n            ema.update()\n        if batch_idx % PRINT_STEP == 0:\n            pbar.set_description(f\"Epoch {epoch} Iteration {batch_idx}: Loss = {loss:.6f} Lr: {scheduler.get_last_lr()[0]:.6f}\")\n    if scheduler:\n        scheduler.step()\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T20:35:00.741480Z","iopub.execute_input":"2022-12-19T20:35:00.741859Z","iopub.status.idle":"2022-12-19T20:35:00.749706Z","shell.execute_reply.started":"2022-12-19T20:35:00.741828Z","shell.execute_reply":"2022-12-19T20:35:00.748772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_all_embeddings(dataset, model):\n    tester = testers.BaseTester()\n    return tester.get_all_embeddings(dataset, model)\n\n### compute accuracy using AccuracyCalculator from pytorch-metric-learning ###\ndef test(train_set, test_set, model, accuracy_calculator, loss_func):\n    train_embeddings, train_labels = get_all_embeddings(train_set, model)\n    test_embeddings, test_labels = get_all_embeddings(test_set, model)\n    train_labels = train_labels.squeeze(1)\n    test_labels = test_labels.squeeze(1)\n    train_loss = loss_func(train_embeddings, train_labels)\n    test_loss = loss_func(test_embeddings, test_labels)\n    print(\"Computing accuracy\")\n    accuracies = accuracy_calculator.get_accuracy(\n        test_embeddings, train_embeddings, test_labels, train_labels, False)\n    print(\"Train loss = {}\".format(train_loss))\n    print(\"Test loss = {}\".format(test_loss))\n    print(\"Test set accuracy (Precision@1) = {}\".format(accuracies[\"precision_at_1\"]))\n    print(\"Test set accuracy (mean_average_precision) = {}\".format(accuracies[\"mean_average_precision\"]))\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.904083Z","iopub.execute_input":"2022-12-19T18:20:35.904445Z","iopub.status.idle":"2022-12-19T18:20:35.917173Z","shell.execute_reply.started":"2022-12-19T18:20:35.904393Z","shell.execute_reply":"2022-12-19T18:20:35.916130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LiArcFace(ArcFaceLoss):\n    def __init__(self, *args, margin=0.5, scale=12, **kwargs):\n        super().__init__(*args, margin=margin, scale=scale, **kwargs)\n    \n    def modify_cosine_of_target_classes(self, cosine_of_target_classes):\n        angles = self.get_angles(cosine_of_target_classes)\n        return (torch.pi - 2 * (angles + self.margin)) / torch.pi\n    ","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.918709Z","iopub.execute_input":"2022-12-19T18:20:35.919216Z","iopub.status.idle":"2022-12-19T18:20:35.928479Z","shell.execute_reply.started":"2022-12-19T18:20:35.919181Z","shell.execute_reply":"2022-12-19T18:20:35.927222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(torch.nn.Module):\n    def __init__(self, model_name = 'resnet18'):\n        super(Net, self).__init__()\n        model = getattr(models, model_name)(pretrained=True)\n        self.model = torch.nn.Sequential(\n            nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False),\n            *(list(model.children())[1:-1]))\n\n    def forward(self, x):\n        out = self.model(x)\n        return out.squeeze()\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.929973Z","iopub.execute_input":"2022-12-19T18:20:35.930305Z","iopub.status.idle":"2022-12-19T18:20:35.942468Z","shell.execute_reply.started":"2022-12-19T18:20:35.930272Z","shell.execute_reply":"2022-12-19T18:20:35.941550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train/val split\n\ndf = pd.read_csv('/kaggle/input/humpback-whale-identification/train.csv')\ndf = df.drop(df[df.Id == 'new_whale'].index)\n# print(df)\nfac = df['Id'].factorize()\ndf[\"Id_int\"] = fac[0]\n\ntrain_df = df\nvalid_df = pd.DataFrame({'Image' : [], 'Id': [], 'Id_int': []})\nval_part = 0.05\n\nfor k, Id in enumerate(fac[1]):\n    \n    sz = df[df['Id'] == fac[1][k]].count()['Id']\n    if sz > 1:\n        n_samples = math.ceil(val_part*sz)\n        # rows global indeces for spicific class \n        idex_ids = df.loc[train_df['Id'] == fac[1][k]].index\n        sampled_idcs = np.random.choice(idex_ids, n_samples, replace = False)\n        to_val = train_df.loc[sampled_idcs]\n        train_df.drop(sampled_idcs, inplace=True)\n        valid_df = valid_df.append(to_val)\n\nprint(train_df.shape)\nprint(valid_df.shape)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:35.947057Z","iopub.execute_input":"2022-12-19T18:20:35.947401Z","iopub.status.idle":"2022-12-19T18:20:58.149320Z","shell.execute_reply.started":"2022-12-19T18:20:35.947376Z","shell.execute_reply":"2022-12-19T18:20:58.147735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train over-sampling\n\nclass_weights = train_df.groupby('Id_int')['Image'].nunique()\nsample_weights = [0] * train_df.shape[0]\nfor idx, (_, row) in enumerate(train_df.iterrows()):\n\n    sample_weights[idx] = class_weights[int(row['Id_int'])]\n\nsampler = WeightedRandomSampler(sample_weights,\n                                num_samples=len(sample_weights),\n                                replacement=True)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:58.150617Z","iopub.execute_input":"2022-12-19T18:20:58.151463Z","iopub.status.idle":"2022-12-19T18:20:58.817094Z","shell.execute_reply.started":"2022-12-19T18:20:58.151423Z","shell.execute_reply":"2022-12-19T18:20:58.816106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from sklearn.model_selection import train_test_split\n# df = pd.read_csv('/kaggle/input/humpback-whale-identification/train.csv')\n# df = df.drop(df[df.Id ==  'new_whale'].index)\n# train_df, valid_df = train_test_split(df, test_size=0.1)\n\ndevice = torch.device(\"cuda:0\")\n\ntrain_dataset = HumpbackDataset(df=train_df, image_dir=TRAIN_DIR, train=True)\nvalid_dataset = HumpbackDataset(df=valid_df, image_dir=TRAIN_DIR, train=False)\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, \n                              num_workers=N_WORKER,\n                              sampler=sampler)\nvalid_dataloader = DataLoader(valid_dataset, batch_size=VAL_BATCH_SIZE, shuffle=False, \n                              num_workers=N_WORKER)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-19T18:20:58.818457Z","iopub.execute_input":"2022-12-19T18:20:58.819080Z","iopub.status.idle":"2022-12-19T18:20:58.827407Z","shell.execute_reply.started":"2022-12-19T18:20:58.819042Z","shell.execute_reply":"2022-12-19T18:20:58.826183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Net('resnet18').to(device)\nema = ExponentialMovingAverage(model.parameters(), decay=0.995)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func = LiArcFace(num_classes=N_CLASSES, embedding_size=EMBEDDING_SIZE, scale=13, margin=0.5).to(device)\naccuracy_calculator = AccuracyCalculator(include=(\"precision_at_1\", \"mean_average_precision\", ), k=1)\n\n# optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9, weight_decay=1e-5)\n# lambda1 = lambda epoch: 0.8 ** epoch\n# scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda1)\n# scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=0.7)\n\noptimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n# scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=len(train_dataloader)*N_EPOCH, eta_min=1e-6)\nscheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[2,6,11,15], gamma=0.1)\n\n\n\nN_EPOCH=20\n\nfor epoch in range(1, N_EPOCH + 1):\n    train(model, loss_func, device, train_dataloader, optimizer, epoch, ema, scheduler)\n    if epoch%2 == 0:\n        print('Val EMA')\n        with ema.average_parameters():\n            test(train_dataset, valid_dataset, model, accuracy_calculator, loss_func)\n#         if epoch%5 == 0 or epoch == N_EPOCH: \n#             print('Val Default')\n#             test(train_dataset, valid_dataset, model, accuracy_calculator)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-12-19T21:52:06.035314Z","iopub.execute_input":"2022-12-19T21:52:06.036106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### save first model\nema.copy_to()\nmodel.eval()\ntorch.save(model.state_dict(), 'LiArc_rnet18.pth')\n","metadata":{"execution":{"iopub.status.busy":"2022-12-15T11:16:00.409563Z","iopub.execute_input":"2022-12-15T11:16:00.409935Z","iopub.status.idle":"2022-12-15T11:16:00.583255Z","shell.execute_reply.started":"2022-12-15T11:16:00.409904Z","shell.execute_reply":"2022-12-15T11:16:00.582245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ### train model with ArcFace loss \n# net = Net()\n# model = torch.nn.DataParallel(net, device_ids=[0, 1])\n# ### pytorch-metric-learning stuff ###\n# loss_func = ArcFace(num_classes=N_CLASSES, embedding_size=EMBEDDING_SIZE).to(device)\n# optimizer = optim.SGD(model.parameters(), lr=1e-5, momentum=0.9)\n# scheduler = torch.optim.lr_scheduler.CyclicLR(optimizer, base_lr=0.00001, max_lr=0.1,step_size_up=2,\n#                                               step_size_down=7, mode=\"exp_range\",gamma=0.88,cycle_momentum=False)\n# accuracy_calculator = AccuracyCalculator(include=(\"precision_at_1\",), k=1)","metadata":{"execution":{"iopub.status.busy":"2022-12-15T11:13:15.552816Z","iopub.status.idle":"2022-12-15T11:13:15.553637Z","shell.execute_reply.started":"2022-12-15T11:13:15.553368Z","shell.execute_reply":"2022-12-15T11:13:15.553393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for epoch in range(1, N_EPOCH + 1):\n#     train(model, loss_func, device, train_dataloader, optimizer, loss_optimizer, epoch)\n#     test(train_dataset, valid_dataset, model, accuracy_calculator)","metadata":{"execution":{"iopub.status.busy":"2022-12-15T11:13:15.555125Z","iopub.status.idle":"2022-12-15T11:13:15.555925Z","shell.execute_reply.started":"2022-12-15T11:13:15.555626Z","shell.execute_reply":"2022-12-15T11:13:15.555651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ### save second model\n# model.eval()\n# torch.save(model.state_dict(), 'Arc_rnet50.pth')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# os.chdir(r'/kaggle/working')\n# !zip out.zip /kaggle/working/test.pth\n# FileLink(r'out.zip')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}