{"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 os\n\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset as BaseDataset\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torchvision import transforms\nfrom torchvision.transforms import functional\nfrom tqdm import tqdm_notebook as tqdm\nimport sys\nsys.path.append('/kaggle/input/add-liblaly-for-smp')\nimport segmentation_models_pytorch as smp\nimport cv2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-25T07:30:32.545533Z","iopub.execute_input":"2023-06-25T07:30:32.545876Z","iopub.status.idle":"2023-06-25T07:30:39.918293Z","shell.execute_reply.started":"2023-06-25T07:30:32.54585Z","shell.execute_reply":"2023-06-25T07:30:39.917343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_TIMES_BEFORE = 4\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\npath = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\"\n\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])\n\ndef normalize_std(spec):\n    return (spec- np.mean(spec))/np.std(spec)\n\nclass Dataset(BaseDataset):\n    def __init__(self, dir_path_list):\n        self.dir_path_list = dir_path_list\n\n    def __getitem__(self, i):\n        \n        band11 = np.load(os.path.join(path, self.dir_path_list[i], 'band_11.npy'))\n        band14 = np.load(os.path.join(path, self.dir_path_list[i], 'band_14.npy'))\n        band15 = np.load(os.path.join(path, self.dir_path_list[i], 'band_15.npy'))\n        \n        r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = normalize_range(band14, _T11_BOUNDS)\n        x = np.clip(np.stack([r, g, b], axis=2), 0, 1)[:,:,:,N_TIMES_BEFORE].astype(np.float32).transpose(2,0,1)      \n        y = self.dir_path_list[i]\n        \n        return x, y\n\n    def __len__(self):\n        return len(self.dir_path_list)","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:26.200322Z","iopub.execute_input":"2023-06-21T02:30:26.200709Z","iopub.status.idle":"2023-06-21T02:30:26.212485Z","shell.execute_reply.started":"2023-06-21T02:30:26.200676Z","shell.execute_reply":"2023-06-21T02:30:26.211143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 1\n\n# フォルダリストの読み込み\nfiles = os.listdir(path)\nfiles_dir = [f for f in files if os.path.isdir(os.path.join(path, f))]\n\ntest_dataset = Dataset(files_dir)\ntest_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, num_workers=0, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:26.214405Z","iopub.execute_input":"2023-06-21T02:30:26.21484Z","iopub.status.idle":"2023-06-21T02:30:26.231781Z","shell.execute_reply.started":"2023-06-21T02:30:26.214802Z","shell.execute_reply":"2023-06-21T02:30:26.2309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(y_pred, fg_val=1):\n    def list_to_string(x):\n        if x:\n            s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n        else:\n            s = '-'\n        return s\n    dots = np.where(\n        y_pred.T.flatten() == fg_val)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return list_to_string(run_lengths)","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:26.234086Z","iopub.execute_input":"2023-06-21T02:30:26.234493Z","iopub.status.idle":"2023-06-21T02:30:26.248136Z","shell.execute_reply.started":"2023-06-21T02:30:26.234462Z","shell.execute_reply":"2023-06-21T02:30:26.246487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', index_col='record_id')\n\n# モデルのロード\nmodel = torch.load(\"/kaggle/input/unet-model/train_10.pth\", torch.device('cpu'))\nmodel.eval()\n\nfor data, labels in test_loader:\n    outputs = model(data)\n    pred = torch.where(outputs>0.5, 1, 0)\n    submission.loc[int(labels[0]), 'encoded_pixels'] = rle_encode(pred[0,0,...])","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:31.832334Z","iopub.execute_input":"2023-06-21T02:30:31.832704Z","iopub.status.idle":"2023-06-21T02:30:33.718836Z","shell.execute_reply.started":"2023-06-21T02:30:31.832669Z","shell.execute_reply":"2023-06-21T02:30:33.717764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:43.835585Z","iopub.execute_input":"2023-06-21T02:30:43.835946Z","iopub.status.idle":"2023-06-21T02:30:43.86017Z","shell.execute_reply.started":"2023-06-21T02:30:43.83592Z","shell.execute_reply":"2023-06-21T02:30:43.859143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-06-21T02:30:45.106136Z","iopub.execute_input":"2023-06-21T02:30:45.106955Z","iopub.status.idle":"2023-06-21T02:30:45.117064Z","shell.execute_reply.started":"2023-06-21T02:30:45.106901Z","shell.execute_reply":"2023-06-21T02:30:45.11602Z"},"trusted":true},"execution_count":null,"outputs":[]}]}