{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision.models import vgg16\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-02-07T11:52:48.437535Z","iopub.execute_input":"2022-02-07T11:52:48.437846Z","iopub.status.idle":"2022-02-07T11:52:49.396554Z","shell.execute_reply.started":"2022-02-07T11:52:48.437760Z","shell.execute_reply":"2022-02-07T11:52:49.395677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load data","metadata":{}},{"cell_type":"code","source":"train_csv_path = '../input/happy-whale-and-dolphin/train.csv'\ntrain_df = pd.read_csv(train_csv_path)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:52.688887Z","iopub.execute_input":"2022-02-07T11:52:52.689142Z","iopub.status.idle":"2022-02-07T11:52:52.756761Z","shell.execute_reply.started":"2022-02-07T11:52:52.689113Z","shell.execute_reply":"2022-02-07T11:52:52.755860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.describe()","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:53.736057Z","iopub.execute_input":"2022-02-07T11:52:53.736743Z","iopub.status.idle":"2022-02-07T11:52:53.813950Z","shell.execute_reply.started":"2022-02-07T11:52:53.736702Z","shell.execute_reply":"2022-02-07T11:52:53.813122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.dtypes","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:53.946947Z","iopub.execute_input":"2022-02-07T11:52:53.947804Z","iopub.status.idle":"2022-02-07T11:52:53.954741Z","shell.execute_reply.started":"2022-02-07T11:52:53.947758Z","shell.execute_reply":"2022-02-07T11:52:53.954028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get unique species\nunique_species = train_df.species.unique()\nunique_species","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:55.178132Z","iopub.execute_input":"2022-02-07T11:52:55.178697Z","iopub.status.idle":"2022-02-07T11:52:55.189885Z","shell.execute_reply.started":"2022-02-07T11:52:55.178658Z","shell.execute_reply":"2022-02-07T11:52:55.189152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's see the distribution of each species\nsns.countplot(train_df.species)\nplt.xticks(rotation=90)","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:55.332724Z","iopub.execute_input":"2022-02-07T11:52:55.333193Z","iopub.status.idle":"2022-02-07T11:52:55.794638Z","shell.execute_reply.started":"2022-02-07T11:52:55.333153Z","shell.execute_reply":"2022-02-07T11:52:55.793950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Custom Dataset","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, root_dir, df, label_to_id, transform):\n        self.root_dir = root_dir\n        self.df = df\n        self.label_to_id = label_to_id\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        image_path = os.path.join(self.root_dir, self.df.iloc[index, 0])\n        image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        label = self.df.iloc[index, 2]\n        target = self.label_to_id[label]\n        \n        image = self.transform(image)\n        return image, torch.tensor(target)","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:57.828097Z","iopub.execute_input":"2022-02-07T11:52:57.828366Z","iopub.status.idle":"2022-02-07T11:52:57.837036Z","shell.execute_reply.started":"2022-02-07T11:52:57.828336Z","shell.execute_reply":"2022-02-07T11:52:57.836184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = transforms.Compose([transforms.ToPILImage(),\n                                       transforms.Resize((56,56)),\n                                       transforms.RandomHorizontalFlip(),\n                                       transforms.ToTensor(),\n                                       transforms.Normalize([0.5,0.5,0.5],\n                                                            [0.5,0.5,0.5])])","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:58.776586Z","iopub.execute_input":"2022-02-07T11:52:58.777117Z","iopub.status.idle":"2022-02-07T11:52:58.782170Z","shell.execute_reply.started":"2022-02-07T11:52:58.777078Z","shell.execute_reply":"2022-02-07T11:52:58.781303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_individual_ids = train_df.individual_id.unique()\nunique_individual_ids","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:58.986982Z","iopub.execute_input":"2022-02-07T11:52:58.987301Z","iopub.status.idle":"2022-02-07T11:52:58.998998Z","shell.execute_reply.started":"2022-02-07T11:52:58.987272Z","shell.execute_reply":"2022-02-07T11:52:58.998226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_to_id = {}\nid_to_label = {}\nidx = 0\nfor label in unique_individual_ids:\n    label_to_id[label] = idx\n    id_to_label[idx] = label\n    idx += 1","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:52:59.908614Z","iopub.execute_input":"2022-02-07T11:52:59.909342Z","iopub.status.idle":"2022-02-07T11:52:59.921034Z","shell.execute_reply.started":"2022-02-07T11:52:59.909298Z","shell.execute_reply":"2022-02-07T11:52:59.920109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_dir = '../input/happy-whale-and-dolphin/train_images'\n\ndataset = CustomDataset(root_dir,\n                        train_df,\n                        label_to_id,\n                        train_transforms)\n\ntrain_loader = DataLoader(dataset, batch_size=8, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:53:00.062961Z","iopub.execute_input":"2022-02-07T11:53:00.063289Z","iopub.status.idle":"2022-02-07T11:53:00.067802Z","shell.execute_reply.started":"2022-02-07T11:53:00.063246Z","shell.execute_reply":"2022-02-07T11:53:00.066849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load pretrained vgg16","metadata":{}},{"cell_type":"code","source":"model = vgg16(pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:54:48.688865Z","iopub.execute_input":"2022-02-07T11:54:48.689118Z","iopub.status.idle":"2022-02-07T11:54:50.215927Z","shell.execute_reply.started":"2022-02-07T11:54:48.689090Z","shell.execute_reply":"2022-02-07T11:54:50.215131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"last_checkpoint = torch.load('../input/happywhale-pytorch-vgg16/last_checkpoint.pth.tar')","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:55:04.587431Z","iopub.execute_input":"2022-02-07T11:55:04.588149Z","iopub.status.idle":"2022-02-07T11:55:06.234303Z","shell.execute_reply.started":"2022-02-07T11:55:04.588100Z","shell.execute_reply":"2022-02-07T11:55:06.233539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"last_checkpoint.keys()","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:56:37.113068Z","iopub.execute_input":"2022-02-07T11:56:37.113348Z","iopub.status.idle":"2022-02-07T11:56:37.119377Z","shell.execute_reply.started":"2022-02-07T11:56:37.113315Z","shell.execute_reply":"2022-02-07T11:56:37.118240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.classifier = nn.Sequential(\n    nn.Linear(25088, 4096),\n    nn.ReLU(),\n    nn.Dropout(),\n    nn.Linear(4096, len(label_to_id))\n)","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:55:38.222859Z","iopub.execute_input":"2022-02-07T11:55:38.223117Z","iopub.status.idle":"2022-02-07T11:55:39.618665Z","shell.execute_reply.started":"2022-02-07T11:55:38.223087Z","shell.execute_reply":"2022-02-07T11:55:39.617848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for name, param in model.named_parameters():\n    if 'classifier' not in name:\n        param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:55:40.229518Z","iopub.execute_input":"2022-02-07T11:55:40.230373Z","iopub.status.idle":"2022-02-07T11:55:40.235344Z","shell.execute_reply.started":"2022-02-07T11:55:40.230322Z","shell.execute_reply":"2022-02-07T11:55:40.234636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(last_checkpoint['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:55:43.644397Z","iopub.execute_input":"2022-02-07T11:55:43.644690Z","iopub.status.idle":"2022-02-07T11:55:43.802743Z","shell.execute_reply.started":"2022-02-07T11:55:43.644658Z","shell.execute_reply":"2022-02-07T11:55:43.802098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(train_loader))\nimages.shape, targets.shape","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:56:03.618792Z","iopub.execute_input":"2022-02-07T11:56:03.619050Z","iopub.status.idle":"2022-02-07T11:56:04.795977Z","shell.execute_reply.started":"2022-02-07T11:56:03.619020Z","shell.execute_reply":"2022-02-07T11:56:04.795159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = model(images.cpu())\noutput.shape","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:56:06.114046Z","iopub.execute_input":"2022-02-07T11:56:06.114334Z","iopub.status.idle":"2022-02-07T11:56:06.559559Z","shell.execute_reply.started":"2022-02-07T11:56:06.114303Z","shell.execute_reply":"2022-02-07T11:56:06.558730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2022-02-07T11:56:07.351049Z","iopub.execute_input":"2022-02-07T11:56:07.351333Z","iopub.status.idle":"2022-02-07T11:56:07.355635Z","shell.execute_reply.started":"2022-02-07T11:56:07.351302Z","shell.execute_reply":"2022-02-07T11:56:07.354828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train model\n\nmodel.to(device)\nstart_epoch = last_checkpoint['epoch']\nEPOCHS = start_epoch + 4\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\noptimizer.load_state_dict(last_checkpoint['optimizer_state_dict'])\n\nlast_train_loss = 0\n\nfor epoch in range(start_epoch, EPOCHS):\n    print(f'Epoch: {epoch+1}/{EPOCHS}')\n    \n    correct = 0\n    total = 0\n    losses = []\n    \n    for batch_idx, data in enumerate(tqdm(train_loader)):\n        images, targets = data\n        images = images.to(device)\n        targets = targets.to(device)\n        \n        output = model(images)  # (batch_size, num_classes)\n        \n        loss = criterion(output, targets)\n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        _, pred = torch.max(output, 1)\n        correct += (pred == targets).sum().item()\n        total += pred.size(0)\n        \n        losses.append(loss.item())\n        \n    train_loss = np.mean(losses)\n    train_acc = correct * 1.0 / total\n    \n    last_train_loss = train_loss\n    print(f'Train Loss: {train_loss}\\tTrain Acc: {train_acc}')","metadata":{"execution":{"iopub.status.busy":"2022-02-07T12:03:42.784117Z","iopub.execute_input":"2022-02-07T12:03:42.784816Z","iopub.status.idle":"2022-02-07T12:03:53.558647Z","shell.execute_reply.started":"2022-02-07T12:03:42.784777Z","shell.execute_reply":"2022-02-07T12:03:53.557552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save({\n    'epoch': EPOCHS,\n    'model_state_dict': model.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    'loss': last_train_loss\n}, 'last_checkpoint.pth.tar')","metadata":{"execution":{"iopub.status.busy":"2022-02-07T05:38:43.479892Z","iopub.execute_input":"2022-02-07T05:38:43.480158Z","iopub.status.idle":"2022-02-07T05:38:49.967672Z","shell.execute_reply.started":"2022-02-07T05:38:43.480127Z","shell.execute_reply":"2022-02-07T05:38:49.966917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make predictions","metadata":{}},{"cell_type":"code","source":"sample_df = pd.read_csv('../input/happy-whale-and-dolphin/sample_submission.csv')\nsample_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-02-07T05:26:21.140877Z","iopub.execute_input":"2022-02-07T05:26:21.14113Z","iopub.status.idle":"2022-02-07T05:26:21.207027Z","shell.execute_reply.started":"2022-02-07T05:26:21.141102Z","shell.execute_reply":"2022-02-07T05:26:21.206348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_transforms = transforms.Compose([transforms.ToPILImage(),\n                                     transforms.Resize(56),\n                                     transforms.ToTensor(),\n                                     transforms.Normalize([0.5,0.5,0.5],\n                                                          [0.5,0.5,0.5])])","metadata":{"execution":{"iopub.status.busy":"2022-02-07T05:36:27.271868Z","iopub.execute_input":"2022-02-07T05:36:27.27214Z","iopub.status.idle":"2022-02-07T05:36:27.277297Z","shell.execute_reply.started":"2022-02-07T05:36:27.27211Z","shell.execute_reply":"2022-02-07T05:36:27.276535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_img_dir = '../input/happy-whale-and-dolphin/test_images'\n\nres = []\n\nfor i in tqdm(range(sample_df.shape[0])):\n    image_path = os.path.join(test_img_dir, sample_df.iloc[i,0])\n    image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    image = test_transforms(image)\n    image = image.unsqueeze(0)\n\n    output = model(image.to(device))\n    _, tk = torch.topk(output, 5, dim=1)\n    pred = []\n    for j in range(len(tk[0])):\n        pred.append(id_to_label[tk[0][j].item()])\n    pred = ' '.join(pred)\n    \n    sample_df.iloc[i, 1] = pred","metadata":{"execution":{"iopub.status.busy":"2022-02-07T05:36:29.28414Z","iopub.execute_input":"2022-02-07T05:36:29.284391Z","iopub.status.idle":"2022-02-07T05:36:33.671582Z","shell.execute_reply.started":"2022-02-07T05:36:29.284361Z","shell.execute_reply":"2022-02-07T05:36:33.67031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-02-06T07:40:43.169899Z","iopub.execute_input":"2022-02-06T07:40:43.17061Z","iopub.status.idle":"2022-02-06T07:40:43.290755Z","shell.execute_reply.started":"2022-02-06T07:40:43.170565Z","shell.execute_reply":"2022-02-06T07:40:43.290016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Done!')","metadata":{"execution":{"iopub.status.busy":"2022-02-06T07:40:48.052247Z","iopub.execute_input":"2022-02-06T07:40:48.052914Z","iopub.status.idle":"2022-02-06T07:40:48.057633Z","shell.execute_reply.started":"2022-02-06T07:40:48.052877Z","shell.execute_reply":"2022-02-06T07:40:48.056542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}