{"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\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport torch\n\n# Thanks to @Doomsday for the wheels!\n!pip install /kaggle/input/segmodelpytorchwheel/wheel/timm-0.6.12-py3-none-any.whl\n!pip install /kaggle/input/segmodelpytorchwheel/wheel/efficientnet_pytorch-0.7.1-py3-none-any.whl\n!pip install /kaggle/input/segmodelpytorchwheel/wheel/pretrainedmodels-0.7.4-py3-none-any.whl\n!pip install /kaggle/input/segmodelpytorchwheel/wheel/segmentation_models_pytorch-0.3.2-py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-04T15:07:53.234943Z","iopub.execute_input":"2023-06-04T15:07:53.23529Z","iopub.status.idle":"2023-06-04T15:10:06.168684Z","shell.execute_reply.started":"2023-06-04T15:07:53.235241Z","shell.execute_reply":"2023-06-04T15:10:06.167139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get Device","metadata":{}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n    \ndevice","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:10:06.172091Z","iopub.execute_input":"2023-06-04T15:10:06.173083Z","iopub.status.idle":"2023-06-04T15:10:06.232095Z","shell.execute_reply.started":"2023-06-04T15:10:06.173041Z","shell.execute_reply":"2023-06-04T15:10:06.230856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Model","metadata":{}},{"cell_type":"code","source":"# Path for data\ndata_path = Path('/kaggle/input/google-research-identify-contrails-reduce-global-warming')\n# Path for model\nmodel_path = '/kaggle/input/unet-baseline-model/model_epoch_19_dice_0.5647.pt'\n\n#Load model\nmodel = torch.load(model_path,map_location=device)\nmodel.eval()\n\nresize=True\nresize_size = 384\n","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.370427Z","iopub.execute_input":"2023-06-04T15:12:53.371012Z","iopub.status.idle":"2023-06-04T15:12:53.732816Z","shell.execute_reply.started":"2023-06-04T15:12:53.370963Z","shell.execute_reply":"2023-06-04T15:12:53.731677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Function to load and parse the data","metadata":{}},{"cell_type":"code","source":"\n\ndef load_and_parse_data(path):\n    N_TIMES_BEFORE = 4\n\n    with open(os.path.join(path, 'band_11.npy'), 'rb') as f:\n        band11 = np.load(f)\n    with open(os.path.join(path, 'band_14.npy'), 'rb') as f:\n        band14 = np.load(f)\n    with open(os.path.join(path, 'band_15.npy'), 'rb') as f:\n        band15 = np.load(f)\n\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n\n    def normalize_range(data, bounds):\n        \"\"\"Maps data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\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    false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n    false_color = false_color[...,4]\n\n    return false_color","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.734707Z","iopub.execute_input":"2023-06-04T15:12:53.735078Z","iopub.status.idle":"2023-06-04T15:12:53.746471Z","shell.execute_reply.started":"2023-06-04T15:12:53.735042Z","shell.execute_reply":"2023-06-04T15:12:53.745353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RLE Encoding as provided by competition hosts","metadata":{}},{"cell_type":"code","source":"def rle_encode(x, fg_val=1):\n    \"\"\"\n    Args:\n        x:  numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns: run length encoding as list\n    \"\"\"\n\n    dots = np.where(\n        x.T.flatten() == fg_val)[0]  # .T sets Fortran order down-then-right\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 run_lengths\n\n\ndef list_to_string(x):\n    \"\"\"\n    Converts list to a string representation\n    Empty list returns '-'\n    \"\"\"\n    if x: # non-empty list\n        s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n    else:\n        s = '-'\n    return s\n\n\ndef rle_decode(mask_rle, shape=(256, 256)):\n    '''\n    mask_rle: run-length as string formatted (start length)\n              empty predictions need to be encoded with '-'\n    shape: (height, width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    if mask_rle != '-': \n        s = mask_rle.split()\n        starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n        starts -= 1\n        ends = starts + lengths\n        for lo, hi in zip(starts, ends):\n            img[lo:hi] = 1\n    return img.reshape(shape, order='F')  # Needed to align to RLE direction","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.748043Z","iopub.execute_input":"2023-06-04T15:12:53.749219Z","iopub.status.idle":"2023-06-04T15:12:53.762845Z","shell.execute_reply.started":"2023-06-04T15:12:53.749176Z","shell.execute_reply":"2023-06-04T15:12:53.761785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function for predicting\ndef predict_data(path):\n    false_color = load_and_parse_data(path)\n    false_color = np.expand_dims(false_color,0)\n    false_color = torch.from_numpy(false_color)\n    \n    false_color=torch.moveaxis(false_color,-1,1)\n    \n\n    # Resize the false_color img if resize was activated during training\n    if resize:\n        false_color = torch.nn.functional.interpolate(false_color, \n                                            size=384,\n                                            mode='bilinear'\n                                           )\n    \n    \n    false_color = false_color.to(device)\n    pred = model.predict(false_color)\n    \n    if resize:\n        pred = torch.nn.functional.interpolate(pred, \n                                    size=256,\n                                    mode='bilinear'\n                                   )\n    \n    pred = torch.moveaxis(pred, 1, -1)\n    pred = pred[0]\n    pred = pred[...,0]\n    \n    pred[pred > 0.5] = 1\n    pred[pred<=0.5]=0\n    \n    pred = pred.cpu()\n    \n    return pred\n    ","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.765774Z","iopub.execute_input":"2023-06-04T15:12:53.766317Z","iopub.status.idle":"2023-06-04T15:12:53.777229Z","shell.execute_reply.started":"2023-06-04T15:12:53.766278Z","shell.execute_reply":"2023-06-04T15:12:53.776167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make the submission","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv(data_path / 'sample_submission.csv', index_col='record_id')\ntest = 'test'\ntest_recs = os.listdir(data_path / test)\nfor record in test_recs:\n    path = os.path.join(data_path,test, record)\n    predicted = predict_data(path)\n    submission.loc[int(record), 'encoded_pixels'] = list_to_string(rle_encode(predicted))\n\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.778823Z","iopub.execute_input":"2023-06-04T15:12:53.779345Z","iopub.status.idle":"2023-06-04T15:12:53.889383Z","shell.execute_reply.started":"2023-06-04T15:12:53.779304Z","shell.execute_reply":"2023-06-04T15:12:53.888172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-06-04T15:12:53.891046Z","iopub.execute_input":"2023-06-04T15:12:53.891735Z","iopub.status.idle":"2023-06-04T15:12:53.898261Z","shell.execute_reply.started":"2023-06-04T15:12:53.891689Z","shell.execute_reply":"2023-06-04T15:12:53.89706Z"},"trusted":true},"execution_count":null,"outputs":[]}]}