{"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 Library","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt \nimport seaborn as sns\n\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensor\n\nimport torch\nimport cv2\n\nimport torch.nn as nn\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\n\n\nfrom sklearn import metrics, model_selection\n\n%matplotlib inline","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\") ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.label.value_counts().sort_index().plot.barh()\ndf.label.value_counts().sort_index()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization According to Label(0,1,2,3,4)","metadata":{}},{"cell_type":"code","source":"label0_img = df[df['label']==0].image_id.values\nlabel1_img = df[df['label']==1].image_id.values\nlabel2_img = df[df['label']==2].image_id.values\nlabel3_img = df[df['label']==3].image_id.values\nlabel4_img = df[df['label']==4].image_id.values","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '../input/cassava-leaf-disease-classification/train_images/'\n\nlabel0_img_path = [os.path.join(img_path, x) for x in label0_img]\nlabel1_img_path = [os.path.join(img_path, x) for x in label1_img]\nlabel2_img_path = [os.path.join(img_path, x) for x in label2_img]\nlabel3_img_path = [os.path.join(img_path, x) for x in label3_img]\nlabel4_img_path = [os.path.join(img_path, x) for x in label4_img]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\n\nfor i in range(4):\n    \n    plt.subplot(2,2,i+1)\n        \n    img = cv2.imread(label0_img_path[i])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.title(\"label:0\")\n    plt.imshow(img)\n    \nplt.show()\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\n\nfor i in range(4):\n    \n    plt.subplot(2,2,i+1)\n        \n    img = cv2.imread(label1_img_path[i])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.title(\"label:1\")\n    plt.imshow(img)\n    \nplt.show()\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\n\nfor i in range(4):\n    \n    plt.subplot(2,2,i+1)\n        \n    img = cv2.imread(label2_img_path[i])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.title(\"label:2\")\n    plt.imshow(img)\n    \nplt.show()\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\n\nfor i in range(4):\n    \n    plt.subplot(2,2,i+1)\n        \n    img = cv2.imread(label3_img_path[i])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.title(\"label:3\")\n    plt.imshow(img)\n    \nplt.show()\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\n\nfor i in range(4):\n    \n    plt.subplot(2,2,i+1)\n        \n    img = cv2.imread(label4_img_path[i])\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.title(\"label:4\")\n    plt.imshow(img)\n    \nplt.show()\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Split Train ,Val Data","metadata":{}},{"cell_type":"code","source":"df_train, df_val = model_selection.train_test_split(df, test_size=0.1, random_state=42, stratify=df.label.values)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = df_train.reset_index(drop=True)\ndf_val = df_val.reset_index(drop=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape, df_val.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '../input/cassava-leaf-disease-classification/train_images/'\n\ntrain_img_path = [os.path.join(img_path, x) for x in df_train.image_id.values]\nval_img_path = [os.path.join(img_path, x) for x in df_val.image_id.values]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_img_path), len(val_img_path)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_target = df_train.label.values\nval_target = df_val.label.values","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Dataset","metadata":{}},{"cell_type":"code","source":"class LeafDataset(Dataset):\n    def __init__(self, img_ids, targets, transform):\n        self.img_ids = img_ids\n        self.targets = targets\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.img_ids)\n    \n    def __getitem__(self, index):\n        img_id = self.img_ids[index]\n        img = cv2.imread(img_id)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        target = self.targets[index]\n        \n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n            \n        return img, target\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augumentation","metadata":{}},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Resize(256,256),\n    A.Rotate(15,p=0.2),\n    A.VerticalFlip(p=0.2),\n    A.HorizontalFlip(p=0.2),\n    ToTensor()\n])\n\nval_transform=A.Compose([\n    ToTensor()\n])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = LeafDataset(img_ids = train_img_path, targets = train_target, transform=train_transform)\nval_dataset = LeafDataset(img_ids = val_img_path, targets = val_target, transform=val_transform)\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=2)\nval_dataloader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataset),len(val_dataset)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install pretrainedmodels","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pretrainedmodels\n\nmodel_name = 'resnet34'\nmodel = pretrainedmodels.__dict__[model_name](pretrained='imagenet')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"in_features = model.last_linear.in_features","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.last_linear = nn.Linear(in_features, len(np.unique(df.label.values)), bias=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(),lr=0.001)\nloss_fn = nn.CrossEntropyLoss()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available else 'cpu'  ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm_notebook\nfrom sklearn.metrics import accuracy_score","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nbest_score = -1\n\nfor epoch in tqdm_notebook(range(10)):\n    model = model.to(device)\n    model.train()\n    train_loss=[]\n    for inputs, outputs in train_dataloader:\n        inputs = inputs.to(device)\n        outputs = outputs.to(device)\n        \n        optimizer.zero_grad()\n        \n        logit = model(inputs)\n        \n        loss = loss_fn(logit, outputs)\n        train_loss.append(loss.item())\n        \n        loss.backward()\n        optimizer.step()\n        \n    val_loss=[]\n    val_true=[]\n    val_pred=[]\n    \n    model.eval()\n    with torch.no_grad():\n        for inputs, outputs in val_dataloader:\n            inputs = inputs.to(device)\n            outputs = outputs.to(device)\n            \n            logit = model(inputs)\n            \n            loss = loss_fn(logit, outputs)\n            \n            val_loss.append(loss.item())\n\n            val_pred.append(np.argmax(logit.cpu().data.numpy(),axis=1))\n            val_true.append(outputs.cpu().data.numpy())\n        \n    \n    val_pred = np.concatenate(val_pred, axis=0)\n    val_true = np.concatenate(val_true, axis=0)\n\n    \n    score = accuracy_score(val_pred, val_true)\n    \n    print(f\" epoch: {epoch+1}, train_loss: {np.mean(train_loss)}, val_loss:{np.mean(val_loss)}, accuracy:{score}\")\n    \n    if score>best_score:\n        best_score = score\n        \n        state_dict = model.cpu().state_dict()\n        torch.save(state_dict, '../output/kaggle/working/model.pt')\n            \n        \n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = model.load_state_dict(torch.load('../output'))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_image_path = '../input/cassava-leaf-disease-classification/test_images/2216849948.jpg'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.imread(test_image_path)\nimg = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nimg = img/255\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = torch.FloatTensor(img)\nimg = img.permute(1,2,0)\nimg = img.to(device)\n\nresult = model(img)\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}