{"cells":[{"metadata":{"_uuid":"a34796c9-3e96-417f-b028-6065fe0420fd","_cell_guid":"f9fdf29a-2f31-4174-96c0-0b95c197afee","trusted":true},"cell_type":"code","source":"import sys\npackage_path = '../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master'\nsys.path.append(package_path)\nimport os\nimport pandas as pd\nimport skimage.io\nimport numpy as np\nfrom matplotlib import pyplot as plt\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations import Normalize, Compose, RandomRotate90, Resize, Transpose, VerticalFlip, HorizontalFlip\n\nfrom tqdm import tqdm\nfrom efficientnet_pytorch import model as enet\n\nfrom sklearn.model_selection import StratifiedKFold\n\nTRAIN = False\n\ntrain_path = \"../input/prostate-cancer-grade-assessment/train.csv\"\ntest_path = \"../input/prostate-cancer-grade-assessment/test.csv\"\nimage_path = \"../input/prostate-cancer-grade-assessment/train_images\"\ntest_image_path = \"../input/prostate-cancer-grade-assessment/test_images\"\nmask_path = \"../input/prostate-cancer-grade-assessment/train_label_masks\"\nsample_sub_path = \"../input/prostate-cancer-grade-assessment/sample_submission.csv\"\ntrained_model_path = \"../input/b1-net-64eps/enet-b1\"\nsub_path = \"submission.csv\"\n\ntrain_data = pd.read_csv(train_path)\ntest_data = pd.read_csv(test_path)\nsample_data = pd.read_csv(sample_sub_path)\ndummy_df = train_data[:20]\n\ntile_size = 256\nimage_size = 256\nn_tiles = 36\n\n\n\ndef get_tiles(img, mode=0):\n    result = []\n    h, w, c = img.shape\n    pad_h = (tile_size - h % tile_size) + ((tile_size * mode) // 2)  # %tile_size?\n    pad_w = (tile_size - w % tile_size) + ((tile_size * mode) // 2)\n\n    img = np.pad(img,[[pad_h // 2, pad_h - pad_h // 2], [pad_w // 2,pad_w - pad_w//2], [0,0]], constant_values=255)\n    img = img.reshape(img.shape[0]//tile_size, tile_size, img.shape[1]//tile_size, tile_size, 3)\n    img = img.transpose(0, 2, 1, 3, 4).reshape(-1, tile_size, tile_size, 3)\n    n_tiles_with_info = (img.reshape(img.shape[0], -1).sum(1) < tile_size ** 2 * 3 * 255).sum()\n\n    index = np.argsort(img.reshape(img.shape[0],-1).sum(-1))[:n_tiles]\n\n    img = img[index]\n    for i in range(len(img)):\n        result.append({\"img\":img[i], \"index\": i})\n    return result, n_tiles_with_info >=n_tiles\n\n\nclass Dataset(Dataset):\n    def __init__(self, df, image_size=image_size, n_tiles=n_tiles, tile_mode=0):\n        self.df = df.reset_index(drop=True)\n        self.image_size = image_size\n        self.n_tiles = n_tiles\n        self.tile_mode = tile_mode\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, index):\n        if TRAIN:\n            img_path = image_path\n        if not TRAIN:\n            if os.path.exists(test_image_path):\n                img_path = test_image_path\n            else:\n                img_path = image_path\n\n        tiff_path = os.path.join(img_path, self.df[\"image_id\"][index] + \".tiff\")\n        image = skimage.io.MultiImage(tiff_path)[1]  # Middle size one\n        tiles, all_with_info = get_tiles(image, self.tile_mode)\n\n        idxes = list(range(self.n_tiles))\n        idxes = np.asarray(idxes)\n\n        n_row_tiles = int(np.sqrt(self.n_tiles))  # 6\n        images = np.zeros((image_size * n_row_tiles, image_size * n_row_tiles, 3))  # a pic of stack of tiles\n\n        for h in range(n_row_tiles):\n            for w in range(n_row_tiles):\n                i = h * n_row_tiles + w\n                if len(tiles) > idxes[i]:\n                    this_img = tiles[idxes[i]][\"img\"]\n                else:\n                    this_img = np.ones((self.image_size, self.image_size, 3)).astype(np.int)*255\n                this_img = 255 - this_img\n\n                h1 = h * image_size\n                w1 = w * image_size\n                images[h1:h1 + image_size, w1:w1 + image_size] = this_img\n\n        images = images.astype(np.float32)\n        images /= 255\n\n        if TRAIN:\n            transform = Compose([Transpose(p=0.5), VerticalFlip(p=0.5), HorizontalFlip(p=0.5), RandomRotate90(p=0.5)])\n            images = transform(image=images)[\"image\"]\n            images = images.transpose(2, 0, 1)\n            label = np.zeros(5).astype(np.float32)\n            label[:self.df[\"isup_grade\"][index]] = 1.\n\n            return torch.tensor(images), torch.tensor(label)\n\n        if not TRAIN:\n            images = images.transpose(2, 0, 1)\n            return torch.tensor(images)\n\n\nclass eNet(nn.Module):\n    def __init__(self, backbone, out_dim):\n        super(eNet, self).__init__()\n        self.enet = enet.EfficientNet.from_name(backbone)\n        self.lastfc = nn.Linear(self.enet._fc.in_features, out_dim) # Todo: in_features?\n        self.enet._fc = nn.Identity()\n\n    def extract(self, x):\n        return self.enet(x)\n\n    def forward(self, x):\n        x = self.extract(x)\n        x = self.lastfc(x)\n        return x\n\n\nbackbone = \"efficientnet-b1\"\nmodel = eNet(backbone, 5).cuda()\npre_trained_model_path = trained_model_path\nmodel.load_state_dict(torch.load(pre_trained_model_path, map_location=\"cpu\"))\nmodel.eval()\n\n\ndef prediction(df):\n\n    final_pred = []\n    test_dataset = Dataset(df, image_size, n_tiles, 0)\n    test_bar = tqdm(DataLoader(test_dataset, batch_size=1, shuffle=False), desc=\"Outputting: \", unit=\" batches\")\n\n    for images in test_bar:\n        images = images.cuda()\n        pred = model(images).cuda()\n        grade = pred.sigmoid().round().cpu().detach().numpy()\n        grade = np.sum(grade == 1.)\n        final_pred.append(grade)\n\n    image_id = df[\"image_id\"]\n    isup_grade = final_pred\n    sub_df = pd.DataFrame({\"image_id\": image_id, \"isup_grade\": isup_grade})\n    sub_df.to_csv(sub_path, index=False)\n\n\n\nif TRAIN:\n    epoch, alpha = 64, 3e-4\n    batch_size = 1\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=alpha)\n    loss = nn.BCEWithLogitsLoss().cuda()\n\n    train_dataset = Dataset(train_data, image_size, n_tiles)\n    for i in range(epoch):\n        train_bar = tqdm(DataLoader(train_dataset, batch_size=batch_size), desc=\"Training: \", unit=\" batches\")\n        for images, label in train_bar:\n            # Put into GPU\n            images = images.cuda()\n            label = label.cuda()\n            # Training\n            pred = model(images).cuda()\n            loss_value = loss(pred, label)\n            optimizer.zero_grad()\n            loss_value.backward()\n            optimizer.step()\n            # Feedback\n            loss_show = loss_value.cpu().detach().numpy()\n            print(\"Epoch: \", i + 1, \"loss: \", loss_show)\n        torch.save(model.state_dict(), trained_model_path)  # Saving model for every epoch\n\nif not TRAIN:\n    if os.path.exists(test_image_path):\n        prediction(test_data)\n    else:\n        prediction(dummy_df)\n","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}