{"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\nimport torchvision.transforms as T\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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":132.440666,"end_time":"2023-06-18T13:50:31.228166","exception":false,"start_time":"2023-06-18T13:48:18.7875","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:47:05.221607Z","iopub.execute_input":"2023-06-25T10:47:05.221939Z","iopub.status.idle":"2023-06-25T10:49:18.221216Z","shell.execute_reply.started":"2023-06-25T10:47:05.221889Z","shell.execute_reply":"2023-06-25T10:49:18.219819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get Device","metadata":{"papermill":{"duration":0.006111,"end_time":"2023-06-18T13:50:31.242001","exception":false,"start_time":"2023-06-18T13:50:31.23589","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n    \ndevice","metadata":{"papermill":{"duration":0.08087,"end_time":"2023-06-18T13:50:31.329016","exception":false,"start_time":"2023-06-18T13:50:31.248146","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:18.225067Z","iopub.execute_input":"2023-06-25T10:49:18.226344Z","iopub.status.idle":"2023-06-25T10:49:18.309101Z","shell.execute_reply.started":"2023-06-25T10:49:18.226305Z","shell.execute_reply":"2023-06-25T10:49:18.307806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\nmodel = smp.Unet(\n    encoder_name = 'timm-resnest26d',\n    encoder_weights=None,    # use `imagenet` pre-trained weights for encoder initialization\n    in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n    classes=1,        # model output channels (number of classes in your dataset)\n    activation=\"sigmoid\",\n    )","metadata":{"execution":{"iopub.status.busy":"2023-06-25T10:49:18.312091Z","iopub.execute_input":"2023-06-25T10:49:18.313332Z","iopub.status.idle":"2023-06-25T10:49:21.011577Z","shell.execute_reply.started":"2023-06-25T10:49:18.313293Z","shell.execute_reply":"2023-06-25T10:49:21.010386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Model","metadata":{"papermill":{"duration":0.00596,"end_time":"2023-06-18T13:50:31.341497","exception":false,"start_time":"2023-06-18T13:50:31.335537","status":"completed"},"tags":[]}},{"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/simple-baseline-train/model_state_dict_epoch_19_dice_0.5922.pth'\n\n#Load model\nmodel.load_state_dict(torch.load(model_path,map_location=device))\nmodel = model.to(device)\nmodel.eval()\n\nresize=False\n\nresize_size = 256\nif resize:\n    resize_size = 384","metadata":{"papermill":{"duration":5.968223,"end_time":"2023-06-18T13:50:37.315936","exception":false,"start_time":"2023-06-18T13:50:31.347713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:21.014847Z","iopub.execute_input":"2023-06-25T10:49:21.015371Z","iopub.status.idle":"2023-06-25T10:49:25.435623Z","shell.execute_reply.started":"2023-06-25T10:49:21.015328Z","shell.execute_reply":"2023-06-25T10:49:25.434408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Function to load and parse the data","metadata":{"papermill":{"duration":0.006211,"end_time":"2023-06-18T13:50:37.329081","exception":false,"start_time":"2023-06-18T13:50:37.32287","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.021969,"end_time":"2023-06-18T13:50:37.357594","exception":false,"start_time":"2023-06-18T13:50:37.335625","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:25.437412Z","iopub.execute_input":"2023-06-25T10:49:25.437791Z","iopub.status.idle":"2023-06-25T10:49:25.4495Z","shell.execute_reply.started":"2023-06-25T10:49:25.43775Z","shell.execute_reply":"2023-06-25T10:49:25.448126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RLE Encoding as provided by competition hosts","metadata":{"papermill":{"duration":0.006122,"end_time":"2023-06-18T13:50:37.369972","exception":false,"start_time":"2023-06-18T13:50:37.36385","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.023885,"end_time":"2023-06-18T13:50:37.400008","exception":false,"start_time":"2023-06-18T13:50:37.376123","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:25.451655Z","iopub.execute_input":"2023-06-25T10:49:25.45218Z","iopub.status.idle":"2023-06-25T10:49:25.466988Z","shell.execute_reply.started":"2023-06-25T10:49:25.452098Z","shell.execute_reply":"2023-06-25T10:49:25.465856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function for predicting\ndef predict_data(path):\n    normalize = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\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    false_color = normalize(false_color)\n    false_color = false_color.to(device)\n    pred = model(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.01] = 1\n    pred[pred<=0.01]=0\n    \n    pred = pred.cpu()\n    \n    return pred\n    ","metadata":{"papermill":{"duration":0.019474,"end_time":"2023-06-18T13:50:37.425727","exception":false,"start_time":"2023-06-18T13:50:37.406253","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:25.468695Z","iopub.execute_input":"2023-06-25T10:49:25.469572Z","iopub.status.idle":"2023-06-25T10:49:25.480489Z","shell.execute_reply.started":"2023-06-25T10:49:25.469534Z","shell.execute_reply":"2023-06-25T10:49:25.479361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make the submission","metadata":{"papermill":{"duration":0.006126,"end_time":"2023-06-18T13:50:37.438099","exception":false,"start_time":"2023-06-18T13:50:37.431973","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":2.589033,"end_time":"2023-06-18T13:50:40.033418","exception":false,"start_time":"2023-06-18T13:50:37.444385","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:25.482131Z","iopub.execute_input":"2023-06-25T10:49:25.482604Z","iopub.status.idle":"2023-06-25T10:49:27.949407Z","shell.execute_reply.started":"2023-06-25T10:49:25.482567Z","shell.execute_reply":"2023-06-25T10:49:27.948387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"papermill":{"duration":0.0178,"end_time":"2023-06-18T13:50:40.057783","exception":false,"start_time":"2023-06-18T13:50:40.039983","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T10:49:27.950976Z","iopub.execute_input":"2023-06-25T10:49:27.951335Z","iopub.status.idle":"2023-06-25T10:49:27.958836Z","shell.execute_reply.started":"2023-06-25T10:49:27.951297Z","shell.execute_reply":"2023-06-25T10:49:27.957502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}