{"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":"# 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)\nimport matplotlib.pyplot as plt\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 gc\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":"2022-10-06T12:12:21.548873Z","iopub.execute_input":"2022-10-06T12:12:21.549355Z","iopub.status.idle":"2022-10-06T12:12:21.556082Z","shell.execute_reply.started":"2022-10-06T12:12:21.549300Z","shell.execute_reply":"2022-10-06T12:12:21.555137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## PyTorch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.data as data\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\n## Torchvision\nimport torchvision\nfrom torchvision import transforms\n\n## Augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\nimport cv2\n\n## Utils\nimport joblib\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.571230Z","iopub.execute_input":"2022-10-06T12:12:21.571529Z","iopub.status.idle":"2022-10-06T12:12:21.586330Z","shell.execute_reply.started":"2022-10-06T12:12:21.571501Z","shell.execute_reply":"2022-10-06T12:12:21.585153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = {\n    'img_size': 448,\n    'seed': 22,\n    'n_fold': 5,\n    'train_batch_size': 32,\n    'test_batch_size': 64,\n    'num_classes': 15587,\n    'patches_size': 32,\n    'device': torch.device('cuda:0' if torch.cuda.is_available() else 'cpu'),\n}","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.589026Z","iopub.execute_input":"2022-10-06T12:12:21.589709Z","iopub.status.idle":"2022-10-06T12:12:21.662440Z","shell.execute_reply.started":"2022-10-06T12:12:21.589672Z","shell.execute_reply":"2022-10-06T12:12:21.661330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.664343Z","iopub.execute_input":"2022-10-06T12:12:21.665106Z","iopub.status.idle":"2022-10-06T12:12:21.675042Z","shell.execute_reply.started":"2022-10-06T12:12:21.665056Z","shell.execute_reply":"2022-10-06T12:12:21.674117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '../input/happy-whale-and-dolphin'\nTRAIN_DIR = '../input/happy-whale-and-dolphin/train_images'\nTEST_DIR = '../input/happy-whale-and-dolphin/test_images'","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.678232Z","iopub.execute_input":"2022-10-06T12:12:21.678979Z","iopub.status.idle":"2022-10-06T12:12:21.685307Z","shell.execute_reply.started":"2022-10-06T12:12:21.678944Z","shell.execute_reply":"2022-10-06T12:12:21.684273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(id):\n    return f\"{TRAIN_DIR}/{id}\"","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.687036Z","iopub.execute_input":"2022-10-06T12:12:21.687901Z","iopub.status.idle":"2022-10-06T12:12:21.693205Z","shell.execute_reply.started":"2022-10-06T12:12:21.687861Z","shell.execute_reply":"2022-10-06T12:12:21.692198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f\"{ROOT_DIR}/train.csv\")\ndf['file_path'] = df['image'].apply(get_train_file_path)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.694806Z","iopub.execute_input":"2022-10-06T12:12:21.695648Z","iopub.status.idle":"2022-10-06T12:12:21.873954Z","shell.execute_reply.started":"2022-10-06T12:12:21.695612Z","shell.execute_reply":"2022-10-06T12:12:21.872884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = LabelEncoder()\ndf['individual_id'] = encoder.fit_transform(df['individual_id'])\n\nwith open(\"le.pkl\", \"wb\") as fp:\n    joblib.dump(encoder, fp)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.875620Z","iopub.execute_input":"2022-10-06T12:12:21.876325Z","iopub.status.idle":"2022-10-06T12:12:21.946977Z","shell.execute_reply.started":"2022-10-06T12:12:21.876270Z","shell.execute_reply":"2022-10-06T12:12:21.946131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.948534Z","iopub.execute_input":"2022-10-06T12:12:21.949208Z","iopub.status.idle":"2022-10-06T12:12:21.962692Z","shell.execute_reply.started":"2022-10-06T12:12:21.949168Z","shell.execute_reply":"2022-10-06T12:12:21.961723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HappyWhaleDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df['individual_id'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.file_names[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        label = self.labels[index]\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return {\n            'image': img,\n            'label': torch.tensor(label, dtype=torch.long)\n        }","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.966320Z","iopub.execute_input":"2022-10-06T12:12:21.966969Z","iopub.status.idle":"2022-10-06T12:12:21.975264Z","shell.execute_reply.started":"2022-10-06T12:12:21.966934Z","shell.execute_reply":"2022-10-06T12:12:21.974279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.ShiftScaleRotate(shift_limit=0.1, \n                           scale_limit=0.15, \n                           rotate_limit=60, \n                           p=0.5),\n        A.HueSaturationValue(\n                hue_shift_limit=0.2, \n                sat_shift_limit=0.2, \n                val_shift_limit=0.2, \n                p=0.5\n            ),\n        A.RandomBrightnessContrast(\n                brightness_limit=(-0.1,0.1), \n                contrast_limit=(-0.1, 0.1), \n                p=0.5\n            ),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.),\n    \n    \"test\": A.Compose([\n        A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(\n                mean=[0.485, 0.456, 0.406], \n                std=[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:21.977016Z","iopub.execute_input":"2022-10-06T12:12:21.977730Z","iopub.status.idle":"2022-10-06T12:12:21.988663Z","shell.execute_reply.started":"2022-10-06T12:12:21.977695Z","shell.execute_reply":"2022-10-06T12:12:21.987626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def img_to_patch(x, patch_size, flatten_channels=True):\n    \"\"\"\n    Inputs:\n        x - torch.Tensor representing the image of shape [B, C, H, W]\n        patch_size - Number of pixels per dimension of the patches (integer)\n        flatten_channels - If True, the patches will be returned in a flattened format\n                           as a feature vector instead of a image grid.\n    \"\"\"\n    B, C, H, W = x.shape\n    x = x.reshape(B, C, H//patch_size, patch_size, W//patch_size, patch_size)\n    x = x.permute(0, 2, 4, 1, 3, 5) # [B, H', W', C, p_H, p_W]\n    x = x.flatten(1,2)              # [B, H'*W', C, p_H, p_W]\n    if flatten_channels:\n        x = x.flatten(2,4)          # [B, H'*W', C*p_H*p_W]\n    return x","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:22.005355Z","iopub.execute_input":"2022-10-06T12:12:22.005732Z","iopub.status.idle":"2022-10-06T12:12:22.016484Z","shell.execute_reply.started":"2022-10-06T12:12:22.005695Z","shell.execute_reply":"2022-10-06T12:12:22.015506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n\n    def __init__(self, embed_dim, hidden_dim, num_heads, dropout=0.0):\n        \"\"\"\n        Inputs:\n            embed_dim - Dimensionality of input and attention feature vectors\n            hidden_dim - Dimensionality of hidden layer in feed-forward network\n                         (usually 2-4x larger than embed_dim)\n            num_heads - Number of heads to use in the Multi-Head Attention block\n            dropout - Amount of dropout to apply in the feed-forward network\n        \"\"\"\n        super().__init__()\n\n        self.layer_norm_1 = nn.LayerNorm(embed_dim)\n        self.attn = nn.MultiheadAttention(embed_dim, num_heads,\n                                          dropout=dropout)\n        self.layer_norm_2 = nn.LayerNorm(embed_dim)\n        self.linear = nn.Sequential(\n            nn.Linear(embed_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, embed_dim),\n            nn.Dropout(dropout)\n        )\n\n\n    def forward(self, x):\n        inp_x = self.layer_norm_1(x)\n        x = x + self.attn(inp_x, inp_x, inp_x)[0]\n        x = x + self.linear(self.layer_norm_2(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:22.039813Z","iopub.execute_input":"2022-10-06T12:12:22.040340Z","iopub.status.idle":"2022-10-06T12:12:22.050505Z","shell.execute_reply.started":"2022-10-06T12:12:22.040305Z","shell.execute_reply":"2022-10-06T12:12:22.049519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VisionTransformer(nn.Module):\n\n    def __init__(self, embed_dim, hidden_dim, num_channels, num_heads, num_layers, num_classes, patch_size, num_patches, dropout=0.0):\n        \"\"\"\n        Inputs:\n            embed_dim - Dimensionality of the input feature vectors to the Transformer\n            hidden_dim - Dimensionality of the hidden layer in the feed-forward networks\n                         within the Transformer\n            num_channels - Number of channels of the input (3 for RGB)\n            num_heads - Number of heads to use in the Multi-Head Attention block\n            num_layers - Number of layers to use in the Transformer\n            num_classes - Number of classes to predict\n            patch_size - Number of pixels that the patches have per dimension\n            num_patches - Maximum number of patches an image can have\n            dropout - Amount of dropout to apply in the feed-forward network and\n                      on the input encoding\n        \"\"\"\n        super().__init__()\n\n        self.patch_size = patch_size\n\n        # Layers/Networks\n        self.input_layer = nn.Linear(num_channels*(patch_size**2), embed_dim)\n        self.transformer = nn.Sequential(*[AttentionBlock(embed_dim, hidden_dim, num_heads, dropout=dropout) for _ in range(num_layers)])\n        self.mlp_head = nn.Sequential(\n            nn.LayerNorm(embed_dim),\n            nn.Linear(embed_dim, num_classes)\n        )\n        self.dropout = nn.Dropout(dropout)\n\n        # Parameters/Embeddings\n        self.cls_token = nn.Parameter(torch.randn(1,1,embed_dim))\n        self.pos_embedding = nn.Parameter(torch.randn(1,1+num_patches,embed_dim))\n\n\n    def forward(self, x):\n        # Preprocess input\n        x = img_to_patch(x, self.patch_size)\n        B, T, _ = x.shape\n        x = self.input_layer(x)\n\n        # Add CLS token and positional encoding\n        cls_token = self.cls_token.repeat(B, 1, 1)\n        x = torch.cat([cls_token, x], dim=1)\n        x = x + self.pos_embedding[:,:T+1]\n\n        # Apply Transforrmer\n        x = self.dropout(x)\n        x = x.transpose(0, 1)\n        x = self.transformer(x)\n\n        # Perform classification prediction\n        cls = x[0]\n        out = self.mlp_head(cls)\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:22.052376Z","iopub.execute_input":"2022-10-06T12:12:22.053147Z","iopub.status.idle":"2022-10-06T12:12:22.069805Z","shell.execute_reply.started":"2022-10-06T12:12:22.053108Z","shell.execute_reply":"2022-10-06T12:12:22.068828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG['device']","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:22.071423Z","iopub.execute_input":"2022-10-06T12:12:22.072302Z","iopub.status.idle":"2022-10-06T12:12:22.082863Z","shell.execute_reply.started":"2022-10-06T12:12:22.072266Z","shell.execute_reply":"2022-10-06T12:12:22.081765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = VisionTransformer(**{\n        'embed_dim': 784,\n        'hidden_dim': 1568,\n        'num_heads': 8,\n        'num_layers': 6,\n        'patch_size': 32,\n        'num_channels': 3,\n        'num_patches': 196,\n        'num_classes': CONFIG['num_classes'],\n        'dropout': 0.2\n    }\n)\nmodel.to(CONFIG['device'])","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:22.084411Z","iopub.execute_input":"2022-10-06T12:12:22.084930Z","iopub.status.idle":"2022-10-06T12:12:26.422931Z","shell.execute_reply.started":"2022-10-06T12:12:22.084895Z","shell.execute_reply":"2022-10-06T12:12:26.421804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = HappyWhaleDataset(df, transforms=data_transforms[\"train\"])\ntrain_loader = DataLoader(train_dataset, batch_size=CONFIG['train_batch_size'], num_workers=2, shuffle=True, drop_last=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:26.427613Z","iopub.execute_input":"2022-10-06T12:12:26.427998Z","iopub.status.idle":"2022-10-06T12:12:26.437640Z","shell.execute_reply.started":"2022-10-06T12:12:26.427963Z","shell.execute_reply":"2022-10-06T12:12:26.436418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# criterion + optimizer + scheduler\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=3e-5)\nlr_scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=[8, 12], gamma=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:26.442586Z","iopub.execute_input":"2022-10-06T12:12:26.442937Z","iopub.status.idle":"2022-10-06T12:12:26.449991Z","shell.execute_reply.started":"2022-10-06T12:12:26.442903Z","shell.execute_reply":"2022-10-06T12:12:26.448899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_loss = []","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:26.459160Z","iopub.execute_input":"2022-10-06T12:12:26.459572Z","iopub.status.idle":"2022-10-06T12:12:26.465669Z","shell.execute_reply.started":"2022-10-06T12:12:26.459480Z","shell.execute_reply":"2022-10-06T12:12:26.464635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(path):\n    model = VisionTransformer(**{\n            'embed_dim': 784,\n            'hidden_dim': 1568,\n            'num_heads': 8,\n            'num_layers': 6,\n            'patch_size': 32,\n            'num_channels': 3,\n            'num_patches': 196,\n            'num_classes': CONFIG['num_classes'],\n            'dropout': 0.2\n        }\n    )\n    model.to(CONFIG['device'])\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=3e-5)\n    lr_scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=[8, 12], gamma=0.1)\n\n    checkpoint = torch.load(path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    lr_scheduler.load_state_dict(checkpoint['lr_scheduler_state_dict'])\n    epoch = checkpoint['epoch']\n    training_loss = checkpoint['train_loss']\n    return model, optimizer, criterion, lr_scheduler, epoch, training_loss    ","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:26.467357Z","iopub.execute_input":"2022-10-06T12:12:26.468799Z","iopub.status.idle":"2022-10-06T12:12:26.481836Z","shell.execute_reply.started":"2022-10-06T12:12:26.468763Z","shell.execute_reply":"2022-10-06T12:12:26.480734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model, optimizer, criterion, lr_scheduler, current_epoch, training_loss = load_model('model-e15.pt')","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:12:29.791673Z","iopub.execute_input":"2022-10-06T12:12:29.792040Z","iopub.status.idle":"2022-10-06T12:12:30.791435Z","shell.execute_reply.started":"2022-10-06T12:12:29.792008Z","shell.execute_reply":"2022-10-06T12:12:30.790451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_epoch = current_epoch\nend_epoch = start_epoch + 3\nfor epoch in range(start_epoch, end_epoch):\n    train_epoch_loss = 0\n    gc.collect()\n    model.train()\n    bar = tqdm(enumerate(train_loader), total=len(train_loader))\n    for step, batch in bar:\n        X_train_batch, y_train_batch = batch['image'], batch['label']\n        X_train_batch = X_train_batch.to(CONFIG['device'])\n        y_train_batch = y_train_batch.to(CONFIG['device'])\n\n        optimizer.zero_grad()\n\n        y_train_pred = model(X_train_batch).squeeze()\n\n        train_loss = criterion(y_train_pred, y_train_batch)\n        #train_acc = multi_acc(y_train_pred, y_train_batch)\n\n        train_loss.backward()\n        optimizer.step()\n        #lr_scheduler.step()\n\n        train_epoch_loss += train_loss.item()\n        bar.set_postfix(Epoch=epoch, Train_Loss=train_epoch_loss/len(train_loader),LR=optimizer.param_groups[0]['lr']) \n\n    training_loss.append(train_epoch_loss/len(train_loader))\n    gc.collect()\n    \n\n#     print(f'Epoch {epoch+0:02}: | Train Loss: {train_epoch_loss/len(train_loader):.5f}')","metadata":{"execution":{"iopub.status.busy":"2022-10-06T00:40:07.083795Z","iopub.execute_input":"2022-10-06T00:40:07.084364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(training_loss)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T14:05:56.050656Z","iopub.execute_input":"2022-10-06T14:05:56.051034Z","iopub.status.idle":"2022-10-06T14:05:56.245968Z","shell.execute_reply.started":"2022-10-06T14:05:56.051000Z","shell.execute_reply":"2022-10-06T14:05:56.244986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_model(path):\n    torch.save({\n                'epoch': 15,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'lr_scheduler_state_dict': lr_scheduler.state_dict(),\n                'train_loss': training_loss,\n                }, path)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_model('model-e15.pt')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd /kaggle/working\nfrom IPython.display import FileLink\nFileLink(r'model-e15.pt')","metadata":{"execution":{"iopub.status.busy":"2022-10-06T00:34:27.200356Z","iopub.execute_input":"2022-10-06T00:34:27.200810Z","iopub.status.idle":"2022-10-06T00:34:27.216462Z","shell.execute_reply.started":"2022-10-06T00:34:27.200779Z","shell.execute_reply":"2022-10-06T00:34:27.214769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = []\nfor dirname, _, filenames in os.walk('/kaggle/input/happy-whale-and-dolphin/test_images'):\n    for filename in filenames:\n        path = os.path.join(dirname, filename)\n        test_data.append([filename, path])","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:01:02.589742Z","iopub.execute_input":"2022-10-06T12:01:02.590538Z","iopub.status.idle":"2022-10-06T12:01:31.048935Z","shell.execute_reply.started":"2022-10-06T12:01:02.590497Z","shell.execute_reply":"2022-10-06T12:01:31.047686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = pd.DataFrame(test_data, columns=['filename', 'path'])","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:03:40.114822Z","iopub.execute_input":"2022-10-06T12:03:40.115311Z","iopub.status.idle":"2022-10-06T12:03:40.128822Z","shell.execute_reply.started":"2022-10-06T12:03:40.115276Z","shell.execute_reply":"2022-10-06T12:03:40.127556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:03:47.167810Z","iopub.execute_input":"2022-10-06T12:03:47.168954Z","iopub.status.idle":"2022-10-06T12:03:47.183765Z","shell.execute_reply.started":"2022-10-06T12:03:47.168908Z","shell.execute_reply":"2022-10-06T12:03:47.182446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HappyWhaleTestDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.file_names = df['path'].values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.file_names[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        if self.transforms:\n            img = self.transforms(image=img)[\"image\"]\n            \n        return {\n            'image': img\n        }","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:10:08.886839Z","iopub.execute_input":"2022-10-06T12:10:08.887302Z","iopub.status.idle":"2022-10-06T12:10:08.896271Z","shell.execute_reply.started":"2022-10-06T12:10:08.887268Z","shell.execute_reply":"2022-10-06T12:10:08.895004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = HappyWhaleTestDataset(df_test, transforms=data_transforms['test'])\ntest_loader = DataLoader(test_dataset, batch_size=CONFIG['test_batch_size'], num_workers=2, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:10:29.878757Z","iopub.execute_input":"2022-10-06T12:10:29.879162Z","iopub.status.idle":"2022-10-06T12:10:29.885493Z","shell.execute_reply.started":"2022-10-06T12:10:29.879127Z","shell.execute_reply":"2022-10-06T12:10:29.884257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred_list = []","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:17:34.402939Z","iopub.execute_input":"2022-10-06T13:17:34.404343Z","iopub.status.idle":"2022-10-06T13:17:34.411163Z","shell.execute_reply.started":"2022-10-06T13:17:34.404280Z","shell.execute_reply":"2022-10-06T13:17:34.410029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import itertools\ny_pred_list_p = list(itertools.chain(*y_pred_list))","metadata":{"execution":{"iopub.status.busy":"2022-10-06T12:43:56.162681Z","iopub.execute_input":"2022-10-06T12:43:56.163036Z","iopub.status.idle":"2022-10-06T12:43:56.169991Z","shell.execute_reply.started":"2022-10-06T12:43:56.163007Z","shell.execute_reply":"2022-10-06T12:43:56.169042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    bar = tqdm(enumerate(test_loader), total=len(test_loader))\n    for step, batch in bar:\n        x_batch = batch['image']\n        x_batch = x_batch.to(CONFIG['device'])\n        y_test_pred = model(x_batch)\n        y_test_pred = torch.softmax(y_test_pred, dim = 1)\n        y_pred_probs, y_pred_tags = torch.topk(y_test_pred, 5, dim = 1)\n        y_pred_probs = y_pred_probs.cpu().numpy()\n        y_pred_tags = y_pred_tags.cpu().numpy()\n        \n        # threshold for new_individual\n        y_pred_tags[:, -1][y_pred_probs[:, -1] < 0.7] = -1 # new_individual     \n        \n        y_pred_list.append(y_pred_tags)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:17:36.689433Z","iopub.execute_input":"2022-10-06T13:17:36.689794Z","iopub.status.idle":"2022-10-06T13:45:16.513553Z","shell.execute_reply.started":"2022-10-06T13:17:36.689763Z","shell.execute_reply":"2022-10-06T13:45:16.511946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred_list[:2]","metadata":{"execution":{"iopub.status.busy":"2022-10-06T14:06:37.335779Z","iopub.execute_input":"2022-10-06T14:06:37.336247Z","iopub.status.idle":"2022-10-06T14:06:37.352815Z","shell.execute_reply.started":"2022-10-06T14:06:37.336208Z","shell.execute_reply":"2022-10-06T14:06:37.352002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions_col = []","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:02.120035Z","iopub.execute_input":"2022-10-06T13:56:02.120745Z","iopub.status.idle":"2022-10-06T13:56:02.125452Z","shell.execute_reply.started":"2022-10-06T13:56:02.120708Z","shell.execute_reply":"2022-10-06T13:56:02.124442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for y_pred_batch in y_pred_list:\n    for y_pred_top_5 in y_pred_batch:\n        y_pred_string = f'{encoder.inverse_transform([y_pred_top_5[0]])[0]}'\n        for y_pred in y_pred_top_5[1:]:\n            if y_pred == -1:\n                y_pred_string = y_pred_string + '\\nnew_individual'\n            else:\n                y_pred_string = y_pred_string + f'\\n{encoder.inverse_transform([y_pred])[0]}'\n        predictions_col.append(y_pred_string)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:04.789909Z","iopub.execute_input":"2022-10-06T13:56:04.790880Z","iopub.status.idle":"2022-10-06T13:56:39.695446Z","shell.execute_reply.started":"2022-10-06T13:56:04.790841Z","shell.execute_reply":"2022-10-06T13:56:39.694393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions_col[:3]","metadata":{"execution":{"iopub.status.busy":"2022-10-06T14:06:55.420318Z","iopub.execute_input":"2022-10-06T14:06:55.420681Z","iopub.status.idle":"2022-10-06T14:06:55.427046Z","shell.execute_reply.started":"2022-10-06T14:06:55.420651Z","shell.execute_reply":"2022-10-06T14:06:55.426132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(predictions_col)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:44.309462Z","iopub.execute_input":"2022-10-06T13:56:44.309813Z","iopub.status.idle":"2022-10-06T13:56:44.318480Z","shell.execute_reply.started":"2022-10-06T13:56:44.309784Z","shell.execute_reply":"2022-10-06T13:56:44.317353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = df_test.copy()","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:48.035918Z","iopub.execute_input":"2022-10-06T13:56:48.036613Z","iopub.status.idle":"2022-10-06T13:56:48.045989Z","shell.execute_reply.started":"2022-10-06T13:56:48.036575Z","shell.execute_reply":"2022-10-06T13:56:48.044875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.rename(columns={\"filename\": \"image\"}, inplace=True)\ndf_submission.drop(columns=['path'], inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:53.916802Z","iopub.execute_input":"2022-10-06T13:56:53.917199Z","iopub.status.idle":"2022-10-06T13:56:53.926375Z","shell.execute_reply.started":"2022-10-06T13:56:53.917165Z","shell.execute_reply":"2022-10-06T13:56:53.925311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission['predictions'] = predictions_col","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:56:55.287816Z","iopub.execute_input":"2022-10-06T13:56:55.288210Z","iopub.status.idle":"2022-10-06T13:56:55.297741Z","shell.execute_reply.started":"2022-10-06T13:56:55.288174Z","shell.execute_reply":"2022-10-06T13:56:55.296921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.to_csv('./submission_topk.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-06T14:00:25.645778Z","iopub.execute_input":"2022-10-06T14:00:25.646174Z","iopub.status.idle":"2022-10-06T14:00:25.706165Z","shell.execute_reply.started":"2022-10-06T14:00:25.646139Z","shell.execute_reply":"2022-10-06T14:00:25.705272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission","metadata":{"execution":{"iopub.status.busy":"2022-10-06T13:57:16.197129Z","iopub.execute_input":"2022-10-06T13:57:16.197605Z","iopub.status.idle":"2022-10-06T13:57:16.219417Z","shell.execute_reply.started":"2022-10-06T13:57:16.197563Z","shell.execute_reply":"2022-10-06T13:57:16.218592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# References\n- https://www.kaggle.com/code/debarshichanda/pytorch-arcface-gem-pooling-starter\n- https://www.kaggle.com/code/jaykumar2862/happy-whale-competition\n- https://uvadlc-notebooks.readthedocs.io/en/latest/tutorial_notebooks/tutorial6/Transformers_and_MHAttention.html\n- https://uvadlc-notebooks.readthedocs.io/en/latest/tutorial_notebooks/tutorial15/Vision_Transformer.html","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}