{"cells":[{"metadata":{"trusted":true,"_uuid":"27393d3b533b1928609f3bceda21209b7a08f6fa"},"cell_type":"code","source":"!pip install albumentations > /dev/null","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"0cfaa55cedb7e80727b272bf5b648dc38f93bf62"},"cell_type":"code","source":"!pip install pretrainedmodels > /dev/null","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import os\n\nimport albumentations\nfrom albumentations import torch as AT\nimport pretrainedmodels\n\nimport numpy as np\nimport pandas as pd\n\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom PIL import Image\nfrom sklearn.preprocessing import LabelEncoder, OneHotEncoder\n\nfrom tqdm import tqdm\n\nfrom matplotlib import pyplot as plt\n%matplotlib inline\n\nimport warnings\nwarnings.simplefilter(\"ignore\", category=DeprecationWarning)","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"train_df = pd.read_csv(\"../input/train.csv\")\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"8f62fc320145a438147b857169f6725c88c58274"},"cell_type":"code","source":"train_df.shape, train_df.Id.nunique()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b4fffbaf19f8371fe63f07f8ff72ae4afb0e2dc7"},"cell_type":"code","source":"NUM_CLASSES = train_df.Id.nunique()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"cc066a56726375d28143b284158716305d200d99"},"cell_type":"code","source":"train_df.Id.value_counts().iloc[1:].hist(bins=40)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"24acf1683033a1e7d882fa4cfca4bdf7a62dd182"},"cell_type":"code","source":"RESIZE_H = 160\nRESIZE_W = 320\n\ndata_transforms = albumentations.Compose([\n    albumentations.Resize(RESIZE_H, RESIZE_W),\n    albumentations.HorizontalFlip(),\n    albumentations.OneOf([\n        albumentations.RandomContrast(),\n        albumentations.RandomBrightness(),\n    ]),\n    albumentations.ShiftScaleRotate(rotate_limit=10, scale_limit=0.15),\n    albumentations.JpegCompression(80),\n    albumentations.HueSaturationValue(),\n    albumentations.Normalize(),\n    AT.ToTensor()\n])\n\ndata_transforms_test = albumentations.Compose([\n    albumentations.Resize(RESIZE_H, RESIZE_W),\n    albumentations.Normalize(),\n    AT.ToTensor()\n])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"eaef0437168af0b48db1b8cf43589901734606b3"},"cell_type":"code","source":"def prepare_labels(y):\n    # From here: https://www.kaggle.com/pestipeti/keras-cnn-starter\n    values = np.array(y)\n    label_encoder = LabelEncoder()\n    integer_encoded = label_encoder.fit_transform(values)\n\n    onehot_encoder = OneHotEncoder(sparse=False)\n    integer_encoded = integer_encoded.reshape(len(integer_encoded), 1)\n    onehot_encoded = onehot_encoder.fit_transform(integer_encoded)\n\n    y = onehot_encoded\n    return y, label_encoder","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"28538a1976383479893384ed691d116fc878845b"},"cell_type":"code","source":"y, lab_encoder = prepare_labels(train_df['Id'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"efffa7e82f50f339aab100fed67835e49dd8da89"},"cell_type":"code","source":"class WhaleDataset(Dataset):\n    def __init__(self, datafolder, datatype='train', df=None, transform=None, y=None):\n        self.datafolder = datafolder\n        self.datatype = datatype\n        self.y = y\n        if self.datatype == 'train':\n            self.df = df.values\n        self.image_files_list = [s for s in os.listdir(datafolder)]\n        self.transform = transform\n\n\n    def __len__(self):\n        return len(self.image_files_list)\n    \n    def __getitem__(self, idx):\n        if self.datatype == 'train':\n            img_name = os.path.join(self.datafolder, self.df[idx][0])\n            label = self.y[idx]\n            \n        elif self.datatype == 'test':\n            img_name = os.path.join(self.datafolder, self.image_files_list[idx])\n            label = np.zeros((NUM_CLASSES,))\n\n        img = cv2.imread(img_name)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        image = self.transform(image=img)['image']\n        if self.datatype == 'train':\n            return image, label\n        elif self.datatype == 'test':\n            # so that the images will be in a correct order\n            return image, label, self.image_files_list[idx]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"52816552dfa11714cb8e7866b422474971f78d82"},"cell_type":"code","source":"train_dataset = WhaleDataset(\n    datafolder='../input/train/', \n    datatype='train', \n    df=train_df, \n    transform=data_transforms, \n    y=y\n)\n\ntest_set = WhaleDataset(\n    datafolder='../input/test/', \n    datatype='test', \n    transform=data_transforms_test\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"68d634c72ad9eb5b0ebd551661983e3d93e6ce79"},"cell_type":"code","source":"batch_size = 10\nnum_workers = 4\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, pin_memory=True)\ntest_loader = DataLoader(test_set, batch_size=batch_size, num_workers=num_workers, pin_memory=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"2ba1f48468ed2d2bf427b95d9a03abe12f793280"},"cell_type":"code","source":"model = pretrainedmodels.resnext101_64x4d()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"f6141ad16752e62fd4ed9aa3275cf78abf4f992c"},"cell_type":"code","source":"model.avg_pool = nn.AvgPool2d((5,10))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"d7da045079a594c5ce9bc87b1e7e5f7b0a2d0465"},"cell_type":"code","source":"model.last_linear = nn.Linear(model.last_linear.in_features, NUM_CLASSES)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"a7755b4730cdfef20e2e28c18e1d1b21cb010123"},"cell_type":"code","source":"model.cuda();","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"cc485dbfb03f61c128424e7add4c2e487b4ea9bb"},"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(model.parameters(), lr=0.0005)\n\nscheduler = lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"4f3c57883af3f42ff55d5a34ef1c7dcf5001dbcd"},"cell_type":"code","source":"def cuda(x):\n    return x.cuda(non_blocking=True) if torch.cuda.is_available() else x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"23aa1fbfacbf9c73bb3651dd70de1c78c4b21408"},"cell_type":"code","source":"n_epochs = 5\nfor epoch in range(1, n_epochs+1):\n    train_loss = []\n    \n    for batch_i, (data, target) in tqdm(enumerate(train_loader), total = len(train_loader)):\n        data, target = cuda(data), cuda(target)\n\n        optimizer.zero_grad()\n        output = model(data)\n        loss = criterion(output, target.float())\n        train_loss.append(loss.item())\n\n        loss.backward()\n        optimizer.step()\n    \n    scheduler.step()\n    \n    print(f'Epoch {epoch}, train loss: {np.mean(train_loss):.4f}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e3f5def101bccf9065e00c83f0597a96437bfb81"},"cell_type":"code","source":"sub = pd.read_csv('../input/sample_submission.csv')\n\nmodel.eval()\nfor (data, target, name) in tqdm(test_loader):\n    data = cuda(data)\n    output = model(data)\n    output = output.cpu().detach().numpy()\n    for i, (e, n) in enumerate(list(zip(output, name))):\n        sub.loc[sub['Image'] == n, 'Id'] = ' '.join(lab_encoder.inverse_transform(e.argsort()[-5:][::-1]))\n        \nsub.to_csv('submission.csv', index=False)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}