{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport sys\nfrom scipy import misc\nfrom glob import glob\n\nimport albumentations as albu\nfrom albumentations.pytorch import ToTensor\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset, sampler\nimport cv2\nfrom tqdm import tqdm\n\nimport os\n        \nsys.path.append('/kaggle/input/srnet-model-weight/')\n        \nfrom model import Srnet\n# You can write up to 5GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"BATCH_SIZE = 40\nTESTPATH = '/kaggle/input/alaska2-image-steganalysis/Test/'\nWEIGHTS =  '/kaggle/input/srnet-model-weight/SRNet_model_weights.pt'\n\ndf_sub = pd.read_csv('/kaggle/input/alaska2-image-steganalysis/sample_submission.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def transform_test():\n    transform = albu.Compose([\n        albu.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n        ToTensor()\n    ])\n    return transform","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_sub","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class AlaskaDataset(Dataset):\n    def __init__(self, df, data_folder, transform):\n        self.df = df\n        self.root = data_folder\n        self._transform = transform\n        \n    def __getitem__(self, idx):\n        image_id = self.df.Id.iloc[idx]\n        print(image_id)\n        image_path = os.path.join(self.root, image_id)\n        img = cv2.imread(image_path)\n        augment = self._transform(image=img)\n        img = augment['image']\n        img = torch.mean(img, axis=0,keepdim=True) #mean all channels because of SRNet input format\n        return img\n    \n    def __len__(self):\n        return len(self.df)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = Srnet().cuda()\nweights = torch.load(WEIGHTS)\nmodel.load_state_dict(weights['model_state_dict'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_transform = transform_test()\ntest_dataset = AlaskaDataset(df_sub, TESTPATH, test_transform)\ntest_data = DataLoader(\n    test_dataset,\n    batch_size = BATCH_SIZE,\n    num_workers = 2,\n    shuffle = False\n)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from tqdm import tqdm\noutputs = []\nmodel.eval() #turn model into eval mode before inference\nwith torch.no_grad():\n    for inputs in tqdm(test_data):\n        inputs = inputs.cuda()\n        output = model(inputs)\n        pred = output.data.cpu().numpy()\n        pred = np.exp(pred[:,1]) / (np.exp(pred[:,0]) + np.exp(pred[:,1]))\n        outputs.append(pred)\noutputs = np.concatenate(outputs)\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_sub['Label']  = outputs\ndf_sub.to_csv('submissions.csv', index=None)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_sub.head()","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":4}