{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"from albumentations.pytorch import ToTensor\nfrom albumentations import (\n    Compose, HorizontalFlip, CLAHE, HueSaturationValue,\n    RandomBrightness, RandomContrast, RandomGamma, OneOf, Resize,\n    ToFloat, ShiftScaleRotate, GridDistortion, RandomRotate90, Cutout,\n    RGBShift, RandomBrightness, RandomContrast, Blur, MotionBlur, MedianBlur, GaussNoise, CoarseDropout,\n    IAAAdditiveGaussianNoise, GaussNoise, OpticalDistortion, RandomSizedCrop, VerticalFlip, Normalize\n)\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nfrom torch.utils.data import DataLoader\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as t\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\nfrom sklearn import metrics\n\n!pip install --upgrade efficientnet_pytorch\nfrom efficientnet_pytorch import EfficientNet\nfrom tqdm.notebook import tqdm\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import scipy\n\nfrom numpy import pi\nfrom numpy import sin\nfrom numpy import zeros\nfrom numpy import r_\nfrom scipy import signal\nfrom scipy import misc # pip install Pillow\nimport matplotlib.pylab as pylab\n\n%matplotlib inline\npylab.rcParams['figure.figsize'] = (20.0, 7.0)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import random\n\nseed = 42\nprint(f'setting everything to seed {seed}')\nrandom.seed(seed)\nos.environ['PYTHONHASHSEED'] = str(seed)\nnp.random.seed(seed)\ntorch.manual_seed(seed)\ntorch.cuda.manual_seed(seed)\ntorch.backends.cudnn.deterministic = True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data_dir = '../input/alaska2-image-steganalysis'\nfolder_names = ['JMiPOD/', 'JUNIWARD/', 'UERD/']\nclass_names = ['Normal', 'JMiPOD_75', 'JMiPOD_90', 'JMiPOD_95', \n               'JUNIWARD_75', 'JUNIWARD_90', 'JUNIWARD_95',\n                'UERD_75', 'UERD_90', 'UERD_95']\nclass_labels = { name: i for i, name in enumerate(class_names)}\nnum_classes = len(class_labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df = pd.read_csv('../input/alaska2trainvalsplit/alaska2_train_df.csv')\nval_df = pd.read_csv('../input/alaska2trainvalsplit/alaska2_val_df.csv')\n\nprint(train_df.head(10))\ntrain_df.Label.hist()\nplt.title('Distribution of Classes')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#train_df = train_df.sample(1000)\n#val_df = val_df.sample(500)\n#train_df,val_df","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import scipy\nimport os\nimport numpy as np\nimport pandas as pd\nfrom numpy import pi\nfrom numpy import sin\nfrom numpy import zeros\nfrom numpy import r_\nfrom scipy import signal\nfrom scipy import misc #pip install Pillow\nfrom scipy import fftpack\nimport matplotlib.pylab as pylab\n\n%matplotlib inline\npylab.rcParams['figure.figsize'] = (20.0, 7.0)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def dct2(a):\n    return scipy.fftpack.dct( scipy.fftpack.dct( a, axis=0, norm='ortho' ), axis=1, norm='ortho' )\n\ndef idct2(a):\n    return scipy.fftpack.idct( scipy.fftpack.idct( a, axis=0 , norm='ortho'), axis=1 , norm='ortho')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def dct_ext(img):\n    imsize = img.shape\n    dct = np.zeros(imsize)\n    for i in r_[:imsize[0]:8]:\n        for j in r_[:imsize[1]:8]:\n            dct[i:(i+8),j:(j+8)] = dct2( img[i:(i+8),j:(j+8)] )\n\n    thresh = 0.02\n    dct_thresh = dct * (abs(dct) > (thresh*np.max(dct)))\n    return dct_thresh","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from torch.utils.data import Dataset\nimport cv2\n\nclass Alaska(Dataset):\n    \n    def __init__(self, dataframe, trans = None):\n        self.data = dataframe\n        self.transform = trans\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        fname, target = self.data.iloc[idx]\n        img = cv2.imread(fname)[:, :, ::-1]\n        #img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32)\n        #img/= 255\n        #img = dct_ext(img)\n        \n        if self.transform:\n            img = self.transform(image = img)\n        x = (img['image'], target)\n        \n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"augmentations_train = Compose([\n    #Resize(512, 512, p=1), \n    VerticalFlip(p=0.5),\n    HorizontalFlip(p=0.5),\n    ToFloat(max_value=255),\n    ToTensor()\n],p=1)\n\naugmentations_test = Compose([\n    #Resize(512, 512, p=1),\n    ToFloat(max_value=255),\n    ToTensor()\n])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_ds = Alaska(train_df, trans = augmentations_train)\nval_ds = Alaska(val_df, trans = augmentations_test)\ntrain_ds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"len(train_ds)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"img, lab = train_ds[200]\nplt.imshow(img.permute(1,2,0))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"batch_size = 64\nnum_workers = 0\n\ntemp_dl = DataLoader(train_ds, batch_size = batch_size, num_workers = num_workers, shuffle=True)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\n\nimages, labels = next(iter(temp_dl))\nimages = images.permute(0, 2, 3, 1)\nmax_images = 64\ngrid_width = 16\ngrid_height = int(max_images / grid_width)\nfig, axs = plt.subplots(grid_height, grid_width,\n                        figsize=(grid_width+1, grid_height+1))\n\nfor i, (im, label) in enumerate(zip(images, labels)):\n    ax = axs[int(i / grid_width), i % grid_width]\n    ax.imshow(im.squeeze())\n    ax.set_title(str(label.item()))\n    ax.axis('off')\n\nplt.suptitle(\"0: No Hidden Message, 1: JMiPOD, 2: JUNIWARD, 3:UERD\")\nplt.show()\ndel images, temp_dl\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_default_device():\n    if torch.cuda.is_available():\n        return torch.device('cuda')\n    else:\n        return torch.device('cpu')\n\ndef to_device(data, device):\n    if isinstance(data, (list,tuple)):\n        return [to_device(x, device) for x in data]\n    return data.to(device, non_blocking = True)\n\nclass DeviceDataLoader():\n    def __init__(self, dl, device):\n        self.dl = dl\n        self.device = device\n        \n    def __iter__(self):\n        for b in self.dl:\n            yield to_device(b, self.device)\n            \n    def __len__(self):\n        return len(self.dl)\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = get_default_device()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.model = EfficientNet.from_pretrained('efficientnet-b0')\n        # b0 => 1280\n        self.dense_output = nn.Linear(1280, num_classes)\n\n    def forward(self, x):\n        feat = self.model.extract_features(x)\n        feat = F.avg_pool2d(feat, feat.size()[2:]).reshape(-1, 1280)\n        return self.dense_output(feat)\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = Net()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# https://www.kaggle.com/anokas/weighted-auc-metric-updated\n\ndef alaska_weighted_auc(y_true, y_valid):\n    tpr_thresholds = [0.0, 0.4, 1.0]\n    weights = [2,   1]\n\n    fpr, tpr, thresholds = metrics.roc_curve(y_true, y_valid, pos_label=1)\n\n    # size of subsets\n    areas = np.array(tpr_thresholds[1:]) - np.array(tpr_thresholds[:-1])\n\n    # The total area is normalized by the sum of weights such that the final weighted AUC is between 0 and 1.\n    normalization = np.dot(areas, weights)\n\n    competition_metric = 0\n    for idx, weight in enumerate(weights):\n        y_min = tpr_thresholds[idx]\n        y_max = tpr_thresholds[idx + 1]\n        mask = (y_min < tpr) & (tpr < y_max)\n        # pdb.set_trace()\n\n        x_padding = np.linspace(fpr[mask][-1], 1, 100)\n\n        x = np.concatenate([fpr[mask], x_padding])\n        y = np.concatenate([tpr[mask], [y_max] * len(x_padding)])\n        y = y - y_min  # normalize such that curve starts at y=0\n        score = metrics.auc(x, y)\n        submetric = score * weight\n        best_subscore = (y_max - y_min) * weight\n        competition_metric += submetric\n\n    return competition_metric / normalization","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def loss_batch(model, loss_func, xb, yb, opt = None, metric = None):\n    preds = model(xb)\n    \n    loss = loss_func(preds, yb)\n    if opt is not None:\n        loss.backward()\n        opt.step()\n        opt.zero_grad()\n        \n    metric_result = None\n    if metric is not None:\n        metric_result = metric(preds, yb)\n    return loss.item(), len(xb), metric_result","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def evaluate(model, loss_fn, valid_dl, metric = None):\n    labs, predictions = [], []\n    with torch.no_grad():\n        for imgs, labels in valid_dl:\n            imgs = imgs.to(device, dtype=torch.float)\n            labels = labels.to(device, dtype=torch.long)\n            preds = model(imgs)\n            loss = loss_fn(preds, labels)\n            \n            labs.extend(labels.cpu().numpy().astype(int))\n            predictions.extend(F.softmax(preds, 1).cpu().numpy())\n                \n        results = [loss_batch(model, loss_fn, xb, yb, metric = metric) \n                          for xb, yb in valid_dl]\n        losses, nums, metrics = zip(*results)\n\n        predictions = np.array(predictions)\n        pred_labels = predictions.argmax(1)\n        \n        eval_accuracy = (pred_labels == labs).mean()\n        \n        new_preds = np.zeros(len(predictions))\n        temp = predictions[pred_labels != 0, 1:]\n\n        new_preds[pred_labels != 0] = temp.sum(1)\n        new_preds[pred_labels == 0] = 1 - predictions[pred_labels == 0, 0]\n        labs = np.array(labs)\n        labs[labs != 0] = 1\n        \n        auc_score = alaska_weighted_auc(labs, new_preds)\n        \n        total = np.sum(nums)\n    \n        avg_loss = np.sum(np.multiply(losses, nums))/total\n    \n        avg_metric = None\n        if metric is not None:\n            avg_metric = np.sum(np.multiply(metrics, nums))/total\n    return avg_loss, total, avg_metric, auc_score, eval_accuracy","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def fit(epochs, model, loss_fn, train_dl, valid_dl, opt_fn = None, lr = None, metric = None):\n    train_losses, val_losses, val_metrics, auc_metrics, eval_metrics = [], [], [], [], []\n    \n    if opt_fn is None: opt_fn = torch.optim.SGD\n    opt = opt_fn(model.parameters(), lr = lr)\n    \n    for epoch in range (epochs):\n        model.train()\n        for xb, yb in train_dl:\n            train_loss, _, _ =loss_batch(model, loss_fn, xb, yb, opt)\n            \n        model.eval()\n        result = evaluate(model, loss_fn, valid_dl, metric)\n        val_loss, total, val_metric, auc_score, eval_accuracy = result\n        \n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        val_metrics.append(val_metric)\n        auc_metrics.append(auc_score)\n        eval_metrics.append(eval_accuracy)\n        \n        if metric is None:\n            print('Epoch [{}/{}], train_loss: {:4f}, val_loss: {:.4f}'\n                  .format(epoch+1, epochs, train_loss, val_loss))\n        else:\n            print('Epoch [{}/{}], train_loss: {:4f}, val_loss: {:.4f}, val_{}: {:.4f},auc_score: {:.4f}, eval_accuracy: {:.4f}'\n                 .format(epoch+1, epochs, train_loss, val_loss, metric.__name__, val_metric, auc_score, eval_accuracy))\n    return train_losses, val_losses, val_metrics, auc_metrics, eval_metrics","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def accuracy(outputs, labels):\n    _,preds = torch.max(outputs, dim = 1)\n    return torch.sum(preds == labels).item()/len(preds)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"batch_size = 8\nnum_workers = 8\n\ntrain_dl = DataLoader(train_ds, batch_size = batch_size, num_workers = num_workers, shuffle=True)\nval_dl = DataLoader(val_ds, batch_size = batch_size, num_workers = num_workers, shuffle=False)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = Net() \ntrain_dl = DeviceDataLoader(train_dl, device)\nval_dl = DeviceDataLoader(val_dl, device)\nto_device(model, device)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def load_cp(cp):\n    model.load_state_dict(cp['model_state'])\n    ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for p in model.parameters():\n    print(p)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"cp = torch.load('../input/checkpoint/checkpoint26102020.pth')\nload_cp(cp)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for p in model.parameters():\n    print(p)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"num_epochs = 2\nopt_fn = torch.optim.AdamW\nlr = 1e-4","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"x = fit(num_epochs, model, F.cross_entropy, train_dl, val_dl, opt_fn, lr, metric = accuracy)\ntrain_losses, val_losses, val_metrics, auc_metrics, eval_metrics = x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import glob\nclass Alaska2TestDataset(Dataset):\n\n    def __init__(self, df, augmentations=None):\n\n        self.data = df\n        self.augment = augmentations\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        fn = self.data.loc[idx][0]\n        im = cv2.imread(fn)[:, :, ::-1]\n\n        if self.augment:\n            # Apply transformations\n            im = self.augment(image=im)\n\n        return im\n\n\ntest_filenames = sorted(glob.glob(f\"{data_dir}/Test/*.jpg\"))\ntest_df = pd.DataFrame({'ImageFileName': list(\n    test_filenames)}, columns=['ImageFileName'])\n\nbatch_size = 16\nnum_workers = 4\ntest_dataset = Alaska2TestDataset(test_df, augmentations=augmentations_test)\ntest_loader = torch.utils.data.DataLoader(test_dataset,\n                                          batch_size=batch_size,\n                                          num_workers=num_workers,\n                                          shuffle=False,\n                                          drop_last=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.eval()\n\npreds = []\ntk0 = tqdm(test_loader)\nwith torch.no_grad():\n    for i, im in enumerate(tk0):\n        inputs = im[\"image\"].to(device)\n        # flip vertical\n        im = inputs.flip(2)\n        outputs = model(im)\n        # fliplr\n        im = inputs.flip(3)\n        outputs = (0.25*outputs + 0.25*model(im))\n        outputs = (outputs + 0.5*model(inputs))\n        labels = labels.to(device, dtype=torch.long)\n        \n        preds.extend(F.softmax(outputs, 1).cpu().numpy())\n\npreds = np.array(preds)\nlabels = preds.argmax(1)\nnew_preds = np.zeros((len(preds),))\ntemp = preds[labels != 0, 1:]\nnew_preds[labels != 0] = [temp[i, val] for i, val in enumerate(temp.argmax(1))]\nnew_preds[labels == 0] = preds[labels == 0, 0]\n\ntest_df['Id'] = test_df['ImageFileName'].apply(lambda x: x.split(os.sep)[-1])\ntest_df['Label'] = new_preds\n\ntest_df = test_df.drop('ImageFileName', axis=1)\ntest_df.to_csv('submission.csv', index=False)\nprint(test_df.head())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.save(model.state_dict(), 'alaska2-effnetb0-4eps.pth')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint = {\n    'epochs': 4,\n    'model_state': model.state_dict(),\n    'optim_state': opt_fn.state_dict()\n}\ntorch.save(checkpoint, 'checkpoint07112020.pth')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#model = torch.load(PATH)\n#model.eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"markdown","source":"for p in model.parameters():\n    print(p)"},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}