{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport torchvision\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom PIL import Image\nimport os\nfrom functools import reduce","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"test_df = pd.read_csv('/kaggle/input/bengaliai-cv19/test.csv')\nsubmission_df = pd.read_csv('/kaggle/input/bengaliai-cv19/sample_submission.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class TestDataset(torch.utils.data.Dataset):\n\n    def __init__(self, test_images, image_ids, transforms=None):\n        super(TestDataset, self).__init__()\n        self.image_ids = image_ids\n        self.test_images = test_images\n        self.transforms = transforms\n\n    def __getitem__(self, index):\n        image_id = self.image_ids.iloc[index]\n        img_array = np.zeros((137, 236, 3), dtype='uint8')\n        img_array[:, :, 0] = self.test_images[index].reshape(137, 236)\n        img_array[:, :, 1] = img_array[:, :, 0]\n        img_array[:, :, 2] = img_array[:, :, 0]\n        img = Image.fromarray(img_array)\n        if self.transforms:\n            img = self.transforms(img)\n\n        return img\n\n    def __len__(self,):\n        return len(self.image_ids)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class GraphemeModel(nn.Module):\n\n    def __init__(self):\n        super(GraphemeModel, self).__init__()\n        self.base_model = torchvision.models.resnet18(pretrained=False) # use resnet18 as the base model\n#         self.fc = nn.Linear(1000, 256) \n        self.fc_root = nn.Linear(1000, 168)\n        self.fc_vowel = nn.Linear(1000, 11)\n        self.fc_consonant = nn.Linear(1000, 7)\n        \n    def forward(self, inp):\n        x = self.base_model(inp)\n        x = x.view(x.shape[0], -1)\n#         x = F.relu(self.fc(x))\n        root_output = self.fc_root(x)\n        vowel_output = self.fc_vowel(x)\n        consonant_output = self.fc_consonant(x)\n\n        return (root_output, vowel_output, consonant_output)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = torch.load('/kaggle/input/bengaligraphememodel3/model.pth')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"transforms = torchvision.transforms.Compose([\n                              torchvision.transforms.ToTensor(),\n                              torchvision.transforms.Normalize((0, 0, 0), (1., 1., 1.))                \n            ])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Perform predictions"},{"metadata":{"trusted":true},"cell_type":"code","source":"model.eval()\npredictions = []\n\nfor i in range(4):\n    test_image_data = pd.read_parquet(f'/kaggle/input/bengaliai-cv19/test_image_data_{i}.parquet') # read ith test parquet file\n    test_matrix = test_image_data.drop(columns=['image_id']).values\n    image_ids = test_image_data.image_id\n    test_dataset = TestDataset(test_images=test_matrix, image_ids=image_ids, transforms=transforms)\n    test_dataloader = torch.utils.data.DataLoader(dataset=test_dataset, batch_size=128, shuffle=False)\n    for x in test_dataloader:\n        root, vowel, consonant = model(x.to(device)) # get prediction for the batch\n        root = root.argmax(1).detach().cpu().numpy() # convert to numpy\n        vowel = vowel.argmax(1).detach().cpu().numpy()\n        consonant = consonant.argmax(1).detach().cpu().numpy()\n        predictions += list(reduce(lambda a, b: a + b, zip(consonant, root, vowel)))\n    ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Attach predictions to the dataframe"},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df['target'] = predictions","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Save the dataframe to submission.csv file"},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df.to_csv('submission.csv', index=False)","execution_count":null,"outputs":[]},{"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":1}