{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import pytorch_lightning as pl\nimport pandas as pd\nimport cv2 \nimport os\nfrom torch import nn\nfrom torch.utils.data import Dataset,DataLoader\nimport numpy as np\nimport torch\nfrom sklearn.model_selection import train_test_split","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"IMG_SIZE=64\nPATH=\"../input/cassava-leaf-disease-classification/train_images/\"\nCLASSES=5","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaModel(pl.LightningModule):\n    def __init__(self):\n        super().__init__()\n        self.cnc=nn.Conv2d(3,128,5,4)\n        self.rel=nn.ReLU()\n        self.bn=nn.BatchNorm2d(128)\n        self.mxpool=nn.MaxPool2d(4)\n        self.flat=nn.Flatten()\n        self.fc1=nn.Linear(1152,64)\n        self.fc2=nn.Linear(64,64)\n        self.fc3=nn.Linear(64,CLASSES)\n        self.softmax=nn.Softmax()\n        self.accuracy=pl.metrics.Accuracy()\n            \n    def forward(self,x):\n        x=self.cnc(x)\n        x=self.rel(x)\n        x=self.bn(x)\n        x=self.mxpool(x)\n        x=self.flat(x)\n        x=self.fc1(x)\n        x=self.fc2(x)\n        x=self.fc3(x)\n        return x\n    def loss_fn(self,output,target):\n        return nn.CrossEntropyLoss()(output.view(-1,CLASSES),target)\n    def configure_optimizers(self):\n        LR=1e-3\n        optimizer=torch.optim.AdamW(self.parameters(),lr=LR)\n        return optimizer\n    def training_step(self,batch,batch_idx):\n        x,y=batch['x'],batch['y']\n        img=x.view(-1,3,IMG_SIZE,IMG_SIZE)\n        label=y.view(-1)\n        out=self(img)\n        loss=self.loss_fn(out,label)\n        self.log('train_loss',loss)\n        return loss\n    def validation_step(self,batch,batch_idx):\n        x,y=batch['x'],batch['y']\n        img=x.view(-1,3,IMG_SIZE,IMG_SIZE)\n        label=y.view(-1)\n        out=self(img)\n        loss=self.loss_fn(out,label)\n        out=nn.Softmax(-1)(out)\n        logits=torch.argmax(out,dim=1)\n        accu=self.accuracy(logits,label)\n        self.log('valid_loss',loss)\n        self.log('train_acc_step',accu)\n        return loss,accu\n        \n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self,path,image_ids,labels,image_size):\n        self.image_ids=image_ids\n        self.path=path\n        self.labels=labels\n        self.image_size=image_size\n        \n    def __len__(self):\n        return len(self.image_ids)\n    \n    \n    def __getitem__(self,idx):\n        image_ids=str(self.image_ids[idx])\n        labels=self.labels[idx]\n        img_file=cv2.imread(self.path+image_ids)\n        img=cv2.resize(img_file,(self.image_size,self.image_size))\n        img=img.astype(np.float64)\n        return {\n            \"x\":torch.tensor(img,dtype=torch.float),\n            \"y\":torch.tensor(labels,dtype=torch.long)\n        }\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaLightDataset(pl.LightningDataModule):\n    \n    def __init__(self,batch_size=64):\n        super().__init__()\n        self.batch_size=batch_size\n        \n    def setup(self,stage=None):\n        dfx=pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\n        xtrain,xval,ytrain,yval=train_test_split(dfx['image_id'].values,dfx['label'].values,test_size=0.1)\n        self.train_dataset=CassavaDataset(PATH,xtrain,ytrain,IMG_SIZE)\n        self.validation_dataset=CassavaDataset(PATH,xval,yval,IMG_SIZE)\n        \n        \n    def train_dataloader(self):\n            train_loader=DataLoader(self.train_dataset,batch_size=self.batch_size,shuffle=True)\n            return train_loader\n        \n    def validation_dataloader(self):\n            validation_loader=DataLoader(self.validation_dataset,batch_size=self.batch_size,shuffle=False)\n            return validation_loade\n    ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint_callback = pl.callbacks.ModelCheckpoint(\n    monitor='valid_loss',\n    dirpath='./',\n    filename='models-{epoch:02d}-{valid_loss:.2f}',\n    save_top_k=3,\n    mode='min') \n\nmod = CassavaModel()\ndx = CassavaLightDataset()\ntrainer = pl.Trainer(gpus=-1,max_epochs=10)\ntrainer.fit(model=mod,datamodule=dx) ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"TEST_FILE_PATH=\"../input/cassava-leaf-disease-classification/test_images/\"\nclass CassavaTestDataset(Dataset):\n    def __init__(self,path,image_ids,image_size):\n        self.image_ids=image_ids\n        self.image_size=image_size\n        self.path=path\n        \n    def __len__(self):\n        return len(self.image_ids)\n    def __getitem__(self,item):\n        image_ids=self.image_ids[item]\n        image=cv2.imread(self.path+image_ids)\n        image=cv2.resize(image,(self.image_size,self.image_size))\n        return {'x':torch.tensor(image,dtype=torch.float)}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample=pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\ntest_dataset=CassavaTestDataset(TEST_FILE_PATH,sample.image_id,IMG_SIZE)\ntest_loader=DataLoader(test_dataset,batch_size=1,shuffle=False)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fin_y=[]\nfor data in test_loader:\n    y_hat=mod(data['x'].view(-1,3,IMG_SIZE,IMG_SIZE))\n    y_hat=nn.Softmax(dim=-1)(y_hat)\n    y_hat=torch.argmax(y_hat,dim=1)\n    fin_y.append(y_hat.cpu().detach().numpy())\n    \nsample['label']=np.array(fin_y).reshape(-1)\nsample[['image_id','label']].to_csv('submission.csv',index=False)\nsample.head()","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}