{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":2580782,"sourceType":"datasetVersion","datasetId":1404099},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":154122325,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!ls /kaggle/input/pyvips-python-and-deb-package-gpu\n# intall the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-08T13:50:47.036929Z","iopub.execute_input":"2023-12-08T13:50:47.037215Z","iopub.status.idle":"2023-12-08T13:51:51.616788Z","shell.execute_reply.started":"2023-12-08T13:50:47.03719Z","shell.execute_reply":"2023-12-08T13:51:51.61582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r /kaggle/input/swintransformertf /kaggle/working/\n    \nimport sys\nsys.path.append('/kaggle/working/swintransformertf')","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:51:51.618799Z","iopub.execute_input":"2023-12-08T13:51:51.619108Z","iopub.status.idle":"2023-12-08T13:51:52.598221Z","shell.execute_reply.started":"2023-12-08T13:51:51.619079Z","shell.execute_reply":"2023-12-08T13:51:52.597114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport random\nimport glob\nimport gc\nimport os\nimport re\nimport tensorflow as tf\nfrom tensorflow.keras import callbacks\nfrom tensorflow.keras.layers import Dense\nfrom tensorflow.keras.metrics import SparseCategoricalAccuracy\nfrom sklearn.model_selection import KFold\nfrom swintransformer import SwinTransformer\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\nimport pyvips\nfrom PIL import Image\n\n\nAUTO = tf.data.AUTOTUNE","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-08T13:51:52.599585Z","iopub.execute_input":"2023-12-08T13:51:52.599895Z","iopub.status.idle":"2023-12-08T13:52:04.385966Z","shell.execute_reply.started":"2023-12-08T13:51:52.599866Z","shell.execute_reply":"2023-12-08T13:52:04.385105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 402\nIMAGE_SIZE = [384, 384]\nEPOCHS = 20\nBATCH_SIZE = 4\n\nSWIN_TYPE = 'large'\n\nCLASSES = {\n        'Other': 0,\n        'CC': 1,\n        'EC': 2,\n        'HGSC': 3,\n        'LGSC': 4,\n        'MC': 5\n    }","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.388037Z","iopub.execute_input":"2023-12-08T13:52:04.38854Z","iopub.status.idle":"2023-12-08T13:52:04.397097Z","shell.execute_reply.started":"2023-12-08T13:52:04.388512Z","shell.execute_reply":"2023-12-08T13:52:04.3928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Source: https://www.kaggle.com/code/jirkaborovec/cancer-subtype-lightning-torch-inference-tiles\ndef extract_image_tiles(\n    p_img, folder, size: int = 2048, new_size: int = 512, #scale: float = 0.5,\n    drop_thr: float = 0.6, white_thr: int = 240, max_samples: int = 10\n) -> list:\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    im = pyvips.Image.new_from_file(p_img)\n    w = h = size\n    # https://stackoverflow.com/a/47581978/4521646\n    idxs = [(y, y + h, x, x + w) for y in range(0, im.height, h) for x in range(0, im.width, w)]\n    # random subsample\n    max_samples = max_samples if isinstance(max_samples, int) else int(len(idxs) * max_samples)\n    random.shuffle(idxs)\n    files = []\n    for y, y_, x, x_ in idxs:\n        # https://libvips.github.io/pyvips/vimage.html#pyvips.Image.crop\n        tile = im.crop(x, y, min(w, im.width - x), min(h, im.height - y)).numpy()[..., :3]\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile_size = (h, w) if tile.ndim == 2 else (h, w, tile.shape[2])\n            tile = np.zeros(tile_size, dtype=tile.dtype)\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        black_bg = np.sum(tile, axis=2) == 0\n        tile[black_bg, :] = 255\n        mask_bg = np.mean(tile, axis=2) > white_thr\n        if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * drop_thr):\n            #print(f\"skip almost empty tile: {k:06}_{int(x_ / w)}-{int(y_ / h)}\")\n            continue\n        p_img = os.path.join(folder, f\"{int(x_ / w)}-{int(y_ / h)}.png\")\n        # print(tile.shape, tile.dtype, tile.min(), tile.max())\n        #new_size = int(size * scale), int(size * scale)\n        Image.fromarray(tile).resize((new_size, new_size), Image.LANCZOS).save(p_img)\n        files.append(p_img)\n        # need to set counter check as some empty tiles could be skipped earlier\n        if len(files) >= max_samples:\n            break\n    return files\n\ndef extract_prune_tiles(\n    path_img: str, folder: str, size: int = 2048, new_size: int = 512, #scale: float = 0.25,\n    drop_thr: float = 0.6, max_samples: int = 30\n) -> str:\n    print(f\"processing: {path_img}\")\n    name, _ = os.path.splitext(os.path.basename(path_img))\n    folder = os.path.join(folder, name)\n    os.makedirs(folder, exist_ok=True)\n    tiles = extract_image_tiles(\n        path_img, folder, size=size, new_size=new_size,\n        drop_thr=drop_thr, max_samples=max_samples)\n    return folder","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.398393Z","iopub.execute_input":"2023-12-08T13:52:04.398818Z","iopub.status.idle":"2023-12-08T13:52:04.859548Z","shell.execute_reply.started":"2023-12-08T13:52:04.398779Z","shell.execute_reply":"2023-12-08T13:52:04.858694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_test_data(img_dir):\n    images = []\n    img_paths = glob.glob(f'{img_dir}/*.png')\n    for i in range(len(img_paths)):\n        img = Image.open(img_paths[i])\n        #if IMG_SIZE != 512:\n        #    img = img.resize((IMG_SIZE, IMG_SIZE))\n        img = np.array(img)[..., :3]\n        black_bg = np.sum(img, axis=2) == 0\n        img[black_bg, :] = 255\n        images.append(img)\n        \n    return np.array(images)","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.860682Z","iopub.execute_input":"2023-12-08T13:52:04.861009Z","iopub.status.idle":"2023-12-08T13:52:04.867227Z","shell.execute_reply.started":"2023-12-08T13:52:04.860982Z","shell.execute_reply":"2023-12-08T13:52:04.866305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _image_feature(value):\n    \"\"\"Returns a bytes_list from a string / byte.\"\"\"\n    return tf.train.Feature(\n        bytes_list=tf.train.BytesList(value=[tf.io.encode_png(value).numpy()])\n    )\n\ndef _bytes_feature(value):\n    \"\"\"Returns a bytes_list from a string / byte.\"\"\"\n    if isinstance(value, type(tf.constant(0))):\n        value = value.numpy() # BytesList won't unpack a string from an EagerTensor.\n    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))\n\ndef _float_feature(value):\n    \"\"\"Returns a float_list from a float / double.\"\"\"\n    return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))\n\ndef _int64_feature(value):\n    \"\"\"Returns an int64_list from a bool / enum / int / uint.\"\"\"\n    return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))\n\ndef serialize_example(data):\n    feature = {\n      'image_id': _int64_feature(data['image_id']),\n      'image': _image_feature(data['image'])\n    }\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n    return example_proto.SerializeToString()\n\ndef test_convert_to_tfrec(image_ids, X, out_dir, n_rec=10000):\n    out_dir = Path(out_dir)\n    if not out_dir.is_dir():\n        out_dir.mkdir(parents=True)\n    for i in range(1, int(np.ceil(X.shape[0]/n_rec))+1):\n        i1 = (i-1)*n_rec\n        if i*n_rec > X.shape[0]:\n            i2 = X.shape[0]\n        else:\n            i2 = i*n_rec\n        X_part = X[i1:i2, ...]\n\n        filepath = Path(out_dir)/f'test-{str(i-1).zfill(2)}-{X_part.shape[0]}.tfrec'\n        with tf.io.TFRecordWriter(str(filepath)) as writer:\n            for i in range(len(X_part)):\n                data = {}\n                data['image_id'] = image_ids[i]\n                data['image'] = X_part[i, ...]\n\n                writer.write(serialize_example(data))\n                \n        filesize = filepath.stat().st_size /10**6\n        #print(filepath.name,':',np.around(filesize, 2),'MB')","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.868603Z","iopub.execute_input":"2023-12-08T13:52:04.869027Z","iopub.status.idle":"2023-12-08T13:52:04.883731Z","shell.execute_reply.started":"2023-12-08T13:52:04.868991Z","shell.execute_reply":"2023-12-08T13:52:04.882941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_image(image_data):\n    image = tf.image.decode_png(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image\n\ndef prepare_label(label):    \n    label = tf.cast(label, tf.int32)            \n    label = tf.reshape(label, [1])         \n    return label\n\ndef read_labeled_tfrecord(example, return_image_id=False):\n    tfrec_format = {\n        'image_id': tf.io.FixedLenFeature([], tf.int64),\n        'image': tf.io.FixedLenFeature([], tf.string),\n        'label': tf.io.FixedLenFeature([], tf.int64)\n    }\n    example = tf.io.parse_single_example(example, tfrec_format)\n    \n    image = prepare_image(example['image'])\n    \n    #image_id = example['image_id']\n    \n    label = prepare_label(example['label'])\n    return image, label# if not return_image_id else image_id\n\ndef read_unlabeled_tfrecord(example, return_image_id=False):\n    tfrec_format = {\n        'image_id': tf.io.FixedLenFeature([], tf.int64),\n        'image': tf.io.FixedLenFeature([], tf.string)\n    }\n    example = tf.io.parse_single_example(example, tfrec_format)\n    image = prepare_image(example['image'])\n    image_id = example['image_id']\n    if return_image_id:\n        return image, image_id\n    else:\n        return image\n    \ndef load_dataset(filenames, return_wave_id=False, labeled=True, ordered=False):\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO)\n    dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(lambda example: read_labeled_tfrecord(example, return_wave_id) \\\n                          if labeled else read_unlabeled_tfrecord(example, return_wave_id), \n                          num_parallel_calls=AUTO)\n    return dataset\n\ndef get_test_dataset(test_filenames, batch_size=BATCH_SIZE, ordered=True):\n    dataset = load_dataset(test_filenames, labeled=False, ordered=ordered)\n    dataset = dataset.batch(batch_size)\n    dataset = dataset.prefetch(AUTO)\n    return dataset\n\ndef count_data_items(filenames):\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1)) for filename in filenames]\n    return np.sum(n)","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.885068Z","iopub.execute_input":"2023-12-08T13:52:04.885565Z","iopub.status.idle":"2023-12-08T13:52:04.899774Z","shell.execute_reply.started":"2023-12-08T13:52:04.885539Z","shell.execute_reply":"2023-12-08T13:52:04.898818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model():\n    model = tf.keras.Sequential([\n        SwinTransformer(f'swin_{SWIN_TYPE}_{IMAGE_SIZE[0]}', \n                                     include_top=False, \n                                     pretrained=False),\n        Dense(len(CLASSES), activation='softmax')\n    ])\n\n    model.compile(\n        optimizer='adam',\n        loss = 'sparse_categorical_crossentropy',\n        metrics=[SparseCategoricalAccuracy()]\n    ) \n    return model","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.900937Z","iopub.execute_input":"2023-12-08T13:52:04.90127Z","iopub.status.idle":"2023-12-08T13:52:04.914017Z","shell.execute_reply.started":"2023-12-08T13:52:04.901236Z","shell.execute_reply":"2023-12-08T13:52:04.913229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\n\ndef remove_folder(rm_dir):\n    for root, dirs, files in os.walk(rm_dir):\n        for name in files:\n            os.remove(os.path.join(root, name))\n        for name in dirs:\n            shutil.rmtree(os.path.join(root, name))","metadata":{"execution":{"iopub.status.busy":"2023-12-08T13:52:04.917212Z","iopub.execute_input":"2023-12-08T13:52:04.917895Z","iopub.status.idle":"2023-12-08T13:52:04.925269Z","shell.execute_reply.started":"2023-12-08T13:52:04.917869Z","shell.execute_reply":"2023-12-08T13:52:04.924407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = '/kaggle/input/UBC-OCEAN'\nimg_dir = '/kaggle/working/test_tiles'\nout_dir = '/kaggle/working'\n\ntest_df = pd.read_csv(os.path.join(data_dir, 'test.csv'))\ntest_filenames = sorted(glob.glob(os.path.join(data_dir, \"test_images\", '*.png')))\nprint(f\"found images: {len(test_filenames)}\")\n\nmodel = load_model()\nmodel.load_weights(f'/kaggle/input/ubc-ocean-tf-swin-large-training/checkpoints/swin_large_best_0')\nall_image_ids = []\nall_pred_labels = []\nfor test_fname in tqdm(test_filenames):\n    folder_tiles = extract_prune_tiles(test_fname, img_dir, size=2048, new_size=384)\n    #tile_filenames = glob.glob(f'{folder_tiles}/*.png')\n    images = prepare_test_data(folder_tiles)\n    image_id = os.path.splitext(os.path.basename(test_fname))[0]\n    all_image_ids.append(image_id)\n    test_convert_to_tfrec([int(image_id)]*len(images), images, folder_tiles)\n    del images\n    gc.collect()\n    tfrec_file = glob.glob(f'{folder_tiles}/*.tfrec')\n    preds = model.predict(get_test_dataset(tfrec_file), batch_size=BATCH_SIZE, verbose=False)\n    preds = np.sum(preds, axis=0)/len(preds)\n    all_pred_labels.append(list(CLASSES.keys())[np.argmax(preds)])\n    del preds\n    gc.collect()\n    remove_folder(folder_tiles)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-08T13:52:04.926598Z","iopub.execute_input":"2023-12-08T13:52:04.926925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.DataFrame({\n    'image_id': all_image_ids,\n    'label': all_pred_labels\n})\nsub_df.to_csv('submission.csv', index=False)\nsub_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf swintransformertf","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}