{"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":"code","source":"import numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset,DataLoader\nfrom torchvision import transforms,models\nfrom tqdm import tqdm_notebook as tqdm","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2022-05-25T09:17:04.312106Z","iopub.status.busy":"2022-05-25T09:17:04.311788Z","iopub.status.idle":"2022-05-25T09:17:04.318849Z","shell.execute_reply":"2022-05-25T09:17:04.317906Z","shell.execute_reply.started":"2022-05-25T09:17:04.312056Z"},"id":"7pkmtO8vq1XN","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/bengaliai-cv19/train.csv')\ndata0 = pd.read_feather('/kaggle/usr/lib/resize_and_load_with_feather_format_much_faster/train_data_0.feather')\ndata1 = pd.read_feather('/kaggle/usr/lib/resize_and_load_with_feather_format_much_faster/train_data_1.feather')\ndata2 = pd.read_feather('/kaggle/usr/lib/resize_and_load_with_feather_format_much_faster/train_data_2.feather')\ndata3 = pd.read_feather('/kaggle/usr/lib/resize_and_load_with_feather_format_much_faster/train_data_3.feather')","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.execute_input":"2022-05-25T09:17:04.326512Z","iopub.status.busy":"2022-05-25T09:17:04.326246Z","iopub.status.idle":"2022-05-25T09:17:05.142816Z","shell.execute_reply":"2022-05-25T09:17:05.141892Z","shell.execute_reply.started":"2022-05-25T09:17:04.326467Z"},"id":"hIpHLPpBq1XP","outputId":"eff579bb-0ea1-4276-841c-b36f366ebad9","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_full = pd.concat([data0,data1,data2,data3],ignore_index=True)","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:05.952818Z","iopub.status.busy":"2022-05-25T09:17:05.952493Z","iopub.status.idle":"2022-05-25T09:17:06.289173Z","shell.execute_reply":"2022-05-25T09:17:06.288412Z","shell.execute_reply.started":"2022-05-25T09:17:05.952767Z"},"id":"-RjHd5m2q1XQ","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GraphemeDataset(Dataset):\n    def __init__(self,df,label,_type='train'):\n        self.df = df\n        self.label = label\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self,idx):\n        label1 = self.label.vowel_diacritic.values[idx]\n        label2 = self.label.grapheme_root.values[idx]\n        label3 = self.label.consonant_diacritic.values[idx]\n        image = self.df.iloc[idx][1:].values.reshape(64,64).astype(np.float)\n        return image,label1,label2,label3","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.290949Z","iopub.status.busy":"2022-05-25T09:17:06.290620Z","iopub.status.idle":"2022-05-25T09:17:06.299926Z","shell.execute_reply":"2022-05-25T09:17:06.297877Z","shell.execute_reply.started":"2022-05-25T09:17:06.290893Z"},"id":"VVpuPyYrq1XR","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self,in_channels,out_channels,stride=1,kernel_size=3,padding=1,bias=False):\n        super(ResidualBlock,self).__init__()\n        self.cnn1 =nn.Sequential(\n            nn.Conv2d(in_channels,out_channels,kernel_size,stride,padding,bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(True)\n        )\n        self.cnn2 = nn.Sequential(\n            nn.Conv2d(out_channels,out_channels,kernel_size,1,padding,bias=False),\n            nn.BatchNorm2d(out_channels)\n        )\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels,out_channels,kernel_size=1,stride=stride,bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        else:\n            self.shortcut = nn.Sequential()\n            \n    def forward(self,x):\n        residual = x\n        x = self.cnn1(x)\n        x = self.cnn2(x)\n        x += self.shortcut(residual)\n        x = nn.ReLU(True)(x)\n        return x","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.306206Z","iopub.status.busy":"2022-05-25T09:17:06.305926Z","iopub.status.idle":"2022-05-25T09:17:06.316078Z","shell.execute_reply":"2022-05-25T09:17:06.315109Z","shell.execute_reply.started":"2022-05-25T09:17:06.306155Z"},"id":"GAfEbER6q1XR","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet34(nn.Module):\n    def __init__(self):\n        super(ResNet34,self).__init__()\n        \n        self.block1 = nn.Sequential(\n            nn.Conv2d(1,64,kernel_size=2,stride=2,padding=3,bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(True),\n            nn.MaxPool2d(1,1)\n        )\n        \n        self.block2 = nn.Sequential(\n            ResidualBlock(64,64,2),\n            ResidualBlock(64,64),\n            ResidualBlock(64,64)\n        )\n        \n        self.block3 = nn.Sequential(\n            ResidualBlock(64,128,2),\n            ResidualBlock(128,128),\n            ResidualBlock(128,128),\n            ResidualBlock(128,128)\n        )\n        \n        self.block4 = nn.Sequential(\n            ResidualBlock(128,256,2),\n            ResidualBlock(256,256),\n            ResidualBlock(256,256),\n            ResidualBlock(256,256),\n            ResidualBlock(256,256),\n            ResidualBlock(256,256)\n        )\n        self.block5 = nn.Sequential(\n            ResidualBlock(256,512,2),\n            ResidualBlock(512,512),\n            ResidualBlock(512,512)\n        )\n        \n        self.avgpool = nn.AvgPool2d(2)\n        # vowel_diacritic\n        self.fc1 = nn.Linear(512,11)\n        # grapheme_root\n        self.fc2 = nn.Linear(512,168)\n        # consonant_diacritic\n        self.fc3 = nn.Linear(512,7)\n        \n    def forward(self,x):\n        x = self.block1(x)\n        x = self.block2(x)\n        x = self.block3(x)\n        x = self.block4(x)\n        x = self.block5(x)\n        x = self.avgpool(x)\n        x = x.view(x.size(0),-1)\n        x1 = self.fc1(x)\n        x2 = self.fc2(x)\n        x3 = self.fc3(x)\n        return x1,x2,x3","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.336170Z","iopub.status.busy":"2022-05-25T09:17:06.335591Z","iopub.status.idle":"2022-05-25T09:17:06.353274Z","shell.execute_reply":"2022-05-25T09:17:06.352264Z","shell.execute_reply.started":"2022-05-25T09:17:06.336057Z"},"id":"RZeFaokEq1XR","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install torchsummary\n#from torchsummary import summary","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.355263Z","iopub.status.busy":"2022-05-25T09:17:06.354758Z","iopub.status.idle":"2022-05-25T09:17:06.375093Z","shell.execute_reply":"2022-05-25T09:17:06.374280Z","shell.execute_reply.started":"2022-05-25T09:17:06.355004Z"},"id":"8yAhmnHqq1XT","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.377103Z","iopub.status.busy":"2022-05-25T09:17:06.376549Z","iopub.status.idle":"2022-05-25T09:17:06.384627Z","shell.execute_reply":"2022-05-25T09:17:06.383776Z","shell.execute_reply.started":"2022-05-25T09:17:06.376918Z"},"id":"m6YrjZpoq1XT","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model_resnet34 = ResNet34().to(device)\n#summary(model_resnet34, (1, 64, 64))","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.386570Z","iopub.status.busy":"2022-05-25T09:17:06.386045Z","iopub.status.idle":"2022-05-25T09:17:06.393827Z","shell.execute_reply":"2022-05-25T09:17:06.393108Z","shell.execute_reply.started":"2022-05-25T09:17:06.386377Z"},"id":"9KrvM1dUq1XT","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ResNet34().to(device)\noptimizer = optimizer = torch.optim.Adam(model.parameters(), lr=4e-4)\n#scheduler = torch.optim.lr_scheduler.CyclicLR(optimizer, base_lr=1e-4, max_lr=0.05)\ncriterion = nn.CrossEntropyLoss()\nbatch_size=32","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.395193Z","iopub.status.busy":"2022-05-25T09:17:06.394945Z","iopub.status.idle":"2022-05-25T09:17:06.573026Z","shell.execute_reply":"2022-05-25T09:17:06.572191Z","shell.execute_reply.started":"2022-05-25T09:17:06.395143Z"},"id":"HXfDi8tgq1XU","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 50 # original 50\nmodel.train()\nlosses = []\naccs = []\nfor epoch in range(epochs):\n    reduced_index =train.groupby(['grapheme_root', 'vowel_diacritic', 'consonant_diacritic']).apply(lambda x: x.sample(5)).image_id.values\n    reduced_train = train.loc[train.image_id.isin(reduced_index)]\n    reduced_data = data_full.loc[data_full.image_id.isin(reduced_index)]\n    train_image = GraphemeDataset(reduced_data,reduced_train)\n    train_loader = torch.utils.data.DataLoader(train_image,batch_size=batch_size,shuffle=True)\n    \n    print('epochs {}/{} '.format(epoch+1,epochs))\n    running_loss = 0.0\n    running_acc = 0.0\n    for idx, (inputs,labels1,labels2,labels3) in tqdm(enumerate(train_loader),total=len(train_loader)):\n        inputs = inputs.to(device)\n        labels1 = labels1.to(device)\n        labels2 = labels2.to(device)\n        labels3 = labels3.to(device)\n        \n        optimizer.zero_grad()\n        outputs1,outputs2,outputs3 = model(inputs.unsqueeze(1).float())\n        loss1 = criterion(outputs1,labels1)\n        loss2 = criterion(outputs2,labels2)\n        loss3 = criterion(outputs3,labels3)\n        running_loss += loss1+loss2+loss3\n        running_acc += (outputs1.argmax(1)==labels1).float().mean()\n        running_acc += (outputs2.argmax(1)==labels2).float().mean()\n        running_acc += (outputs3.argmax(1)==labels3).float().mean()\n        (loss1+loss2+loss3).backward()\n        optimizer.step()\n    #scheduler.step()\n    losses.append(running_loss/len(train_loader))\n    accs.append(running_acc/(len(train_loader)*3))\n    print('acc : {:.2f}'.format(running_acc/(len(train_loader)*3)))\n    print('loss : {:.4f}'.format(running_loss/len(train_loader)))\ntorch.save(model.state_dict(), 'resnet34_50epochs_saved_weights.pth')","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:17:06.575044Z","iopub.status.busy":"2022-05-25T09:17:06.574598Z","iopub.status.idle":"2022-05-25T09:47:41.673665Z","shell.execute_reply":"2022-05-25T09:47:41.672980Z","shell.execute_reply.started":"2022-05-25T09:17:06.574994Z"},"id":"saDpOjAwq1XU","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfig,ax = plt.subplots(1,2,figsize=(15,5))\nax[0].plot(losses)\nax[0].set_title('loss')\nax[1].plot(accs)\nax[1].set_title('acc')","metadata":{"execution":{"iopub.execute_input":"2022-05-25T09:47:41.677069Z","iopub.status.busy":"2022-05-25T09:47:41.676839Z","iopub.status.idle":"2022-05-25T09:47:42.122720Z","shell.execute_reply":"2022-05-25T09:47:42.121832Z","shell.execute_reply.started":"2022-05-25T09:47:41.677026Z"},"id":"kePHK59bq1XV","trusted":true},"execution_count":null,"outputs":[]}]}