{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6867765,"sourceType":"datasetVersion","datasetId":3946918},{"sourceId":7361802,"sourceType":"datasetVersion","datasetId":4276321},{"sourceId":7465291,"sourceType":"datasetVersion","datasetId":4276337,"isSourceIdPinned":true}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":149.078434,"end_time":"2023-12-20T23:14:12.469437","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-12-20T23:11:43.391003","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Note: this notebook is made available under the [Apache 2.0 licence](https://www.apache.org/licenses/LICENSE-2.0).","metadata":{}},{"cell_type":"code","source":"# Clear previous runs\n!rm -rf /kaggle/working/build\n!rm -rf /kaggle/working/features\n!rm -rf /kaggle/working/tmp","metadata":{"papermill":{"duration":2.887237,"end_time":"2023-12-20T23:11:49.684376","exception":false,"start_time":"2023-12-20T23:11:46.797139","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:12.269018Z","iopub.execute_input":"2024-06-09T08:24:12.26939Z","iopub.status.idle":"2024-06-09T08:24:15.154291Z","shell.execute_reply.started":"2024-06-09T08:24:12.269362Z","shell.execute_reply":"2024-06-09T08:24:15.152935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 0 - Compile `split_png` code","metadata":{}},{"cell_type":"code","source":"!cd /kaggle/working\n!mkdir build","metadata":{"papermill":{"duration":1.896072,"end_time":"2023-12-20T23:11:51.607547","exception":false,"start_time":"2023-12-20T23:11:49.711475","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:15.156481Z","iopub.execute_input":"2024-06-09T08:24:15.156806Z","iopub.status.idle":"2024-06-09T08:24:17.076496Z","shell.execute_reply.started":"2024-06-09T08:24:15.156777Z","shell.execute_reply":"2024-06-09T08:24:17.075161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SPLIT_PNG_C_CODE = \"\"\"\n#include <png.h>\n#include <pthread.h>\n#include <semaphore.h>\n#include <stdio.h>\n#include <stdlib.h>\n\nstruct dst_image {\n  FILE *fp;\n  png_structp png;\n  png_infop info;\n};\n\nstruct dst_image dst_image_create(const char *filename, png_uint_32 width,\n                                  png_uint_32 height) {\n  struct dst_image dst;\n  dst.fp = fopen(filename, \"wb\");\n  dst.png = png_create_write_struct(PNG_LIBPNG_VER_STRING, NULL, NULL, NULL);\n  dst.info = png_create_info_struct(dst.png);\n\n  png_init_io(dst.png, dst.fp);\n  png_set_IHDR(dst.png, dst.info, width, height, 8, PNG_COLOR_TYPE_RGB,\n               PNG_INTERLACE_NONE, PNG_COMPRESSION_TYPE_DEFAULT,\n               PNG_FILTER_TYPE_DEFAULT);\n  png_write_info(dst.png, dst.info);\n  return dst;\n}\n\nvoid dst_image_destroy(struct dst_image *dst) {\n  png_write_end(dst->png, NULL);\n  fclose(dst->fp);\n  png_destroy_write_struct(&dst->png, &dst->info);\n}\n\nstruct src_image {\n  FILE *fp;\n  png_structp png;\n  png_infop info;\n  int alpha_present;\n  int interlaced;\n  int height;\n  int width;\n};\n\nstruct src_image src_image_open(const char *file_path) {\n  struct src_image ret = {0};\n  ret.fp = fopen(file_path, \"rb\");\n  ret.png = png_create_read_struct(PNG_LIBPNG_VER_STRING, NULL, NULL, NULL);\n  if (!ret.png)\n    abort();\n\n  ret.info = png_create_info_struct(ret.png);\n  if (!ret.info)\n    abort();\n\n  png_init_io(ret.png, ret.fp);\n  png_read_info(ret.png, ret.info);\n\n  png_byte color_type = png_get_color_type(ret.png, ret.info);\n  png_byte bit_depth = png_get_bit_depth(ret.png, ret.info);\n\n  if (bit_depth == 16)\n    png_set_strip_16(ret.png);\n\n  if (color_type == PNG_COLOR_TYPE_PALETTE)\n    png_set_palette_to_rgb(ret.png);\n\n  if (color_type == PNG_COLOR_TYPE_GRAY && bit_depth < 8)\n    png_set_expand_gray_1_2_4_to_8(ret.png);\n\n  if (png_get_valid(ret.png, ret.info, PNG_INFO_tRNS))\n    png_set_tRNS_to_alpha(ret.png);\n\n  ret.alpha_present = 1;\n  if (color_type == PNG_COLOR_TYPE_RGB || color_type == PNG_COLOR_TYPE_GRAY ||\n      color_type == PNG_COLOR_TYPE_PALETTE)\n    ret.alpha_present = 0;\n\n  if (color_type == PNG_COLOR_TYPE_GRAY ||\n      color_type == PNG_COLOR_TYPE_GRAY_ALPHA)\n    png_set_gray_to_rgb(ret.png);\n\n  png_byte interlace_type = png_get_interlace_type(ret.png, ret.info);\n  if (interlace_type == PNG_INTERLACE_ADAM7) {\n    ret.interlaced = 1;\n  } else if (interlace_type == PNG_INTERLACE_NONE) {\n    ret.interlaced = 0;\n  } else {\n    abort();\n  }\n\n  png_read_update_info(ret.png, ret.info);\n\n  ret.height = png_get_image_height(ret.png, ret.info);\n  ret.width = png_get_image_width(ret.png, ret.info);\n\n  return ret;\n}\n\nvoid src_image_close(struct src_image *src) {\n  png_destroy_read_struct(&src->png, &src->info, NULL);\n  fclose(src->fp);\n}\n\nstruct coord {\n  int x;\n  int y;\n  struct dst_image *dst;\n};\n\nstruct coords {\n  struct coord *coords;\n  size_t n;\n};\n\nint compare_coord_by_y(const void *a, const void *b) {\n  struct coord *ca = (struct coord *)a;\n  struct coord *cb = (struct coord *)b;\n  return ca->y - cb->y;\n}\n\nstruct coords coords_from_file(const char *coordinates_file, int max_patches,\n                               int patch_size, int width, int height) {\n  struct coords c;\n  c.n = 0;\n\n  FILE *coord_fp = fopen(coordinates_file, \"r\");\n  if (!coord_fp)\n    abort();\n\n  c.coords = (struct coord *)malloc(sizeof(struct coord) * max_patches);\n\n  int x, y;\n  while (c.n < max_patches && fscanf(coord_fp, \"%d,%d\", &x, &y) == 2) {\n    if (x + patch_size > width || y + patch_size >= height) {\n      continue; // Patch would be out of bounds\n    }\n    c.coords[c.n].x = x;\n    c.coords[c.n].y = y;\n    c.coords[c.n].dst = NULL;\n    c.n++;\n  }\n  qsort(c.coords, c.n, sizeof(struct coord), compare_coord_by_y);\n  fclose(coord_fp);\n  return c;\n};\n\nvoid coords_free(struct coords *c) {\n  free(c->coords);\n  c->coords = NULL;\n  c->n = 0;\n}\n\nvoid row_rgba_to_rgb(png_bytep row, size_t width) {\n  for (int x = 0; x < width; x++) {\n    row[x * 3 + 0] = row[x * 4 + 0] * row[x * 4 + 3] / 255;\n    row[x * 3 + 1] = row[x * 4 + 1] * row[x * 4 + 3] / 255;\n    row[x * 3 + 2] = row[x * 4 + 2] * row[x * 4 + 3] / 255;\n  }\n}\n\nvoid process_row_by_row(struct src_image *src, const struct coords *c,\n                        size_t patch_size, const char *output_dir) {\n  png_bytep row_pointer =\n      (png_byte *)malloc(png_get_rowbytes(src->png, src->info));\n  const int alpha_present = src->alpha_present;\n\n  // Open patches for output\n  for (size_t i = 0; i < c->n; i++) {\n    char filename[100];\n    snprintf(filename, sizeof(filename), \"%s/patch_%d.png\", output_dir, (int)i);\n    c->coords[i].dst = (struct dst_image *)malloc(sizeof(struct dst_image));\n    *c->coords[i].dst = dst_image_create(filename, patch_size, patch_size);\n  }\n\n  size_t output_patch_begin = 0;\n  size_t output_patch_end = 0;\n\n  for (int y = 0; y < src->height; y++) {\n    // Stop considering patches that are finished.\n    while (output_patch_begin < output_patch_end &&\n           c->coords[output_patch_begin].y + patch_size <= y) {\n      output_patch_begin++;\n    }\n    // Consider new patches.\n    while (output_patch_end < c->n && c->coords[output_patch_end].y <= y) {\n      output_patch_end++;\n    }\n\n    if (output_patch_begin == output_patch_end) {\n      // No patches are using this row.\n      png_read_row(src->png, NULL, NULL);\n      continue;\n    }\n\n    png_read_row(src->png, row_pointer, NULL);\n    if (alpha_present) {\n      row_rgba_to_rgb(row_pointer, src->width);\n    }\n    for (size_t i = output_patch_begin; i < output_patch_end; i++) {\n      png_write_row(c->coords[i].dst->png, row_pointer + (c->coords[i].x * 3));\n    }\n  }\n\n  // Close all the patches.\n  for (size_t i = 0; i < c->n; i++) {\n    dst_image_destroy(c->coords[i].dst);\n    free(c->coords[i].dst);\n    c->coords[i].dst = NULL;\n  }\n  free(row_pointer);\n}\n\nstruct thread_params {\n  sem_t *concurrency_sem;\n  sem_t *done_sem;\n  const char *output_dir;\n  png_bytep *row_pointers;\n  int patch_id;\n  int patch_size;\n  int x;\n  int y;\n};\n\nvoid *thread_save_patch(void *thread_arg) {\n  struct thread_params *p = (struct thread_params *)thread_arg;\n  char filename[100];\n  snprintf(filename, sizeof(filename), \"%s/patch_%d.png\", p->output_dir,\n           p->patch_id);\n  struct dst_image dst = dst_image_create(filename, /*width=*/p->patch_size,\n                                          /*height=*/p->patch_size);\n\n  png_bytep *row_pointers_out =\n      (png_bytep *)malloc(sizeof(png_bytep *) * p->patch_size);\n  for (int i = 0; i < p->patch_size; i++) {\n    row_pointers_out[i] = p->row_pointers[p->y + i] + p->x * 3;\n  }\n  png_write_image(dst.png, row_pointers_out);\n\n  free(row_pointers_out);\n\n  dst_image_destroy(&dst);\n\n  sem_post(p->concurrency_sem);\n  sem_post(p->done_sem);\n  free(p);\n  return NULL;\n}\n\nvoid process_whole_multithreaded(struct src_image *src, const struct coords *c,\n                                 size_t patch_size, const char *output_dir) {\n  png_bytep *row_pointers =\n      (png_bytep *)malloc(sizeof(png_bytep) * src->height);\n\n  size_t output_patch_begin = 0;\n  size_t output_patch_end = 0;\n\n  for (int y = 0; y < src->height; y++) {\n    // Stop considering patches that are finished.\n    while (output_patch_begin < output_patch_end &&\n           c->coords[output_patch_begin].y + patch_size <= y) {\n      output_patch_begin++;\n    }\n    // Consider new patches.\n    while (output_patch_end < c->n && c->coords[output_patch_end].y <= y) {\n      output_patch_end++;\n    }\n\n    if (output_patch_begin == output_patch_end) {\n      // No patches are using this row.\n      row_pointers[y] = NULL;\n    } else {\n      row_pointers[y] = (png_byte *)malloc(png_get_rowbytes(src->png, src->info));\n    }\n  }\n\n  png_read_image(src->png, row_pointers);\n\n  if (src->alpha_present) {\n    // Convert RGBA to RGB\n    for (int y = 0; y < src->height; y++) {\n      png_bytep row = row_pointers[y];\n      if (row) {\n        row_rgba_to_rgb(row, src->width);\n      }\n    }\n  }\n\n  int patch_id = 0;\n  sem_t concurrency_sem;\n  sem_t done_sem;\n  int launched_threads = 0;\n  // Allow 8 threads simultaneously.\n  sem_init(&concurrency_sem, 0, 8);\n  sem_init(&done_sem, 0, 0);\n  for (size_t i = 0; i < c->n; i++) {\n    struct thread_params *p =\n        (struct thread_params *)malloc(sizeof(struct thread_params));\n    p->x = c->coords[i].x;\n    p->y = c->coords[i].y;\n    p->concurrency_sem = &concurrency_sem;\n    p->done_sem = &done_sem;\n    p->patch_id = patch_id;\n    p->output_dir = output_dir;\n    p->patch_size = patch_size;\n    p->row_pointers = row_pointers;\n    sem_wait(&concurrency_sem);\n    launched_threads++;\n    pthread_t t;\n    pthread_create(&t, NULL, thread_save_patch, p);\n    pthread_detach(t);\n\n    patch_id++;\n  }\n  // Wait for all threads to finish\n  for (int i = 0; i < launched_threads; i++) {\n    sem_wait(&done_sem);\n  }\n  sem_destroy(&concurrency_sem);\n  sem_destroy(&done_sem);\n\n  for (int y = 0; y < src->height; y++) {\n    free(row_pointers[y]);\n  }\n  free(row_pointers);\n}\n\nvoid process_file(const char *file_path, const char *output_dir,\n                  const char *coordinates_file, int max_patches,\n                  int patch_size) {\n  struct src_image src = src_image_open(file_path);\n\n  struct coords c = coords_from_file(coordinates_file, max_patches, patch_size,\n                                     src.width, src.height);\n\n  if (src.interlaced == 0) {\n    process_row_by_row(&src, &c, patch_size, output_dir);\n  } else {\n    process_whole_multithreaded(&src, &c, patch_size, output_dir);\n  }\n\n  coords_free(&c);\n  src_image_close(&src);\n}\n\nint main(int argc, char *argv[]) {\n  if (argc != 6) {\n    printf(\"Usage: %s <input_png> <output_dir> <coordinates_file> \"\n           \"<max_patches> <patch_size>\",\n           argv[0]);\n    return 1;\n  }\n\n  int patch_size = atoi(argv[5]);\n  int max_patches = atoi(argv[4]);\n  process_file(argv[1], argv[2], argv[3], max_patches, patch_size);\n\n  return 0;\n}\n\"\"\"","metadata":{"papermill":{"duration":0.042613,"end_time":"2023-12-20T23:11:51.676582","exception":false,"start_time":"2023-12-20T23:11:51.633969","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:17.078476Z","iopub.execute_input":"2024-06-09T08:24:17.078808Z","iopub.status.idle":"2024-06-09T08:24:17.09363Z","shell.execute_reply.started":"2024-06-09T08:24:17.078779Z","shell.execute_reply":"2024-06-09T08:24:17.092778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Write C code to file\nwith open('/kaggle/working/split_png.c', 'w') as f:\n    f.write(SPLIT_PNG_C_CODE)","metadata":{"papermill":{"duration":0.032246,"end_time":"2023-12-20T23:11:51.733746","exception":false,"start_time":"2023-12-20T23:11:51.7015","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:17.096541Z","iopub.execute_input":"2024-06-09T08:24:17.097196Z","iopub.status.idle":"2024-06-09T08:24:17.107864Z","shell.execute_reply.started":"2024-06-09T08:24:17.097163Z","shell.execute_reply":"2024-06-09T08:24:17.10712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compile C code\n!gcc -O2 -o /kaggle/working/build/split_png split_png.c -lpng -lpthread","metadata":{"papermill":{"duration":1.391843,"end_time":"2023-12-20T23:11:53.195414","exception":false,"start_time":"2023-12-20T23:11:51.803571","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:17.1089Z","iopub.execute_input":"2024-06-09T08:24:17.109253Z","iopub.status.idle":"2024-06-09T08:24:18.244188Z","shell.execute_reply.started":"2024-06-09T08:24:17.109227Z","shell.execute_reply":"2024-06-09T08:24:18.242921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1 - Imports","metadata":{"papermill":{"duration":0.024013,"end_time":"2023-12-20T23:11:53.244538","exception":false,"start_time":"2023-12-20T23:11:53.220525","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nos.environ['CUDA_VISIBLE_DEVICES'] = '0'\nos.environ['RAY_DEDUP_LOGS'] = '0'","metadata":{"papermill":{"duration":0.033772,"end_time":"2023-12-20T23:11:53.309152","exception":false,"start_time":"2023-12-20T23:11:53.27538","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:18.245732Z","iopub.execute_input":"2024-06-09T08:24:18.246065Z","iopub.status.idle":"2024-06-09T08:24:18.251191Z","shell.execute_reply.started":"2024-06-09T08:24:18.246037Z","shell.execute_reply":"2024-06-09T08:24:18.25026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from functools import partial\n\nimport re\nimport gc\nimport abc\nimport itertools\nimport math\nimport pickle\nimport warnings\nfrom warnings import warn\nfrom typing import List, Tuple, Callable, Optional, Union, Dict, Literal, Any\nfrom math import ceil\nfrom pathlib import Path\nfrom PIL import Image\nMAX_IMAGE_PIXELS = 10_000_000_000 \nImage.MAX_IMAGE_PIXELS = MAX_IMAGE_PIXELS\nfrom PIL.Image import Image as ImageCls\nfrom subprocess import run, PIPE\nfrom tempfile import TemporaryDirectory\nfrom traceback import print_exc\nfrom copy import deepcopy\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom matplotlib.colors import rgb_to_hsv\nfrom skimage.filters import threshold_otsu\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.dataloader import default_collate\nfrom torchvision import transforms\nfrom torchvision.models import resnet18, ResNet18_Weights\nfrom torchvision.models.feature_extraction import create_feature_extractor\nimport joblib\nfrom scipy import stats\nfrom scipy.special import softmax","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":false,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":5.958187,"end_time":"2023-12-20T23:11:59.29224","exception":false,"start_time":"2023-12-20T23:11:53.334053","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:18.252543Z","iopub.execute_input":"2024-06-09T08:24:18.252979Z","iopub.status.idle":"2024-06-09T08:24:18.266556Z","shell.execute_reply.started":"2024-06-09T08:24:18.252925Z","shell.execute_reply":"2024-06-09T08:24:18.265675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n# Install ray==2.7.1 frow wheels\n!pip install --force-reinstall ray==2.7.1 -f /kaggle/input/ray-wheels --no-index","metadata":{"papermill":{"duration":39.186566,"end_time":"2023-12-20T23:12:38.504256","exception":false,"start_time":"2023-12-20T23:11:59.31769","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:18.26782Z","iopub.execute_input":"2024-06-09T08:24:18.268233Z","iopub.status.idle":"2024-06-09T08:24:49.743113Z","shell.execute_reply.started":"2024-06-09T08:24:18.268194Z","shell.execute_reply":"2024-06-09T08:24:49.742048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ray\nprint(ray.__version__)","metadata":{"papermill":{"duration":0.34307,"end_time":"2023-12-20T23:12:38.872981","exception":false,"start_time":"2023-12-20T23:12:38.529911","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.744776Z","iopub.execute_input":"2024-06-09T08:24:49.745651Z","iopub.status.idle":"2024-06-09T08:24:49.750772Z","shell.execute_reply.started":"2024-06-09T08:24:49.745608Z","shell.execute_reply":"2024-06-09T08:24:49.749921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2 - Environment variables","metadata":{"papermill":{"duration":0.025116,"end_time":"2023-12-20T23:12:38.923607","exception":false,"start_time":"2023-12-20T23:12:38.898491","status":"completed"},"tags":[]}},{"cell_type":"code","source":"PATH_TO_EXECUTABLE = '/kaggle/working/build/split_png'\nassert Path(PATH_TO_EXECUTABLE).is_file()","metadata":{"papermill":{"duration":0.032645,"end_time":"2023-12-20T23:12:39.040156","exception":false,"start_time":"2023-12-20T23:12:39.007511","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.754718Z","iopub.execute_input":"2024-06-09T08:24:49.755014Z","iopub.status.idle":"2024-06-09T08:24:49.771088Z","shell.execute_reply.started":"2024-06-09T08:24:49.754982Z","shell.execute_reply":"2024-06-09T08:24:49.770254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set to ``False`` for inference/submissions\nIS_TRAINING = False","metadata":{"papermill":{"duration":0.03277,"end_time":"2023-12-20T23:12:39.098425","exception":false,"start_time":"2023-12-20T23:12:39.065655","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.772073Z","iopub.execute_input":"2024-06-09T08:24:49.772352Z","iopub.status.idle":"2024-06-09T08:24:49.782029Z","shell.execute_reply.started":"2024-06-09T08:24:49.772328Z","shell.execute_reply":"2024-06-09T08:24:49.781191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTLIERS_CLASS = 5\nENTROPY_THRESHOLD = 1.25","metadata":{"papermill":{"duration":0.032694,"end_time":"2023-12-20T23:12:39.156634","exception":false,"start_time":"2023-12-20T23:12:39.12394","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.783039Z","iopub.execute_input":"2024-06-09T08:24:49.783343Z","iopub.status.idle":"2024-06-09T08:24:49.791375Z","shell.execute_reply.started":"2024-06-09T08:24:49.783319Z","shell.execute_reply":"2024-06-09T08:24:49.790497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CHOWDER_KWARGS = {\n    'in_features': 768,\n    'out_features': 5,\n    'n_top': 10,\n    'n_bottom': 10,\n    'tiles_mlp_hidden': [192],\n    'mlp_hidden': [96],\n    'mlp_dropout': 0.3,\n    'mlp_activation': nn.Sigmoid(),\n}","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.792588Z","iopub.execute_input":"2024-06-09T08:24:49.793195Z","iopub.status.idle":"2024-06-09T08:24:49.801111Z","shell.execute_reply.started":"2024-06-09T08:24:49.793164Z","shell.execute_reply":"2024-06-09T08:24:49.800228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_CLS_MAPPING = {\n    'CC': 0,\n    'EC': 1, \n    'HGSC': 2, \n    'LGSC': 3, \n    'MC': 4,\n    'Other': 5,\n}\n\n_CLS_INV_MAPPING = {\n    '0': 'CC', \n    '1': 'EC', \n    '2': 'HGSC', \n    '3': 'LGSC', \n    '4': 'MC',\n    '5': 'Other',\n    }","metadata":{"execution":{"iopub.status.busy":"2024-06-09T08:24:49.80213Z","iopub.execute_input":"2024-06-09T08:24:49.802383Z","iopub.status.idle":"2024-06-09T08:24:49.810997Z","shell.execute_reply.started":"2024-06-09T08:24:49.802361Z","shell.execute_reply":"2024-06-09T08:24:49.810239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/UBC-OCEAN/\"\n_CLASSES = ['CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other']\nNUM_CLASSES = len(_CLASSES)\nWEIGHTS_DIR = \"/kaggle/input/models-weights-final\"\n_DATA_DIR = Path(DATA_DIR)\nOUTPUT_DIR = '/kaggle/working/features/'\nWEIGHTS_DICT = {\n    'ibotvitbasepancancer': None,\n    'resnet18': Path(WEIGHTS_DIR).joinpath('resnet18_imagenet1k_v1.pt'),\n    'tma_detector': Path(WEIGHTS_DIR).joinpath('lr_tma_detector_thumbnail.pkl'),\n    }","metadata":{"papermill":{"duration":0.034713,"end_time":"2023-12-20T23:12:39.216798","exception":false,"start_time":"2023-12-20T23:12:39.182085","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.812018Z","iopub.execute_input":"2024-06-09T08:24:49.812333Z","iopub.status.idle":"2024-06-09T08:24:49.820927Z","shell.execute_reply.started":"2024-06-09T08:24:49.812303Z","shell.execute_reply":"2024-06-09T08:24:49.820247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Note: ⚠️ The weights of Phikon (`ibotvitbasepancancer`) are available on [HuggingFace](https://huggingface.co/owkin/phikon) under Owkin's non-commercial [licence](https://github.com/owkin/HistoSSLscaling/blob/main/LICENSE.txt).","metadata":{}},{"cell_type":"markdown","source":"# 3 - Core (functions & classes)","metadata":{"papermill":{"duration":0.0252,"end_time":"2023-12-20T23:12:39.267604","exception":false,"start_time":"2023-12-20T23:12:39.242404","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## 3.1 - Data","metadata":{"papermill":{"duration":0.024874,"end_time":"2023-12-20T23:12:39.317752","exception":false,"start_time":"2023-12-20T23:12:39.292878","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_metadata(slide_id: int):\n    \"\"\"Returns the metadata of the slide of id slide_id contained in the train.csv file.\"\"\"\n    if IS_TRAINING:\n        csv_path = _DATA_DIR / \"train.csv\"\n    else:\n        csv_path = _DATA_DIR / \"test.csv\"\n    df_ = pd.read_csv(csv_path)\n    df_ = df_.set_index(\"image_id\")\n    return df_.loc[slide_id]","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.822171Z","iopub.execute_input":"2024-06-09T08:24:49.823058Z","iopub.status.idle":"2024-06-09T08:24:49.834108Z","shell.execute_reply.started":"2024-06-09T08:24:49.82303Z","shell.execute_reply":"2024-06-09T08:24:49.833387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_tma(slide_id: int) -> bool:\n    \"\"\"Determines whether a slide is a TMA.\"\"\"\n    slide_metadata = get_metadata(slide_id)\n    return slide_metadata.is_tma\n","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.835141Z","iopub.execute_input":"2024-06-09T08:24:49.835458Z","iopub.status.idle":"2024-06-09T08:24:49.844531Z","shell.execute_reply.started":"2024-06-09T08:24:49.835428Z","shell.execute_reply":"2024-06-09T08:24:49.843726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_path_fullres_image(slide_id: int) -> str:\n    \"\"\"Path to a full-resolution image.\"\"\"\n    if IS_TRAINING:\n        img_dir = Path(DATA_DIR).joinpath('train_images')\n    else:\n        img_dir = Path(DATA_DIR).joinpath('test_images')\n    img_path = img_dir.joinpath(f'{slide_id}.png')\n    return str(img_path)\n","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.845597Z","iopub.execute_input":"2024-06-09T08:24:49.845876Z","iopub.status.idle":"2024-06-09T08:24:49.855044Z","shell.execute_reply.started":"2024-06-09T08:24:49.845853Z","shell.execute_reply":"2024-06-09T08:24:49.854182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_image_dimensions(slide_id: int) -> Tuple[int, int]:\n    \"\"\"Reads width and height from a full-resolution image.\"\"\"\n    try:\n        Image.MAX_IMAGE_PIXELS = 10_000_000_000\n        fullres_fpath = get_path_fullres_image(slide_id)\n        img = Image.open(fullres_fpath)\n        fullres_w, fullres_h = img.size\n        return fullres_w, fullres_h\n    except Exception:  # noqa\n        slide_metadata = get_metadata(slide_id)\n        fullres_w, fullres_h = slide_metadata.image_width, slide_metadata.image_height\n        return fullres_w, fullres_h","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.856138Z","iopub.execute_input":"2024-06-09T08:24:49.856617Z","iopub.status.idle":"2024-06-09T08:24:49.865593Z","shell.execute_reply.started":"2024-06-09T08:24:49.856591Z","shell.execute_reply":"2024-06-09T08:24:49.864721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_path_thumbnail_image(slide_id: int) -> str:\n    \"\"\"Path to a thumbnail image (if it exists).\"\"\"\n    if IS_TRAINING and is_tma(slide_id):\n        img_dir = Path(DATA_DIR).joinpath('train_images')\n        img_path = img_dir.joinpath(f'{slide_id}.png')\n        return str(img_path)\n    elif IS_TRAINING and (not is_tma(slide_id)):\n        img_dir = Path(DATA_DIR).joinpath('train_thumbnails')\n        img_path = img_dir.joinpath(f'{slide_id}_thumbnail.png')\n        return str(img_path)\n    elif not IS_TRAINING:\n        # Note: \n        # - if a thumbnail exist, we use it\n        # - otherwise, we use the full resolution image\n        img_dir = Path(DATA_DIR).joinpath('test_thumbnails')\n        expected_img_path = img_dir.joinpath(f'{slide_id}_thumbnail.png')\n        if expected_img_path.is_file():\n            return str(expected_img_path)\n        else:\n            return get_path_fullres_image(slide_id)","metadata":{"papermill":{"duration":0.040987,"end_time":"2023-12-20T23:12:39.383988","exception":false,"start_time":"2023-12-20T23:12:39.343001","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.866864Z","iopub.execute_input":"2024-06-09T08:24:49.867156Z","iopub.status.idle":"2024-06-09T08:24:49.876275Z","shell.execute_reply.started":"2024-06-09T08:24:49.867134Z","shell.execute_reply":"2024-06-09T08:24:49.875386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    \"\"\"Dataset of PNG images.\"\"\"\n    \n    def __init__(self, \n                 root_dir: Union[str, Path], \n                 transform: Optional[Callable] = None) -> None:\n        self.root_dir = root_dir\n        self.transform = transform\n        self.__build()\n    \n    def __build(self) -> None:\n        self._paths = list(Path(self.root_dir).rglob('*.png'))\n    \n    def __len__(self) -> int:\n        return len(self._paths)\n\n    def __getitem__(self, idx: int) -> Union[ImageCls, torch.Tensor]:\n        img = Image.open(self._paths[idx])\n        if self.transform is not None:\n            return self.transform(img)\n        else:\n            return img","metadata":{"papermill":{"duration":0.035165,"end_time":"2023-12-20T23:12:39.444372","exception":false,"start_time":"2023-12-20T23:12:39.409207","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.877385Z","iopub.execute_input":"2024-06-09T08:24:49.877688Z","iopub.status.idle":"2024-06-09T08:24:49.890493Z","shell.execute_reply.started":"2024-06-09T08:24:49.877655Z","shell.execute_reply":"2024-06-09T08:24:49.889656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ThumbnailsDataset(Dataset):\n    \"\"\"Image Dataset for thumbnails.\"\"\"\n    def __init__(self,\n                 img_paths: List[Union[str, Path]],\n                 transform: Optional[Callable] = None):\n        super().__init__()\n        self.img_paths = img_paths\n        self.transform = transform\n        self.__validate()\n\n    def __validate(self) -> None:\n        assert all([Path(_f).is_file() for _f in self.img_paths])\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    @staticmethod\n    def get_slide_id_from_filename(fname: str) -> Optional[int]:\n        match = re.findall(r'^(\\d{1,6})', str(fname))\n        if match:\n            return int(match[0])\n\n    def __getitem__(self, idx):\n        img = Image.open(self.img_paths[idx])\n        # img = img.convert('RGB')\n        \n        slide_id = self.get_slide_id_from_filename(self.img_paths[idx].name)\n        if slide_id is None:\n            raise ValueError(f'Could not infer slide ID from {self.img_paths[idx]}!')\n\n        if self.transform is not None:\n            img = self.transform(img)\n        \n        return img, slide_id","metadata":{"papermill":{"duration":0.036894,"end_time":"2023-12-20T23:12:39.506521","exception":false,"start_time":"2023-12-20T23:12:39.469627","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.89151Z","iopub.execute_input":"2024-06-09T08:24:49.891791Z","iopub.status.idle":"2024-06-09T08:24:49.901815Z","shell.execute_reply.started":"2024-06-09T08:24:49.891757Z","shell.execute_reply":"2024-06-09T08:24:49.900993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pad_collate_fn(batch: List[Tuple[torch.Tensor, Any]],\n                   batch_first: bool = True,\n                   max_len: Optional[int] = None) -> Tuple[torch.Tensor, torch.BoolTensor, Any]:\n    \"\"\"Pad together sequences of arbitrary lengths.\"\"\"\n    # Expect the sequences to be the first one in the sample tuples\n    sequences = []\n    others = []\n    for sample in batch:\n        sequences.append(sample[0])\n        others.append(sample[1:])\n\n    if max_len is None:\n        max_len = max([s.size(0) for s in sequences])\n\n    trailing_dims = sequences[0].size()[1:]\n\n    if batch_first:\n        padded_dims = (len(sequences), max_len) + trailing_dims\n        masks_dims = (len(sequences), max_len, 1)\n    else:\n        padded_dims = (max_len, len(sequences)) + trailing_dims\n        masks_dims = (max_len, len(sequences), 1)\n\n    padded_sequences = sequences[0].data.new(*padded_dims).fill_(0.0)\n    masks = torch.ones(*masks_dims, dtype=torch.bool)\n\n    for i, tensor in enumerate(sequences):\n        length = tensor.size(0)\n        # use index notation to prevent duplicate references to the tensor\n        if batch_first:\n            padded_sequences[i, :length, ...] = tensor[:max_len, ...]\n            masks[i, :length, ...] = False\n        else:\n            padded_sequences[:length, i, ...] = tensor[:max_len, ...]\n            masks[:length, i, ...] = False\n\n    # Batching other members of the tuple using default_collate\n    others = default_collate(others)\n\n    return padded_sequences, masks, *others","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.902872Z","iopub.execute_input":"2024-06-09T08:24:49.903207Z","iopub.status.idle":"2024-06-09T08:24:49.915934Z","shell.execute_reply.started":"2024-06-09T08:24:49.903177Z","shell.execute_reply":"2024-06-09T08:24:49.915155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeaturesDataset(Dataset):\n    \"\"\"Creates a `FeaturesDataset`.\n\n    Parameters\n    ----------\n    features_fpaths : Union[str, Path]\n        Path to the features. Each path should be of the form: ``[...]/SLIDE_ID/features.npy``.\n    labels : List[Union[str]]\n        List of slide labels. Each label should be one of: 'CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other'.\n    max_tiles : Optional[int], default=2_000\n        Maximum number of tiles per slide.\n    shuffle : bool, default=False\n        Whether to shuffle the tiles.\n    mmap_mode : Optional[str] = None\n        If not ``None``, memory map mode to be used with ``np.load``.\n    \"\"\"\n\n    def __init__(self,\n                 features_fpaths: Union[str, Path],\n                 labels: List[str],\n                 max_tiles: Optional[int] = 2_000,\n                 shuffle: bool = False,\n                 mmap_mode: Optional[str] = None):\n        super().__init__()\n\n        self.features_fpaths = features_fpaths\n        self.labels = labels\n        self.max_tiles = max_tiles\n        self.shuffle = shuffle\n        self.mmap_mode = mmap_mode\n\n        self._random_state = None\n        self._torch_generator = None\n        self._numpy_generator = None\n\n        self.__validate_args()\n        self.__build()\n\n    def __validate_args(self) -> None:\n        \"\"\"Validate class args.\"\"\"\n        if len(self.features_fpaths) != len(self.labels):\n            raise ValueError('Expected as many feature files as labels!')\n        elif not all([yy in list(_CLS_MAPPING.keys()) for yy in self.labels]):\n            raise ValueError('Got invalid slide labels!')\n\n    def __build(self) -> None:\n        self._encoded_labels = [_CLS_MAPPING[yy] for yy in self.labels]\n\n    def __len__(self) -> int:\n        return len(self.labels)\n\n    @property\n    def random_state(self) -> Union[int, None]:\n        return self._random_state\n\n    @random_state.setter\n    def random_state(self, seed: Optional[int] = None) -> None:\n        \"\"\"Sets the dataset's `random_state` attribute and generators.\"\"\"\n        self._random_state = seed\n        self._torch_generator = torch.Generator()\n        self._torch_generator.manual_seed(self._random_state)\n        self._numpy_generator = np.random.RandomState(seed=self._random_state)\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, int]:\n        fpath = self.features_fpaths[idx]\n        input_arr = np.load(str(fpath), mmap_mode=self.mmap_mode)\n\n        n_tiles, _ = input_arr.shape\n        if self.shuffle:\n            if (self.max_tiles is None) or (n_tiles < self.max_tiles):\n                idxs = slice(None)\n            else:\n                _idxs = torch.randperm(n_tiles, generator=self._torch_generator).tolist()\n                idxs = _idxs[:self.max_tiles]\n        else:\n            if (self.max_tiles is None) or (n_tiles < self.max_tiles):\n                idxs = slice(None)\n            else:\n                idxs = np.arange(self.max_tiles)\n        features = input_arr[idxs]\n        slide_features = torch.from_numpy(features)\n        slide_id = int(Path(fpath).parent.name)  # fpath is supposed to be of the form: [...]/SLIDE_ID/features.npy\n\n        return slide_features, self._encoded_labels[idx], slide_id","metadata":{"papermill":{"duration":0.051691,"end_time":"2023-12-20T23:12:39.583349","exception":false,"start_time":"2023-12-20T23:12:39.531658","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.917136Z","iopub.execute_input":"2024-06-09T08:24:49.917394Z","iopub.status.idle":"2024-06-09T08:24:49.933202Z","shell.execute_reply.started":"2024-06-09T08:24:49.917373Z","shell.execute_reply":"2024-06-09T08:24:49.932358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.2 - Matter Detection","metadata":{"papermill":{"duration":0.025029,"end_time":"2023-12-20T23:12:39.633915","exception":false,"start_time":"2023-12-20T23:12:39.608886","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def _compute_matter_mask_from_array(arr: np.ndarray) -> np.ndarray:\n    # Convert to HSV\n    _hsv = rgb_to_hsv(arr)\n\n    # Apply matter detection (using Otsu thresholding)\n    threshold_h = threshold_otsu(_hsv[:, :, 0])\n    threshold_s = threshold_otsu(_hsv[:, :, 1])\n    _mask = np.logical_and(_hsv[:, :, 0] > threshold_h, _hsv[:, :, 1] > threshold_s)\n    _kernel = np.ones((1, 1))\n    mask = cv2.dilate(_mask.astype(np.int32), _kernel, iterations=1)\n\n    return mask","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.934275Z","iopub.execute_input":"2024-06-09T08:24:49.934875Z","iopub.status.idle":"2024-06-09T08:24:49.947675Z","shell.execute_reply.started":"2024-06-09T08:24:49.934842Z","shell.execute_reply":"2024-06-09T08:24:49.946937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_matter_mask(slide_id: int):\n    \"\"\"Applies matter detection to an image (thumbnail).\"\"\"\n    # Load thumbnail\n    path_thumbnail = get_path_thumbnail_image(slide_id)\n    thumbnail_img = Image.open(path_thumbnail)\n    if thumbnail_img.im is None:\n        thumbnail_img.load()\n    thumbnail_arr = np.asarray(thumbnail_img)\n        \n    # Compute matter mask\n    mask = _compute_matter_mask_from_array(thumbnail_arr)\n    return mask","metadata":{"papermill":{"duration":0.047352,"end_time":"2023-12-20T23:12:39.706487","exception":false,"start_time":"2023-12-20T23:12:39.659135","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.949069Z","iopub.execute_input":"2024-06-09T08:24:49.949379Z","iopub.status.idle":"2024-06-09T08:24:49.957996Z","shell.execute_reply.started":"2024-06-09T08:24:49.949355Z","shell.execute_reply":"2024-06-09T08:24:49.957165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.3 - Extractor (Phikon)","metadata":{"papermill":{"duration":0.025291,"end_time":"2023-12-20T23:12:39.756885","exception":false,"start_time":"2023-12-20T23:12:39.731594","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def validate_weights_path(weights_path: Union[str, Path]) -> None:\n    \"\"\"Ensures that model weights exist.\"\"\"\n    if not Path(weights_path).is_file():\n        raise FileNotFoundError(f'Could not load model weights ({weights_path})!')","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.959045Z","iopub.execute_input":"2024-06-09T08:24:49.960198Z","iopub.status.idle":"2024-06-09T08:24:49.971103Z","shell.execute_reply.started":"2024-06-09T08:24:49.960171Z","shell.execute_reply":"2024-06-09T08:24:49.970277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_extractor(module: nn.Module, device: Optional[Union[str, torch.device]] = None) -> nn.Module:\n    \"\"\"Prepares a feature extractor.\"\"\"\n    # Force type ``torch.device``\n    if device is None:\n        _device = torch.device('cpu')\n    else:\n        _device = torch.device(device)\n\n    if (_device.type == 'cuda') and (not torch.cuda.is_available()):\n        raise ValueError('Requested CUDA device but no GPU was found!')\n\n    # If device ID is not specified, use all GPUs\n    gpu_count = torch.cuda.device_count()\n    if gpu_count > 1:\n        module = nn.DataParallel(module, device_ids=list(range(gpu_count)))\n\n    module.to(_device, non_blocking=True)\n    module.eval()\n    module.requires_grad_(False)\n\n    return module","metadata":{"papermill":{"duration":0.036449,"end_time":"2023-12-20T23:12:39.818926","exception":false,"start_time":"2023-12-20T23:12:39.782477","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.980134Z","iopub.execute_input":"2024-06-09T08:24:49.980407Z","iopub.status.idle":"2024-06-09T08:24:49.987309Z","shell.execute_reply.started":"2024-06-09T08:24:49.980383Z","shell.execute_reply":"2024-06-09T08:24:49.986379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BaseExtractor(abc.ABC):\n    \"\"\"Base class for feature extraction.\"\"\"\n\n    def __init__(self, device: Optional[Union[str, torch.device]] = None):\n        self.device = device\n\n    @property\n    def transform(self) -> Callable[[ImageCls], torch.Tensor]:\n        \"\"\"Transform applied to images (``PIL.Image``).\"\"\"\n        raise NotImplementedError\n\n    @abc.abstractmethod\n    def __call__(self, images: torch.Tensor) -> torch.Tensor:\n        \"\"\"Extracts features from a batch of images.\n\n        Parameters\n        ----------\n        images : torch.Tensor\n            Input images, with shape ``(B, C, H, W)``.\n\n        Returns\n        -------\n        features : torch.Tensor\n            Output features, with shape ``(B, F)``.\n        \"\"\"\n        raise NotImplementedError","metadata":{"papermill":{"duration":0.034464,"end_time":"2023-12-20T23:12:39.878862","exception":false,"start_time":"2023-12-20T23:12:39.844398","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:49.988589Z","iopub.execute_input":"2024-06-09T08:24:49.989509Z","iopub.status.idle":"2024-06-09T08:24:49.998746Z","shell.execute_reply.started":"2024-06-09T08:24:49.989477Z","shell.execute_reply":"2024-06-09T08:24:49.997846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _no_grad_trunc_normal_(tensor: torch.Tensor, mean: float, std: float, a: float, b: float) -> torch.Tensor:\n    \"\"\"Truncated normal initialization.\"\"\"\n\n    def _norm_cdf(x: float) -> float:\n        # Computes standard normal cumulative distribution function\n        return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0\n\n    if (mean < a - 2 * std) or (mean > b + 2 * std):\n        warnings.warn(\n            \"mean is more than 2 std from [a, b] in nn.init.trunc_normal_. \"\n            \"The distribution of values may be incorrect.\",\n            stacklevel=2,\n        )\n\n    with torch.no_grad():\n        # Values are generated by using a truncated uniform distribution and\n        # then using the inverse CDF for the normal distribution.\n        # Get upper and lower cdf values.\n        lb = _norm_cdf((a - mean) / std)\n        ub = _norm_cdf((b - mean) / std)\n\n        # Uniformly fill tensor with values from [l, u], then translate to\n        # [2l-1, 2u-1].\n        tensor.uniform_(2 * lb - 1, 2 * ub - 1)\n\n        # Use inverse cdf transform for normal distribution to get truncated\n        # standard normal.\n        tensor.erfinv_()\n\n        # Transform to proper mean, std.\n        tensor.mul_(std * math.sqrt(2.0))\n        tensor.add_(mean)\n\n        # Clamp to ensure it's in the proper range.\n        tensor.clamp_(min=a, max=b)\n        return tensor","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:49.999773Z","iopub.execute_input":"2024-06-09T08:24:50.000058Z","iopub.status.idle":"2024-06-09T08:24:50.012786Z","shell.execute_reply.started":"2024-06-09T08:24:50.000034Z","shell.execute_reply":"2024-06-09T08:24:50.012099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def trunc_normal_(tensor: torch.Tensor,\n                  mean: float = 0.0,\n                  std: float = 1.0,\n                  a: float = -2.0,\n                  b: float = 2.0) -> torch.Tensor:\n    \"\"\"Truncated normal initialization.\n\n    Parameters\n    ----------\n    tensor : torch.Tensor\n        Input tensor (to be initialized).\n\n    mean : float = 0\n        Mean of the Normal distribution.\n\n    std : float = 1\n        Standard deviation of the Normal distribution.\n\n    a : float = 2\n        The minimum cutoff value.\n\n    b : float = 2\n        The maximum cutoff value.\n\n    Returns\n    -------\n    torch.Tensor\n    \"\"\"\n    return _no_grad_trunc_normal_(tensor, mean, std, a, b)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.013714Z","iopub.execute_input":"2024-06-09T08:24:50.013971Z","iopub.status.idle":"2024-06-09T08:24:50.025728Z","shell.execute_reply.started":"2024-06-09T08:24:50.013936Z","shell.execute_reply":"2024-06-09T08:24:50.024995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def drop_path(x: torch.Tensor,\n              drop_prob: float = 0.0,\n              training: bool = False) -> torch.Tensor:\n    \"\"\"Drop paths (Stochastic Depth) per sample.\n\n    Parameters\n    ----------\n    x : torch.Tensor\n        The input tensor.\n    drop_prob : float = 0.0\n        The probability of dropping a path.\n    training : bool = False\n        Whether the model is in training mode or not.\n\n    Returns\n    -------\n    torch.Tensor\n        The output array after applying drop path.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/drop.py#L137\n    \"\"\"\n    if drop_prob == 0.0 or not training:\n        return x\n    keep_prob = 1 - drop_prob\n    shape = (x.shape[0],) + (1,) * (x.ndim - 1)  # work with diff dim tensors, not just 2D ConvNets\n    random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device)\n    random_tensor.floor_()  # binarize\n    output = x.div(keep_prob) * random_tensor\n    return output","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.026724Z","iopub.execute_input":"2024-06-09T08:24:50.027065Z","iopub.status.idle":"2024-06-09T08:24:50.039731Z","shell.execute_reply.started":"2024-06-09T08:24:50.027035Z","shell.execute_reply":"2024-06-09T08:24:50.038865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DropPath(nn.Module):\n    \"\"\"``torch.nn`` module for Drop path implementation. See [1]_ for details.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/drop.py#L157\n    \"\"\"\n\n    def __init__(self, drop_prob: Optional[float] = None) -> None:\n        \"\"\"Initialize the DropPath class.\n\n        Parameters\n        ----------\n        drop_prob : float\n            The probability of dropping each element, by default None.\n        \"\"\"\n        super(DropPath, self).__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Apply drop path to an input tensor.\n\n        Parameters\n        ----------\n        x : torch.Tensor\n            The input tensor.\n\n        Returns\n        -------\n        torch.Tensor\n            Input tensor with dropout applied on a given path.\n        \"\"\"\n        return drop_path(x, self.drop_prob, self.training)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.040813Z","iopub.execute_input":"2024-06-09T08:24:50.041238Z","iopub.status.idle":"2024-06-09T08:24:50.05073Z","shell.execute_reply.started":"2024-06-09T08:24:50.041213Z","shell.execute_reply":"2024-06-09T08:24:50.049986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Mlp(nn.Module):\n    \"\"\"MLP as used in Vision Transformer. See [1]_ for details.\n\n    Parameters\n    ----------\n    in_features : int\n        Input feature size.\n    hidden_features : int = None\n        Hidden feature size, by default None (uses input feature size).\n    out_features : int = None\n        Output feature size, by default None (uses input feature size).\n    act_layer : nn.Module = nn.GELU\n        Activation layer, by default nn.GELU.\n    drop : float = 0.0\n        Dropout rate, by default 0.0.\n\n    Attributes\n    ----------\n    fc1 : nn.Linear\n        First fully connected layer.\n    act : nn.Module\n        Activation layer.\n    fc2 : nn.Linear\n        Second fully connected layer.\n    drop : nn.Dropout\n        Dropout layer.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/mlp.py#L13\n\n    \"\"\"\n\n    def __init__(self,\n                 in_features: int,\n                 hidden_features: Optional[int] = None,\n                 out_features: Optional[int] = None,\n                 act_layer: nn.Module = nn.GELU,\n                 drop: float = 0.0) -> None:\n        super().__init__()\n\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Forward pass of the MLP.\n\n        Parameters\n        ----------\n        x : torch.Tensor\n            Input tensor.\n\n        Returns\n        -------\n        torch.Tensor\n            Output tensor.\n        \"\"\"\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.051828Z","iopub.execute_input":"2024-06-09T08:24:50.052091Z","iopub.status.idle":"2024-06-09T08:24:50.065139Z","shell.execute_reply.started":"2024-06-09T08:24:50.052069Z","shell.execute_reply":"2024-06-09T08:24:50.06425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Attention(nn.Module):\n    \"\"\"Attention module for Vision Transformer as implemented in [1]_.\n\n    Parameters\n    ----------\n    dim : int\n        Input dimension.\n    num_heads : int = 8\n        Number of attention heads, by default 8.\n    qkv_bias : bool = False\n        Whether to include biases in the linear transformations of q, k, and v,\n        by default False.\n    qk_scale : float = None\n        Scale factor for qk dot product, by default None.\n    attn_drop : float = 0.0\n        Dropout rate for attention weights, by default 0.0.\n    proj_drop : float = 0.0\n        Dropout rate for output tensor, by default 0.0.\n\n    Attributes\n    ----------\n    num_heads : int\n        Number of attention heads.\n    scale : float\n        Scale factor for qk dot product.\n    qkv : nn.Linear\n        Linear transformation for q, k, and v.\n    attn_drop : nn.Dropout\n        Dropout layer for attention weights.\n    proj : nn.Linear\n        Linear transformation for output tensor.\n    proj_drop : nn.Dropout\n        Dropout layer for output tensor.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py\n    \"\"\"\n\n    def __init__(self,\n                 dim: int,\n                 num_heads: int = 8,\n                 qkv_bias: bool = False,\n                 qk_scale: Optional[float] = None,\n                 attn_drop: float = 0.0,\n                 proj_drop: float = 0.0) -> None:\n        super().__init__()\n\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = qk_scale or head_dim**-0.5\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"Forward pass of the attention module.\n\n        Parameters\n        ----------\n        x : torch.Tensor\n            Input tensor.\n\n        Returns\n        -------\n        Tuple[torch.Tensor, torch.Tensor]\n            Output tensor.\n        \"\"\"\n        B, N, C = x.shape\n        qkv = (\n            self.qkv(x)\n            .reshape(B, N, 3, self.num_heads, C // self.num_heads)\n            .permute(2, 0, 3, 1, 4)\n        )\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x, attn","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.066379Z","iopub.execute_input":"2024-06-09T08:24:50.066686Z","iopub.status.idle":"2024-06-09T08:24:50.08038Z","shell.execute_reply.started":"2024-06-09T08:24:50.066662Z","shell.execute_reply":"2024-06-09T08:24:50.079586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Block(nn.Module):\n    \"\"\"Vision Transformer Block as implemented in [1]_.\n\n    Parameters\n    ----------\n    dim : int\n        The dimension of the input.\n    num_heads : int\n        The number of attention heads.\n    mlp_ratio : float = 4.0\n        The ratio of hidden dimension to input dimension in the MLP, by default 4.0.\n    qkv_bias : bool = False\n        Whether to include biases in the query, key, and value projections, by default False.\n    qk_scale : float = None\n        The scaling factor for query-key dot product, by default None.\n    drop : float = 0.0\n        The dropout probability, by default 0.0.\n    attn_drop : float = 0.0\n        The dropout probability for attention weights, by default 0.0.\n    drop_path_rate : float = 0.0\n        The dropout probability for the residual connection, by default 0.0.\n    act_layer : torch.nn.Module = torch.nn.GELU\n        The activation layer used in the MLP, by default ``torch.nn.GELU``.\n    norm_layer : torch.nn.Module = torch.nn.LayerNorm\n        The normalization layer, by default ``torch.nn.LayerNorm``.\n    init_values : int\n        The initial values for gamma_1 and gamma_2, by default 0.\n\n    Attributes\n    ----------\n    norm1 : torch.nn.Module\n        The normalization layer before the attention module.\n    attn : Attention\n        The attention module.\n    drop_path : Union[DropPath, nn.Identity]\n        The dropout layer for the residual connection.\n    norm2 : torch.nn.Module\n        The normalization layer after the MLP.\n    mlp : Mlp\n        The MLP module.\n    gamma_1 : Optional[torch.nn.Parameter]\n        The learnable parameter gamma_1.\n    gamma_2 : Optional[torch.nn.Parameter]\n        The learnable parameter gamma_2.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py\n    \"\"\"\n\n    def __init__(self,\n                 dim: int,\n                 num_heads: int,\n                 mlp_ratio: float = 4.0,\n                 qkv_bias: bool = False,\n                 qk_scale: Optional[float] = None,\n                 drop: float = 0.0,\n                 attn_drop: float = 0.0,\n                 drop_path_rate: float = 0.0,\n                 act_layer: nn.Module = nn.GELU,\n                 norm_layer: nn.Module = nn.LayerNorm,\n                 init_values: int = 0) -> None:\n        super().__init__()\n\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention(dim,\n                              num_heads=num_heads,\n                              qkv_bias=qkv_bias,\n                              qk_scale=qk_scale,\n                              attn_drop=attn_drop,\n                              proj_drop=drop)\n        self.drop_path = (DropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity())\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim,\n                       hidden_features=mlp_hidden_dim,\n                       act_layer=act_layer,\n                       drop=drop)\n\n        if init_values > 0:\n            self.gamma_1 = nn.Parameter(init_values * torch.ones(dim), requires_grad=True)\n            self.gamma_2 = nn.Parameter(init_values * torch.ones(dim), requires_grad=True)\n        else:\n            self.gamma_1, self.gamma_2 = None, None\n\n    def forward(self, x: torch.Tensor, return_attention: bool = False) -> torch.Tensor:\n        \"\"\"\n        Forward pass of the Vision Transformer Block module.\n\n        Parameters\n        ----------\n        x : torch.Tensor\n            The input tensor.\n        return_attention : bool\n            Whether to return the attention weights, by default False.\n\n        Returns\n        -------\n        torch.Tensor\n            The output tensor.\n        \"\"\"\n        y, attn = self.attn(self.norm1(x))\n        if return_attention:\n            return attn\n        if self.gamma_1 is None:\n            x = x + self.drop_path(y)\n            x = x + self.drop_path(self.mlp(self.norm2(x)))\n        else:\n            x = x + self.drop_path(self.gamma_1 * y)\n            x = x + self.drop_path(self.gamma_2 * self.mlp(self.norm2(x)))\n        return x","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.081426Z","iopub.execute_input":"2024-06-09T08:24:50.081714Z","iopub.status.idle":"2024-06-09T08:24:50.097177Z","shell.execute_reply.started":"2024-06-09T08:24:50.08169Z","shell.execute_reply":"2024-06-09T08:24:50.096494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PatchEmbed(nn.Module):\n    \"\"\"2D Image to Patch Embedding as implemented in [1]_.\n\n    Parameters\n    ----------\n    img_size:  int = 224\n        Size of the input image (both height and width), by default 224.\n    patch_size: int = 16\n        Size of each patch, by default 16.\n    in_chans: int = 3\n        Number of input channels, by default 3.\n    embed_dim: int = 768\n        Dimension of the embedded patch representation, by default 768.\n\n    Attributes\n    ----------\n    img_size: int\n        Size of the input image (both height and width).\n    patch_size: int\n        Size of each patch.\n    num_patches: int\n        Total number of patches in the image.\n    proj: nn.Conv2d\n        Convolutional layer used for projection.\n\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/layers/patch_embed.py#L25\n    \"\"\"\n\n    def __init__(\n        self,\n        img_size: int = 224,\n        patch_size: int = 16,\n        in_chans: int = 3,\n        embed_dim: int = 768,\n    ) -> None:\n        super().__init__()\n        num_patches = (img_size // patch_size) * (img_size // patch_size)\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = num_patches\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Forward pass of the PatchEmbed module.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor of shape (``B``, ``C``, ``H``, ``W``) where ``B`` is the\n            batch size, ``C`` is the number of channels, ``H`` is the input image\n            height, and ``W`` is the input image width.\n\n        Returns\n        -------\n        torch.Tensor\n            Output tensor of shape (``B``, ``embed_dim``, ``num_patches``) where\n            ``B`` is the batch size, ``embed_dim`` is the dimension of the embedded\n            patch  representation, and ``num_patches`` is the total number of\n            patches in the image.\n        \"\"\"\n        #  B, C, H, W = x.shape\n        return self.proj(x)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.098577Z","iopub.execute_input":"2024-06-09T08:24:50.098887Z","iopub.status.idle":"2024-06-09T08:24:50.112646Z","shell.execute_reply.started":"2024-06-09T08:24:50.098864Z","shell.execute_reply":"2024-06-09T08:24:50.111887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VisionTransformer(nn.Module):\n    \"\"\"Vision Transformer in PyTorch as implemented in [1]_.\n\n    Parameters\n    ----------\n    img_size : List[int] = [224]\n        The size of the input image.\n    patch_size : int = 16\n        The size of each patch in the input image.\n    in_chans : int = 3\n        The number of input channels.\n    num_classes : int = 0\n        The number of output classes.\n    embed_dim : int = 768\n        The dimension of the token embeddings.\n    depth : int = 12\n        The number of transformer blocks.\n    num_heads : int = 12\n        The number of attention heads.\n    mlp_ratio : float = 4.0\n        The ratio of the hidden dimension to the input dimension in the feed-forward network.\n    qkv_bias : bool = False\n        Whether to include biases to the query, key, and value linear layers.\n    qk_scale : Optional[float] = None\n        The scale factor for query and key.\n    drop_rate : float = 0.0\n        The dropout rate.\n    attn_drop_rate : float = 0.0\n        The dropout rate for attention probabilities.\n    drop_path_rate : float = 0.0\n        The dropout rate for residual connections.\n    norm_layer : Callable = partial(torch.nn.LayerNorm, eps=1e-6)\n        The normalization layer to be used.\n    return_all_tokens : bool = False\n        Whether to return all tokens or just the first token.\n    init_values : int = 0\n        The initial values for convolutional layers.\n    use_mean_pooling : bool = False\n        Whether to use mean pooling or not.\n    masked_im_modeling : bool = False\n        Whether to perform masked image modeling or not.\n\n    References\n    ----------\n    .. [1] https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/vision_transformer.py\n    \"\"\"\n\n    def __init__(\n            self,\n            img_size: Tuple[int] = (224,),\n            patch_size: int = 16,\n            in_chans: int = 3,\n            num_classes: int = 0,\n            embed_dim: int = 768,\n            depth: int = 12,\n            num_heads: int = 12,\n            mlp_ratio: float = 4.0,\n            qkv_bias: bool = False,\n            qk_scale: Optional[float] = None,\n            drop_rate: float = 0.0,\n            attn_drop_rate: float = 0.0,\n            drop_path_rate: float = 0.0,\n            norm_layer: Callable = partial(nn.LayerNorm, eps=1e-6),\n            return_all_tokens: bool = False,\n            init_values: int = 0,\n            use_mean_pooling: bool = False,\n            masked_im_modeling: bool = False) -> None:\n        super().__init__()\n\n        self.num_features = self.embed_dim = embed_dim\n        self.return_all_tokens = return_all_tokens\n        self.patch_embed = PatchEmbed(img_size=img_size[0],\n                                      patch_size=patch_size,\n                                      in_chans=in_chans,\n                                      embed_dim=embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        # Stochastic depth decay rule.\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]\n        self.blocks = nn.ModuleList(\n            [\n                Block(dim=embed_dim,\n                      num_heads=num_heads,\n                      mlp_ratio=mlp_ratio,\n                      qkv_bias=qkv_bias,\n                      qk_scale=qk_scale,\n                      drop=drop_rate,\n                      attn_drop=attn_drop_rate,\n                      drop_path_rate=dpr[i],\n                      norm_layer=norm_layer,\n                      init_values=init_values)\n                for i in range(depth)\n            ]\n        )\n\n        self.norm = (nn.Identity() if use_mean_pooling else norm_layer(embed_dim))\n        self.fc_norm = norm_layer(embed_dim) if use_mean_pooling else None\n\n        # Classifier head.\n        self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()\n\n        trunc_normal_(self.pos_embed, std=0.02)\n        trunc_normal_(self.cls_token, std=0.02)\n        self.apply(self._init_weights)\n\n        # Masked image modeling.\n        self.masked_im_modeling = masked_im_modeling\n        if masked_im_modeling:\n            self.masked_embed = nn.Parameter(torch.zeros(1, embed_dim))\n\n    @staticmethod\n    def _init_weights(m: nn.Module) -> None:\n        \"\"\"Initialize the weights of the module.\n\n        Parameters\n        ----------\n        m: torch.nn.Module\n            Module to initialize weights for.\n        \"\"\"\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=0.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    def interpolate_pos_encoding(self, x: torch.Tensor, w: int, h: int) -> torch.Tensor:\n        \"\"\"Interpolate the positional encoding to match the size of the input\n        tokens.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor of shape (batch_size, num_tokens, embed_dim)\n        w: int\n            Width of the input image.\n        h: int\n            Height of the input image.\n\n        Returns\n        -------\n        torch.Tensor\n            Interpolated positional encoding tensor.\n        \"\"\"\n        npatch = x.shape[1] - 1\n        N = self.pos_embed.shape[1] - 1\n        if npatch == N and w == h:\n            return self.pos_embed\n        class_pos_embed = self.pos_embed[:, 0]\n        patch_pos_embed = self.pos_embed[:, 1:]\n        dim = x.shape[-1]\n        w0 = w // self.patch_embed.patch_size\n        h0 = h // self.patch_embed.patch_size\n        # we add a small number to avoid floating point error in the interpolation\n        # see discussion at https://github.com/facebookresearch/dino/issues/8\n        w0, h0 = w0 + 0.1, h0 + 0.1\n        patch_pos_embed = nn.functional.interpolate(\n            patch_pos_embed.reshape(\n                1, int(math.sqrt(N)), int(math.sqrt(N)), dim\n            ).permute(0, 3, 1, 2),\n            scale_factor=(w0 / math.sqrt(N), h0 / math.sqrt(N)),\n            mode=\"bicubic\",\n        )\n        assert (int(w0) == patch_pos_embed.shape[-2] and int(h0) == patch_pos_embed.shape[-1])\n        patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)\n        return torch.cat(\n            (class_pos_embed.unsqueeze(0), patch_pos_embed), dim=1\n        )\n\n    def prepare_tokens(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:\n        \"\"\"Prepare the input tokens for the Vision Transformer.\n\n        Parameters\n        -----------\n            x: torch.Tensor\n                Input tensor of shape (batch_size, num_channels, width, height).\n            mask: Optional[torch.Tensor] = None\n                Mask tensor of shape (batch_size, width, height) for masked\n                image modeling.\n\n        Returns\n        -------\n        torch.Tensor\n            Prepared tokens tensor.\n        \"\"\"\n        B, nc, w, h = x.shape  # pylint: disable=unused-variable\n        # patch linear embedding\n        x = self.patch_embed(x)\n\n        # mask image modeling\n        if mask is not None:\n            x = self.mask_model(x, mask)\n        x = x.flatten(2).transpose(1, 2)\n\n        # add the [CLS] token to the embed patch tokens\n        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n\n        # add positional encoding to each token\n        x = x + self.interpolate_pos_encoding(x, w, h)\n\n        return self.pos_drop(x)\n\n    def forward(self,\n                x: torch.Tensor,\n                return_all_tokens: Optional[bool] = None,\n                mask: Optional[torch.Tensor] = None) -> torch.Tensor:\n        \"\"\"Forward pass of the Vision Transformer.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor of shape (batch_size, num_channels, width, height).\n        return_all_tokens: Optional[bool] = None\n            Whether to return the embeddings of all tokens.\n        mask: Optional[torch.Tensor] = None\n            Mask tensor of shape (batch_size, width, height) for masked image\n            modeling.\n\n        Returns\n        -------\n            torch.Tensor: Output tensor of shape (batch_size, embed_dim).\n\n        \"\"\"\n        # mim\n        if self.masked_im_modeling:\n            assert mask is not None\n            x = self.prepare_tokens(x, mask=mask)\n        else:\n            x = self.prepare_tokens(x)\n\n        for blk in self.blocks:\n            x = blk(x)\n\n        x = self.norm(x)\n        if self.fc_norm is not None:\n            x[:, 0] = self.fc_norm(  # pylint: disable=not-callable\n                x[:, 1:, :].mean(1)\n            )\n\n        return_all_tokens = (\n            self.return_all_tokens\n            if return_all_tokens is None\n            else return_all_tokens\n        )\n        if return_all_tokens:\n            return x\n        return x[:, 0]\n\n    def extract_feature_maps(self, x: torch.Tensor, output_layers: List[str]) -> Dict[str, torch.Tensor]:\n        \"\"\"Extract feature maps from the given input tensor.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor of shape (B, C, H, W).\n        output_layers: List[str]\n            List of output layer names.\n\n        Returns\n        -------\n        Dict[str, torch.Tensor]\n            Dictionary containing extracted feature maps for each output layer.\n        \"\"\"\n        out_indices = [int(layer[5:]) for layer in output_layers]\n        B, C, H, W = x.shape  # pylint: disable=unused-variable\n        x = self.prepare_tokens(x)\n        Hp = H // self.patch_embed.patch_size\n        Wp = W // self.patch_embed.patch_size\n        features = {}\n        for i, blk in enumerate(self.blocks):\n            x = blk(x)\n            if i in out_indices:\n                xp = x[:, 1:, :].permute(0, 2, 1).reshape(B, -1, Hp, Wp)\n                features[\"block\" + str(i)] = xp\n        return features\n\n    def get_last_selfattention(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Return the self-attention of the last block.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor.\n\n        Returns\n        -------\n        torch.Tensor\n            Self-attention tensor.\n        \"\"\"\n        x = self.prepare_tokens(x)\n        for i, blk in enumerate(self.blocks):\n            if i < len(self.blocks) - 1:\n                x = blk(x)\n            else:\n                # Return attention of the last block.\n                return blk(x, return_attention=True)\n\n    def get_intermediate_layers(self, x: torch.Tensor, n: int = 1) -> List[torch.Tensor]:\n        \"\"\"Return the output tokens from the last ``n`` last blocks.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor.\n        n: int = 1\n            Number of last blocks to extract output from.\n\n        Returns\n        -------\n        List[np.ndarray]\n            List of output tokens from the ``n`` last blocks.\n        \"\"\"\n        x = self.prepare_tokens(x)\n        # Return the output tokens from the ``n`` last blocks.\n        output = []\n        for i, blk in enumerate(self.blocks):\n            x = blk(x)\n            if len(self.blocks) - i <= n:\n                output.append(self.norm(x))\n        return output\n\n    def get_num_layers(self) -> int:\n        \"\"\"Return the number of blocks (layers) in the model.\n\n        Returns\n        -------\n        int\n            Number of blocks.\n        \"\"\"\n        return len(self.blocks)\n\n    def mask_model(self, x: torch.Tensor, mask: torch.BoolTensor):\n        \"\"\"Apply a mask to the input tensor.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor.\n        mask: torch.BoolTensor\n            Mask tensor.\n\n        Returns\n        -------\n        torch.Tensor\n            Masked tensor.\n        \"\"\"\n        x.permute(0, 2, 3, 1)[mask, :] = self.masked_embed.to(x.dtype)\n        return x","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.113841Z","iopub.execute_input":"2024-06-09T08:24:50.114161Z","iopub.status.idle":"2024-06-09T08:24:50.336985Z","shell.execute_reply.started":"2024-06-09T08:24:50.114137Z","shell.execute_reply":"2024-06-09T08:24:50.336181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vit_base(patch_size: int = 16, **kwargs):\n    \"\"\"Create a Base Vision Transformer model (ViT-B).\n\n    Parameters\n    ----------\n    patch_size : int = 16\n        Size of the patches in the input image.\n    **kwargs : keyword arguments\n        Additional arguments to be passed to the VisionTransformer constructor.\n\n    Returns\n    -------\n    model : VisionTransformer\n        The Base Vision Transformer model.\n    \"\"\"\n    model = VisionTransformer(patch_size=patch_size,\n                              embed_dim=768,\n                              depth=12,\n                              num_heads=12,\n                              mlp_ratio=4,\n                              qkv_bias=True,\n                              **kwargs)\n    return model","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.338046Z","iopub.execute_input":"2024-06-09T08:24:50.338321Z","iopub.status.idle":"2024-06-09T08:24:50.351099Z","shell.execute_reply.started":"2024-06-09T08:24:50.338297Z","shell.execute_reply":"2024-06-09T08:24:50.350317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class iBOTViTBase(BaseExtractor):  # noqa\n    \"\"\"Phikon.\n\n    Adapted from HistoSSLscaling (https://github.com/owkin/HistoSSLscaling/blob/main/rl_benchmarks/models/feature_extractors/ibot_vit.py).\n\n    Parameters\n    ----------\n    weights_path : Union[str, Path]\n        Path to the model weights.\n    encoder : str, default='teacher'\n        Name of the model for which we want to load the weights. Valid values are: 'student' and 'teacher'.\n    device : Optional[Union[str, torch.device]], default=None\n        The type of device used for feature extracton. Valid values are: 'cpu' for CPU,\n        'cuda' (uses all GPUs available) or 'cuda:0', 'cuda:0,2' to specify the indices of GPUs to use.\n    \"\"\"\n\n    def __init__(self,\n                 weights_path: Union[str, Path],\n                 encoder: str = \"teacher\",\n                 device: Optional[Union[str, torch.device]] = None) -> None:\n        super().__init__(device=device)\n\n        self.weights_path = weights_path\n        self.encoder = encoder\n\n        self.__build()\n\n    def __build(self) -> None:\n        vit = vit_base(patch_size=16, num_classes=0, use_mean_pooling=False)\n\n        # Restore state dict\n        validate_weights_path(self.weights_path)\n        state_dict = torch.load(self.weights_path, map_location=\"cpu\")\n        state_dict = state_dict[self.encoder]\n        state_dict = {k.replace(\"module.\", \"\"): v for k, v in state_dict.items()}\n        state_dict = {k.replace(\"backbone.\", \"\"): v for k, v in state_dict.items()}\n        vit.load_state_dict(state_dict, strict=False)\n\n        # Create feature extractor\n        self.extractor = prepare_extractor(vit, device=self.device)\n\n    @property\n    def transform(self) -> Callable[[ImageCls], torch.Tensor]:\n        transform_ops = [\n            transforms.ToTensor(),\n            transforms.Normalize(\n                mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225),\n            ),\n        ]\n        return transforms.Compose(transform_ops)\n\n    def __call__(self, images: torch.Tensor) -> np.ndarray:\n        features = self.extractor.get_intermediate_layers(images, 1)\n        features = torch.cat([x[:, 0] for x in features], dim=-1)\n        features_np = features.detach().cpu().numpy()\n        return features_np","metadata":{"papermill":{"duration":0.119108,"end_time":"2023-12-20T23:12:40.092292","exception":false,"start_time":"2023-12-20T23:12:39.973184","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.352245Z","iopub.execute_input":"2024-06-09T08:24:50.352504Z","iopub.status.idle":"2024-06-09T08:24:50.367645Z","shell.execute_reply.started":"2024-06-09T08:24:50.352481Z","shell.execute_reply":"2024-06-09T08:24:50.36682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.4 - Tiling and feature extraction","metadata":{"execution":{"iopub.execute_input":"2023-10-25T14:22:28.950596Z","iopub.status.busy":"2023-10-25T14:22:28.950133Z","iopub.status.idle":"2023-10-25T14:22:28.976573Z","shell.execute_reply":"2023-10-25T14:22:28.975273Z","shell.execute_reply.started":"2023-10-25T14:22:28.950562Z"},"papermill":{"duration":0.025582,"end_time":"2023-12-20T23:12:40.143902","exception":false,"start_time":"2023-12-20T23:12:40.11832","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def _compute_tiles_coordinates_from_mask(mask: np.ndarray,\n                                         fullres_w: int,\n                                         fullres_h: int,\n                                         tile_size: int = 224,\n                                         matter_threshold: float = 0.6,\n                                         seed: Optional[int] = None) -> List[Tuple[int, int]]:\n    matter_mask_as_pil = Image.fromarray((mask * 255).astype(\"uint8\"))\n\n    # Resize mask & get coords of in-matter tiles\n    _num_tiles = (ceil(fullres_w / tile_size), ceil(fullres_h / tile_size))\n    resized_matter_mask = matter_mask_as_pil.resize(size=_num_tiles, resample=Image.BILINEAR)\n    _mask_as_array = np.asarray(resized_matter_mask)\n    coords = list(zip(*np.where(_mask_as_array.T > matter_threshold * 255)))\n\n    rng = np.random.RandomState(seed=seed)\n    rng.shuffle(coords)\n\n    return coords","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.368596Z","iopub.execute_input":"2024-06-09T08:24:50.368848Z","iopub.status.idle":"2024-06-09T08:24:50.381464Z","shell.execute_reply.started":"2024-06-09T08:24:50.368826Z","shell.execute_reply":"2024-06-09T08:24:50.380735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_tiles_coordinates(slide_id: int,\n                          tile_size: int = 224,\n                          matter_threshold: float = 0.6,\n                          seed: Optional[int] = None) -> List[Tuple[int, int]]:\n    \"\"\"Computes the coordinates of tiles with 'enough' matter.\"\"\"\n    fullres_w, fullres_h = get_image_dimensions(slide_id=slide_id)\n    mask = compute_matter_mask(slide_id=slide_id)\n    coords = _compute_tiles_coordinates_from_mask(mask=mask,\n                                                  fullres_w=fullres_w,\n                                                  fullres_h=fullres_h,\n                                                  tile_size=tile_size,\n                                                  matter_threshold=matter_threshold,\n                                                  seed=seed)\n\n    return coords","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.382579Z","iopub.execute_input":"2024-06-09T08:24:50.382894Z","iopub.status.idle":"2024-06-09T08:24:50.392752Z","shell.execute_reply.started":"2024-06-09T08:24:50.38286Z","shell.execute_reply":"2024-06-09T08:24:50.391873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _midpoints(points: List[Tuple[int, int]], swap_coords: bool = False) -> List[Tuple[float, float]]:\n    coords = np.asarray(points)\n    if swap_coords:\n        coords = np.flip(coords, axis=1)\n    idxs = np.lexsort(coords.T)\n    sorted_coords = np.take_along_axis(coords, idxs[:, None], axis=0)\n    diffs = np.diff(sorted_coords, axis=0)\n    idxs = np.where(np.logical_and(diffs[:, 0] == 1, diffs[:, 1] == 0))[0]\n    mid_points = np.asarray([(0.5 * (sorted_coords[i][0] + sorted_coords[i + 1][0]), sorted_coords[i][1]) for i in idxs])\n    return np.flip(mid_points, axis=1) if swap_coords else mid_points","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.393818Z","iopub.execute_input":"2024-06-09T08:24:50.394253Z","iopub.status.idle":"2024-06-09T08:24:50.40713Z","shell.execute_reply.started":"2024-06-09T08:24:50.394223Z","shell.execute_reply":"2024-06-09T08:24:50.406336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mid_tiles_coordinates(tiles_coords: List[Tuple[int, int]], seed: Optional[int] = None) -> List[Tuple[float, float]]:\n    \"\"\"Computes the coordinates of 'middle tiles' (middle of two adjacent tiles).\"\"\"\n    midpoints = np.r_[_midpoints(tiles_coords, False), _midpoints(tiles_coords, True)]\n    np.random.RandomState(seed=seed).shuffle(midpoints)\n    return [(row_[0], row_[1]) for row_ in midpoints]\n","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.408032Z","iopub.execute_input":"2024-06-09T08:24:50.408273Z","iopub.status.idle":"2024-06-09T08:24:50.41788Z","shell.execute_reply.started":"2024-06-09T08:24:50.408251Z","shell.execute_reply":"2024-06-09T08:24:50.416991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def from_tiles_coords_to_pixels_coords(tiles_coords: List[Tuple[Union[int, float], Union[int, float]]],\n                                       fullres_w: int,\n                                       fullres_h: int,\n                                       tile_size: int = 224) -> List[Tuple[int, int]]:\n    \"\"\"Converts tiles coordinates to pixels coordinates (compatible with full resolution image).\"\"\"\n    _px_coords = [(int(tup[0] * tile_size), int(tup[1] * tile_size)) for tup in tiles_coords]\n    px_coords = list(filter(lambda tup: ((tup[0] + tile_size < fullres_w) and (tup[1] + tile_size < fullres_h)), _px_coords))\n    return px_coords","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.419073Z","iopub.execute_input":"2024-06-09T08:24:50.419312Z","iopub.status.idle":"2024-06-09T08:24:50.428754Z","shell.execute_reply.started":"2024-06-09T08:24:50.419291Z","shell.execute_reply":"2024-06-09T08:24:50.428001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def write_coordinates_to_file(px_coords: List[Tuple[int, int]], output_dir: Union[str, Path]) -> None:\n    \"\"\"Writes tiles coordinates as a text file.\"\"\"\n    _output_dir = Path(output_dir)\n    if not _output_dir.is_dir():\n        raise NotADirectoryError(f'Expected {output_dir} to be a directory!')\n\n    fname = str(_output_dir.joinpath('coords.txt'))\n    with open(fname, 'w') as f_:\n        for tup in px_coords:\n            f_.write(f\"{tup[0]},{tup[1]}\\n\")","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.429758Z","iopub.execute_input":"2024-06-09T08:24:50.430127Z","iopub.status.idle":"2024-06-09T08:24:50.442786Z","shell.execute_reply.started":"2024-06-09T08:24:50.430103Z","shell.execute_reply":"2024-06-09T08:24:50.442108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tiling_with_split_png(path_to_png_image: str, output_dir: str, \n                          path_to_coordinates_file: str, \n                          max_tiles_per_slide: int, \n                          tile_size: int) -> None:\n    \"\"\"Uses ``split_png`` to tile a (large) PNG image.\"\"\"\n    command = f\"{PATH_TO_EXECUTABLE} {path_to_png_image} {output_dir} {path_to_coordinates_file} {max_tiles_per_slide} {tile_size}\"\n    result = run(command,\n                 stdout=PIPE,\n                 stderr=PIPE,\n                 universal_newlines=True,\n                 shell=True)\n    if result.returncode != 0:\n        raise RuntimeError(f'Tiling failed with:\\n stderr: {result.stderr}\\n stdout: {result.stdout}')","metadata":{"papermill":{"duration":0.034691,"end_time":"2023-12-20T23:12:40.266618","exception":false,"start_time":"2023-12-20T23:12:40.231927","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.443794Z","iopub.execute_input":"2024-06-09T08:24:50.444077Z","iopub.status.idle":"2024-06-09T08:24:50.453832Z","shell.execute_reply.started":"2024-06-09T08:24:50.444054Z","shell.execute_reply":"2024-06-09T08:24:50.452989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_extractor(extractor_name: Literal[\"ibotvitbasepancancer\"],\n                  extractor_weights_path: Union[str, Path],\n                  device: Optional[Union[str, torch.device]] = None) -> iBOTViTBase:\n    \"\"\"Instantiates the feature extractor.\"\"\"\n    if extractor_name == 'ibotvitbasepancancer':\n        return iBOTViTBase(weights_path=extractor_weights_path, device=device)\n    else:\n        raise ValueError(f\"Got an invalid `extractor_name` ({extractor_name}). Valid values are: `'ibotvitbasepancancer'.\")","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.454859Z","iopub.execute_input":"2024-06-09T08:24:50.455114Z","iopub.status.idle":"2024-06-09T08:24:50.46424Z","shell.execute_reply.started":"2024-06-09T08:24:50.455092Z","shell.execute_reply":"2024-06-09T08:24:50.463401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _inference(extractor: iBOTViTBase, loader: DataLoader) -> np.ndarray:\n    \"\"\"Runs inference on batches of data with ``extractor``.\"\"\"\n    features = []\n\n    with torch.inference_mode():\n\n        for images in loader:\n\n            batch_size = images.shape[0]\n            if extractor.device is not None:\n                images = images.to(extractor.device, non_blocking=True)\n            batch_features = extractor(images)\n            batch_features = batch_features.reshape((batch_size, -1))\n\n            features.append(batch_features)\n\n    return np.concatenate(features, axis=0)","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.465253Z","iopub.execute_input":"2024-06-09T08:24:50.465521Z","iopub.status.idle":"2024-06-09T08:24:50.478325Z","shell.execute_reply.started":"2024-06-09T08:24:50.465498Z","shell.execute_reply":"2024-06-09T08:24:50.477563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _extract_from_single_slide(slide_id: int,\n                               extractor: iBOTViTBase,\n                               tile_size: int = 224,\n                               matter_threshold: float = 0.6,\n                               batch_size: int = 8,\n                               num_workers: int = 1,\n                               max_tiles: Optional[int] = None,\n                               seed: Optional[int] = None) -> Optional[np.ndarray]:\n    \"\"\"Runs feature extraction for a SINGLE slide.\"\"\"\n    with TemporaryDirectory(dir='/kaggle/working') as tmp_dir:\n\n        try:\n            # Step 1: coordinates of tiles with matter\n            tiles_coords = get_tiles_coordinates(slide_id=slide_id,\n                                                 tile_size=tile_size,\n                                                 matter_threshold=matter_threshold,\n                                                 seed=seed)\n\n            if tile_size == 448:\n                # Adding a few more tiles to TMAs\n                mid_tiles_coords = get_mid_tiles_coordinates(tiles_coords)\n                tiles_coords += mid_tiles_coords[:30]  # limit to 30 more tiles\n\n            # Step 2: convert to pixel coordinates & write coordinates to file\n            fullres_w, fullres_h = get_image_dimensions(slide_id=slide_id)\n            px_coords = from_tiles_coords_to_pixels_coords(tiles_coords=tiles_coords,\n                                                           fullres_w=fullres_w,\n                                                           fullres_h=fullres_h,\n                                                           tile_size=tile_size)\n\n            write_coordinates_to_file(px_coords=px_coords[:max_tiles], output_dir=tmp_dir)\n\n            # Step 3: tiling\n            path_fullresolution = get_path_fullres_image(slide_id=slide_id)\n            assert Path(path_fullresolution).is_file() and Path(path_fullresolution).suffix.lower() == '.png'\n            path_coordinates_file = Path(tmp_dir).joinpath('coords.txt')\n            assert Path(path_coordinates_file).is_file()\n            tiling_with_split_png(path_to_png_image=path_fullresolution,\n                                  output_dir=tmp_dir,\n                                  path_to_coordinates_file=path_coordinates_file,\n                                  max_tiles_per_slide=max_tiles,\n                                  tile_size=tile_size)\n\n            # Step 4: feature extraction\n            if tile_size == 448:\n                _transform = transforms.Compose([transforms.Resize((224, 224), Image.BICUBIC), extractor.transform])\n            else:\n                _transform = extractor.transform\n            dataset = ImageDataset(root_dir=tmp_dir, transform=_transform)\n            loader = DataLoader(dataset=dataset,\n                                batch_size=batch_size,\n                                shuffle=True,\n                                num_workers=num_workers,\n                                drop_last=False,\n                                pin_memory=True)\n\n            features = _inference(extractor=extractor, loader=loader)\n\n            # Clear\n            del dataset\n            gc.collect()\n\n            return features\n\n        except Exception:  # noqa\n\n            print(f'--- Preprocessing of slide {slide_id} failed with:')\n            print_exc()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.479448Z","iopub.execute_input":"2024-06-09T08:24:50.480097Z","iopub.status.idle":"2024-06-09T08:24:50.493096Z","shell.execute_reply.started":"2024-06-09T08:24:50.480065Z","shell.execute_reply":"2024-06-09T08:24:50.492142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_features(slide_ids: List[int],\n                     extractor_name: Literal[\"ibotvitbasepancancer\"],\n                     extractor_weights_path: Union[str, Path],\n                     output_dir: Union[str, Path],\n                     tile_size: int = 224,\n                     max_tiles: Optional[int] = None,\n                     seed: Optional[int] = None,\n                     batch_size: int = 8,\n                     num_workers: int = 1,\n                     device: Optional[Union[str, torch.device]] = None, \n                     average_pooling: bool = False, \n                     tma_predictions: Optional[pd.DataFrame] = None) -> None:\n    \"\"\"Runs feature extraction on multiple images.\"\"\"\n    extractor = get_extractor(extractor_name=extractor_name,\n                              extractor_weights_path=extractor_weights_path, \n                              device=device)\n\n    for slide_id in slide_ids:\n        \n        try:\n            \n            # Output dir (and skipping)\n            _output_dir = Path(output_dir).joinpath(f\"{slide_id}\")\n            output_fname = _output_dir.joinpath('features.npy')\n            if not _output_dir.is_dir():\n                _output_dir.mkdir(exist_ok=True, parents=True)\n            elif output_fname.is_file():\n                print(f'Skipping slide ID {slide_id} as features already exist!')\n                continue\n            \n            # Set ``tile_size`` using TMA predictions (if any)\n            if tma_predictions is not None:\n                _slide_ids_preds = tma_predictions['image_id'].tolist()\n                if slide_id in _slide_ids_preds:\n                    # Case 1: get the prediction for this slide_id\n                    slide_is_tma = tma_predictions.loc[tma_predictions['image_id'] == slide_id, 'is_tma'].iloc[0]\n                    _tile_size = (2 * tile_size) if slide_is_tma else tile_size\n                    _matter_threshold = 0.3\n                else:\n                    # Case 2: no predictions for this slide; we assume it's NOT a TMA (?!)\n                    _tile_size = tile_size\n                    _matter_threshold = 0.6\n            else:\n                _tile_size = tile_size\n                _matter_threshold = 0.6\n\n            # Run extraction\n            features = _extract_from_single_slide(slide_id=slide_id,\n                                                  extractor=extractor,\n                                                  tile_size=_tile_size,\n                                                  matter_threshold=_matter_threshold,\n                                                  batch_size=batch_size,\n                                                  num_workers=num_workers,\n                                                  max_tiles=max_tiles,\n                                                  seed=seed)\n\n            # Apply average pooling (if needed) and save features\n            features = features.astype(np.float32)\n            if average_pooling:\n                features = np.mean(features, axis=0, keepdims=True)\n            else:\n                # Check shape\n                n_tiles, _ = features.shape\n                assert n_tiles > 1, 'No tiles in `features.npy`?'\n\n            # Save features\n            np.save(output_fname, features)\n\n            # Clear\n            del features\n            gc.collect()\n\n        except Exception:  # noqa\n\n            continue","metadata":{"papermill":{"duration":0.058418,"end_time":"2023-12-20T23:12:40.495231","exception":false,"start_time":"2023-12-20T23:12:40.436813","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.494521Z","iopub.execute_input":"2024-06-09T08:24:50.494884Z","iopub.status.idle":"2024-06-09T08:24:50.509204Z","shell.execute_reply.started":"2024-06-09T08:24:50.494854Z","shell.execute_reply":"2024-06-09T08:24:50.508355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@ray.remote(num_gpus=0.5, num_cpus=2, max_calls=1)\ndef _remote_extract_features(slide_ids: List[int],\n                             extractor_name: Literal[\"ibotvitbasepancancer\"],\n                             extractor_weights_path: Union[str, Path],\n                             output_dir: Union[str, Path],\n                             tile_size: int = 224,\n                             max_tiles: Optional[int] = None,\n                             seed: Optional[int] = None,\n                             batch_size: int = 8,\n                             num_workers: int = 1,\n                             device: Optional[Union[str, torch.device]] = None, \n                             average_pooling: bool = False, \n                             tma_predictions: Optional[pd.DataFrame] = None) -> None:\n    \"\"\"Helper function as ``ray.remote``.\"\"\"\n    extract_features(slide_ids=slide_ids,\n                     extractor_name=extractor_name,\n                     extractor_weights_path=extractor_weights_path,\n                     output_dir=output_dir,\n                     tile_size=tile_size,\n                     max_tiles=max_tiles,\n                     seed=seed,\n                     batch_size=batch_size,\n                     num_workers=num_workers,\n                     device=device,\n                     average_pooling=average_pooling, \n                     tma_predictions=tma_predictions)","metadata":{"papermill":{"duration":0.036527,"end_time":"2023-12-20T23:12:40.616183","exception":false,"start_time":"2023-12-20T23:12:40.579656","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.510415Z","iopub.execute_input":"2024-06-09T08:24:50.510748Z","iopub.status.idle":"2024-06-09T08:24:50.524625Z","shell.execute_reply.started":"2024-06-09T08:24:50.51071Z","shell.execute_reply":"2024-06-09T08:24:50.52369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.5 - Models","metadata":{"papermill":{"duration":0.02568,"end_time":"2023-12-20T23:12:40.66755","exception":false,"start_time":"2023-12-20T23:12:40.64187","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### 3.5.1 - Chowder","metadata":{"papermill":{"duration":0.025661,"end_time":"2023-12-20T23:12:40.720214","exception":false,"start_time":"2023-12-20T23:12:40.694553","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class MaskedLinear(nn.Linear):\n    \"\"\"Masked Linear layer.\"\"\"\n\n    def __init__(self,\n                 in_features: int,\n                 out_features: int,\n                 mask_value: Union[str, float],\n                 bias: bool = True) -> None:\n        super().__init__(in_features=in_features, \n                         out_features=out_features, \n                         bias=bias)\n        self.mask_value = mask_value\n\n    def forward(self, \n                x: torch.Tensor, \n                mask: Optional[torch.BoolTensor] = None) -> torch.Tensor:\n        \"\"\"Forward pass.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor, shape (B, SEQ_LEN, IN_FEATURES).\n        mask: Optional[torch.BoolTensor] = None\n            True for values that were padded, shape (B, SEQ_LEN, 1),\n\n        Returns\n        -------\n        x: torch.Tensor\n            (B, SEQ_LEN, OUT_FEATURES)\n        \"\"\"\n        x = super(MaskedLinear, self).forward(x)\n        if mask is not None:\n            x = x.masked_fill(mask, float(self.mask_value))\n        return x","metadata":{"papermill":{"duration":0.036485,"end_time":"2023-12-20T23:12:40.782535","exception":false,"start_time":"2023-12-20T23:12:40.74605","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.52586Z","iopub.execute_input":"2024-06-09T08:24:50.526407Z","iopub.status.idle":"2024-06-09T08:24:50.538979Z","shell.execute_reply.started":"2024-06-09T08:24:50.526376Z","shell.execute_reply":"2024-06-09T08:24:50.538155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TilesMLP(nn.Module):\n    \"\"\"TilesMLP layer.\"\"\"\n\n    def __init__(self,\n                 in_features: int,\n                 out_features: int = 1,\n                 hidden: Optional[List[int]] = None,\n                 bias: bool = True,\n                 activation: torch.nn.Module = nn.Sigmoid(),\n                 dropout: Optional[float] = None,\n                 mask_value: float = float(\"-inf\")):\n        super().__init__()\n\n        self.in_features = in_features\n        self.out_features = out_features\n        self.hidden = hidden\n        self.bias = bias\n        self.activation = activation\n        self.dropout = dropout\n        self.mask_value = mask_value\n\n        self.__build()\n\n    def __build(self):\n        self.hidden_layers = torch.nn.ModuleList()\n        in_features = self.in_features\n        if self.hidden is not None:\n            for h in self.hidden:\n                masked_lin = MaskedLinear(in_features=in_features,\n                                          out_features=h,\n                                          bias=self.bias,\n                                          mask_value=self.mask_value)\n                self.hidden_layers.append(masked_lin)\n                self.hidden_layers.append(self.activation)\n                if self.dropout is not None:\n                    self.hidden_layers.append(nn.Dropout(self.dropout))\n                in_features = h\n\n        final_linear = nn.Linear(in_features=in_features,\n                                 out_features=self.out_features,\n                                 bias=self.bias)\n        \n        self.hidden_layers.append(final_linear)\n\n    def forward(self, x: torch.Tensor, mask: Optional[torch.BoolTensor] = None):\n        \"\"\"\n        Parameters\n        ----------\n        x: torch.Tensor\n            (B, N_TILES, IN_FEATURES)\n        mask: Optional[torch.BoolTensor] = None\n            (B, N_TILES), True for values that were padded.\n\n        Returns\n        -------\n        x: torch.Tensor\n            (B, N_TILES, OUT_FEATURES)\n        \"\"\"\n        for layer in self.hidden_layers:\n            if isinstance(layer, MaskedLinear):\n                x = layer(x, mask)\n            else:\n                x = layer(x)\n        return x","metadata":{"papermill":{"duration":0.040787,"end_time":"2023-12-20T23:12:40.849375","exception":false,"start_time":"2023-12-20T23:12:40.808588","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.542414Z","iopub.execute_input":"2024-06-09T08:24:50.543084Z","iopub.status.idle":"2024-06-09T08:24:50.555285Z","shell.execute_reply.started":"2024-06-09T08:24:50.543059Z","shell.execute_reply":"2024-06-09T08:24:50.554455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLP(torch.nn.Sequential):\n    \"\"\"MLP module.\"\"\"\n\n    def __init__(self,\n                 in_features: int,\n                 out_features: int,\n                 hidden: Optional[List[int]] = None,\n                 dropout: Optional[float] = None,\n                 activation: Optional[torch.nn.Module] = torch.nn.Sigmoid(),\n                 bias: bool = True) -> None:\n        \n        # Initialize MLP\n        d_model = in_features\n        layers = []\n\n        if hidden is not None:\n            for i, h in enumerate(hidden):\n                seq = [nn.Linear(d_model, h, bias=bias)]\n                d_model = h\n                \n                if activation is not None:\n                    seq.append(activation)\n\n                if dropout is not None:\n                    seq.append(nn.Dropout(dropout))\n\n                layers.append(torch.nn.Sequential(*seq))\n\n        layers.append(nn.Linear(d_model, out_features))\n\n        super().__init__(*layers)","metadata":{"papermill":{"duration":0.03719,"end_time":"2023-12-20T23:12:40.912935","exception":false,"start_time":"2023-12-20T23:12:40.875745","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.556517Z","iopub.execute_input":"2024-06-09T08:24:50.557058Z","iopub.status.idle":"2024-06-09T08:24:50.568768Z","shell.execute_reply.started":"2024-06-09T08:24:50.557032Z","shell.execute_reply":"2024-06-09T08:24:50.568122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ExtremeLayer(nn.Module):\n    \"\"\"Extreme layer.\n\n    Parameters\n    ----------\n    n_top: Optional[int] = None\n        Number of top tiles to select\n    n_bottom: Optional[int] = None\n        Number of bottom tiles to select\n    dim: int = 1\n        Dimension to select top/bottom tiles from\n    return_indices: bool = False\n        Whether to return the indices of the extreme tiles\n\n    Raises\n    ------\n    ValueError\n        If ``n_top`` and ``n_bottom`` are set to ``None`` or both are 0.\n    \"\"\"\n\n    def __init__(self,\n                 n_top: Optional[int] = None,\n                 n_bottom: Optional[int] = None,\n                 dim: int = 1,\n                 return_indices: bool = False) -> None:\n        super(ExtremeLayer, self).__init__()\n\n        if not (n_top is not None or n_bottom is not None):\n            raise ValueError(\"one of n_top or n_bottom must have a value.\")\n\n        if not ((n_top is not None and n_top > 0) or (n_bottom is not None and n_bottom > 0)):\n            raise ValueError(\"one of n_top or n_bottom must have a value > 0.\")\n\n        self.n_top = n_top\n        self.n_bottom = n_bottom\n        self.dim = dim\n        self.return_indices = return_indices\n\n    def forward(self, \n                x: torch.Tensor, \n                mask: Optional[torch.BoolTensor] = None) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:\n        \"\"\"Forward pass.\n\n        Parameters\n        ----------\n        x: torch.Tensor\n            Input tensor, shape (B, N_TILES, IN_FEATURES).\n        mask: Optional[torch.BoolTensor]\n            True for values that were padded, shape (B, N_TILES, 1).\n\n        Warnings\n        --------\n        If top tiles or bottom tiles is superior to the true number of tiles in\n        the input then padded tiles will be selected and their value will be 0.\n\n        Returns\n        -------\n        values: torch.Tensor\n            Extreme tiles, shape (B, N_TOP + N_BOTTOM).\n        indices: torch.Tensor\n            If ``self.return_indices=True``, return extreme tiles' indices.\n        \"\"\"\n\n        if self.n_top and self.n_bottom and ((self.n_top + self.n_bottom) > x.shape[self.dim]):\n            warn(f\"Sum of tops is larger than the input tensor shape for dimension {self.dim}: {self.n_top + self.n_bottom} > {x.shape[self.dim]}. Values will appear twice (in top and in bottom)\")\n\n        top, bottom = None, None\n        top_idx, bottom_idx = None, None\n        if mask is not None:\n            if self.n_top:\n                top, top_idx = x.masked_fill(mask, float(\"-inf\")).topk(k=self.n_top, sorted=True, dim=self.dim)\n                top_mask = top.eq(float(\"-inf\"))\n                if top_mask.any():\n                    warn(\"The top tiles contain masked values, they will be set to zero.\")\n                    top[top_mask] = 0\n\n            if self.n_bottom:\n                bottom, bottom_idx = x.masked_fill(mask, float(\"inf\")).topk(k=self.n_bottom, largest=False, sorted=True, dim=self.dim)\n                bottom_mask = bottom.eq(float(\"inf\"))\n                if bottom_mask.any():\n                    warn(\"The bottom tiles contain masked values, they will be set to zero.\")\n                    bottom[bottom_mask] = 0\n        else:\n            if self.n_top:\n                top, top_idx = x.topk(k=self.n_top, sorted=True, dim=self.dim)\n            if self.n_bottom:\n                bottom, bottom_idx = x.topk(k=self.n_bottom, largest=False, sorted=True, dim=self.dim)\n\n        if top is not None and bottom is not None:\n            values = torch.cat([top, bottom], dim=self.dim)\n            indices = torch.cat([top_idx, bottom_idx], dim=self.dim)\n        elif top is not None:\n            values = top\n            indices = top_idx\n        elif bottom is not None:\n            values = bottom\n            indices = bottom_idx\n        else:\n            raise ValueError\n\n        if self.return_indices:\n            return values, indices\n        else:\n            return values\n\n    def extra_repr(self) -> str:\n        \"\"\"Format representation.\"\"\"\n        return f\"n_top={self.n_top}, n_bottom={self.n_bottom}\"","metadata":{"papermill":{"duration":0.047221,"end_time":"2023-12-20T23:12:41.176505","exception":false,"start_time":"2023-12-20T23:12:41.129284","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.57069Z","iopub.execute_input":"2024-06-09T08:24:50.571002Z","iopub.status.idle":"2024-06-09T08:24:50.588926Z","shell.execute_reply.started":"2024-06-09T08:24:50.570933Z","shell.execute_reply":"2024-06-09T08:24:50.588027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Chowder(nn.Module):\n    \"\"\"Chowder model.\"\"\"\n\n    def __init__(self,\n                 in_features: int,\n                 out_features: int,\n                 n_extreme: Optional[int] = None,\n                 n_top: Optional[int] = None,\n                 n_bottom: Optional[int] = None,\n                 tiles_mlp_hidden: Optional[List[int]] = None,\n                 tiles_mlp_dropout: Optional[float] = None,\n                 tiles_mlp_activation: Optional[nn.Module] = nn.Sigmoid(),\n                 mlp_hidden: Optional[List[int]] = None,\n                 mlp_dropout: Optional[float] = None,\n                 mlp_activation: Optional[nn.Module] = nn.Sigmoid(),\n                 bias: bool = True, \n                 mask_value: float = float(\"-inf\"),) -> None:\n        super().__init__()\n\n        if n_extreme is not None:\n            warn(DeprecationWarning(f\"Use `n_extreme=None, n_top={n_extreme if n_top is None else n_top}, n_bottom={n_extreme if n_bottom is None else n_bottom}` instead.\"))\n\n            if n_top is not None:\n                warn(DeprecationWarning(f\"Overriding `n_top={n_top}` with `n_top=n_extreme={n_extreme}`.\"))\n\n            if n_bottom is not None:\n                warn(DeprecationWarning(f\"Overriding `n_bottom={n_bottom}` with `n_bottom=n_extreme={n_extreme}`.\"))\n\n            n_top = n_extreme\n            n_bottom = n_extreme\n\n        if n_top is None and n_bottom is None:\n            raise ValueError(\"At least one of `n_top` or `n_bottom` must not be None.\")\n\n        self.score_model = TilesMLP(in_features=in_features,\n                                    out_features=out_features,\n                                    hidden=tiles_mlp_hidden,\n                                    bias=bias,\n                                    activation=tiles_mlp_activation,\n                                    dropout=tiles_mlp_dropout,\n                                    mask_value=mask_value)\n\n        self.score_model.apply(self.weight_initialization)\n\n        self.extreme_layer = ExtremeLayer(n_top=n_top, n_bottom=n_bottom)\n\n        self.mlp = MLP(in_features=n_top + n_bottom,\n                       out_features=1,\n                       hidden=mlp_hidden,\n                       dropout=mlp_dropout,\n                       activation=mlp_activation)\n        \n        self.mlp.apply(self.weight_initialization)\n\n    @staticmethod\n    def weight_initialization(module: torch.nn.Module) -> None:\n        if isinstance(module, torch.nn.Linear):\n            torch.nn.init.xavier_uniform_(module.weight)\n            if module.bias is not None:\n                module.bias.data.fill_(0.0)\n\n    def forward(self,\n                features: torch.Tensor,\n                mask: Optional[torch.BoolTensor] = None) -> torch.Tensor:\n        \"\"\"\n        Parameters\n        ----------\n        features: torch.Tensor\n            (B, N_TILES, IN_FEATURES)\n        mask: Optional[torch.BoolTensor] = None\n            (B, N_TILES, 1), True for values that were padded.\n\n        Returns\n        -------\n        logits, extreme_scores: Tuple[torch.Tensor, torch.Tensor]:\n            (B, OUT_FEATURES), (B, N_TOP + N_BOTTOM, OUT_FEATURES)\n        \"\"\"\n        scores = self.score_model(x=features, mask=mask)\n        extreme_scores = self.extreme_layer(x=scores, mask=mask)\n        y = self.mlp(extreme_scores.transpose(1, 2))\n        return y.squeeze(2), scores","metadata":{"papermill":{"duration":0.044483,"end_time":"2023-12-20T23:12:41.247735","exception":false,"start_time":"2023-12-20T23:12:41.203252","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.590148Z","iopub.execute_input":"2024-06-09T08:24:50.590466Z","iopub.status.idle":"2024-06-09T08:24:50.605813Z","shell.execute_reply.started":"2024-06-09T08:24:50.590418Z","shell.execute_reply":"2024-06-09T08:24:50.605001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3.5.2 - ModelEnsemble","metadata":{"papermill":{"duration":0.026314,"end_time":"2023-12-20T23:12:41.300705","exception":false,"start_time":"2023-12-20T23:12:41.274391","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class ModelEnsemble(nn.ModuleList):\n    \n    def __init__(self, models: List[nn.Module]) -> None:\n        super().__init__(modules=models)\n\n    def forward(self, \n                x: torch.Tensor, \n                mask: Optional[torch.BoolTensor] = None) -> torch.Tensor:\n        \"\"\"Forward pass.\"\"\"        \n        predictions_, scores_ = [], []\n        for model in self:\n            logits, attn_scores = model(x, mask)\n            predictions_.append(logits.unsqueeze(-1))\n            scores_.append(torch.mean(attn_scores, dim=1, keepdim=True).unsqueeze(-1))\n        predictions_ = torch.cat(predictions_, dim=2)\n        scores_ = torch.cat(scores_, dim=2)\n        return predictions_, scores_","metadata":{"papermill":{"duration":0.036937,"end_time":"2023-12-20T23:12:41.364473","exception":false,"start_time":"2023-12-20T23:12:41.327536","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.606868Z","iopub.execute_input":"2024-06-09T08:24:50.607229Z","iopub.status.idle":"2024-06-09T08:24:50.620682Z","shell.execute_reply.started":"2024-06-09T08:24:50.607198Z","shell.execute_reply":"2024-06-09T08:24:50.619908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ModelsBunch(nn.ModuleList):\n    \"\"\"A bunch/list of (trained) MIL models.\"\"\"\n\n    def __init__(self,\n                 base_model: Union[nn.Module, ModelEnsemble],\n                 models_dir: Union[str, Path],\n                 reps_folds_to_use: Optional[Tuple[Tuple[int, int]]] = None,\n                 device: Optional[str] = None):\n\n        models = []\n        \n        if reps_folds_to_use is None:\n            _reps_folds_to_use = list(itertools.product(range(3), range(5)))\n        else:\n            _reps_folds_to_use = reps_folds_to_use\n\n        for tup in _reps_folds_to_use:\n            state_dict_fpath = Path(models_dir).joinpath(f'rep_{tup[0]}_fold_{tup[1]}_chowder_ensemble.pt')\n            if not state_dict_fpath.is_file():\n                raise FileNotFoundError(f'Could not find {state_dict_fpath}!')\n\n            # Load state dict\n            state_dict = torch.load(state_dict_fpath)\n\n            # Restore state_dict\n            model = deepcopy(base_model)\n            model.load_state_dict(state_dict, strict=True)\n\n            # Prepare model\n            model.to(device, non_blocking=True)\n            model.eval()\n\n            models.append(model)\n    \n        super().__init__(modules=models)\n\n    def forward(self):\n        raise NotImplementedError('This class is only intended to be used as an iterable!')","metadata":{"papermill":{"duration":0.038176,"end_time":"2023-12-20T23:12:41.429349","exception":false,"start_time":"2023-12-20T23:12:41.391173","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.6216Z","iopub.execute_input":"2024-06-09T08:24:50.621847Z","iopub.status.idle":"2024-06-09T08:24:50.631837Z","shell.execute_reply.started":"2024-06-09T08:24:50.621815Z","shell.execute_reply":"2024-06-09T08:24:50.63086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(model: Union[nn.Module, ModelEnsemble],\n            loader: DataLoader,\n            device: Optional[str] = None) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"Inference with a trained MIL model.\"\"\"\n    model.eval()\n\n    predictions_, metadata_ = [], []\n    with torch.inference_mode():\n        for batch in loader:\n\n            features, masks, _, slide_ids = batch\n            masks = masks.type(torch.BoolTensor).to(device, non_blocking=True)\n            features = features.float().to(device, non_blocking=True)\n            batch_size = features.shape[0]\n\n            # Forward pass\n            logits, _ = model.forward(features, masks)\n\n            logits_np = logits.detach().cpu().numpy()\n            if batch_size == 1:\n                if isinstance(model, ModelEnsemble):\n                    logits_np = logits_np.reshape((batch_size, 5, len(model)))\n                else:\n                    logits_np = logits_np.reshape((batch_size, 5))\n            predictions_.append(logits_np)\n            metadata_.append(slide_ids.detach().cpu().numpy())\n\n    predictions_ = np.concatenate(predictions_, axis=0)\n    metadata_ = np.concatenate(metadata_, axis=0)\n    return predictions_, metadata_","metadata":{"papermill":{"duration":0.045106,"end_time":"2023-12-20T23:12:41.501039","exception":false,"start_time":"2023-12-20T23:12:41.455933","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.632922Z","iopub.execute_input":"2024-06-09T08:24:50.633219Z","iopub.status.idle":"2024-06-09T08:24:50.646614Z","shell.execute_reply.started":"2024-06-09T08:24:50.633195Z","shell.execute_reply":"2024-06-09T08:24:50.645768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_with_ensemble(dataset: FeaturesDataset, \n                          models_dir: Union[str, Path], \n                          batch_size: int = 8, \n                          num_workers: int = 4,\n                          max_tiles: int = 200,\n                          device: Optional[str] = None) -> Tuple[np.ndarray, np.ndarray]:\n\n    # Loader\n    _test_collate_fn = partial(pad_collate_fn, **{'max_len': max_tiles})\n    test_loader = DataLoader(dataset=dataset,\n                             batch_size=batch_size,\n                             shuffle=False,\n                             num_workers=num_workers,\n                             pin_memory=True,\n                             drop_last=False,\n                             collate_fn=_test_collate_fn)\n\n    base_model = ModelEnsemble([Chowder(**CHOWDER_KWARGS) for _ in range(50)])\n    \n    models = ModelsBunch(base_model=base_model, \n                         models_dir=models_dir, \n                         device=device, \n                         reps_folds_to_use=((0, 0), (2, 3), (2, 4)))\n\n    # Inference with models ensemble\n    all_logits = []\n    metadata = None\n    for model in models:\n\n        _logits, metadata = predict(model=model, loader=test_loader, device=device)\n        all_logits.append(_logits[None,...])\n\n    all_logits = np.concatenate(all_logits, axis=0)\n    return all_logits, metadata","metadata":{"papermill":{"duration":0.044766,"end_time":"2023-12-20T23:12:41.576126","exception":false,"start_time":"2023-12-20T23:12:41.53136","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.647674Z","iopub.execute_input":"2024-06-09T08:24:50.647963Z","iopub.status.idle":"2024-06-09T08:24:50.661627Z","shell.execute_reply.started":"2024-06-09T08:24:50.647921Z","shell.execute_reply":"2024-06-09T08:24:50.660801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calibrate_predictions(logits: np.ndarray, fname_suffix: str) -> np.ndarray:\n    \"\"\"Calibrates model predictions.\"\"\"\n    assert logits.ndim == 3  # (n_samples, n_classes, n_models)\n\n    with open(f'/kaggle/input/models-ensembles-final/lr_calibration_chowder_{fname_suffix}.pkl', 'rb') as fb_:\n        clfs = pickle.load(fb_)\n\n    return np.stack([clfs[model_idx].predict_proba(logits[:, :, model_idx]) for model_idx in range(logits.shape[-1])], axis=2)","metadata":{"execution":{"iopub.status.busy":"2024-06-09T08:24:50.662754Z","iopub.execute_input":"2024-06-09T08:24:50.663117Z","iopub.status.idle":"2024-06-09T08:24:50.674997Z","shell.execute_reply.started":"2024-06-09T08:24:50.663081Z","shell.execute_reply":"2024-06-09T08:24:50.674164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_predictions_entropy(logits: np.ndarray, temperature: float = 1.) -> np.ndarray:\n    \"\"\"Computes the entropy of predictions.\"\"\"\n    assert logits.ndim == 4  # logits: (reps_folds, n_samples, n_classes, n_models)\n    entropy_ = []\n    for idx in range(logits.shape[0]):\n        \n        probas = softmax(temperature * logits[idx], axis=1)\n        probas = probas.mean(2)  # average across models\n        entropy_.append(stats.entropy(probas, axis=1))\n    \n    return np.c_[entropy_]","metadata":{"papermill":{"duration":0.037875,"end_time":"2023-12-20T23:12:41.650331","exception":false,"start_time":"2023-12-20T23:12:41.612456","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.676037Z","iopub.execute_input":"2024-06-09T08:24:50.676465Z","iopub.status.idle":"2024-06-09T08:24:50.685134Z","shell.execute_reply.started":"2024-06-09T08:24:50.67644Z","shell.execute_reply":"2024-06-09T08:24:50.684344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def majority_voting(y_pred: np.ndarray) -> np.ndarray:\n    \"\"\"Applies majority voting to predictions.\"\"\"\n    _func = lambda x: np.argmax(np.bincount(x))\n    return np.apply_along_axis(_func, axis=1, arr=y_pred)","metadata":{"papermill":{"duration":0.036468,"end_time":"2023-12-20T23:12:41.713644","exception":false,"start_time":"2023-12-20T23:12:41.677176","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.686123Z","iopub.execute_input":"2024-06-09T08:24:50.686353Z","iopub.status.idle":"2024-06-09T08:24:50.699288Z","shell.execute_reply.started":"2024-06-09T08:24:50.686333Z","shell.execute_reply":"2024-06-09T08:24:50.698397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3.5.3 - TMA detector","metadata":{"papermill":{"duration":0.027156,"end_time":"2023-12-20T23:12:41.888282","exception":false,"start_time":"2023-12-20T23:12:41.861126","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def run_tma_detector(slide_ids: List[int], \n                     batch_size: int = 64, \n                     num_workers: int = 4, \n                     pin_memory: bool = True) -> pd.DataFrame:\n    \"\"\"Runs the TMA detector.\"\"\"\n    # ResNet18 for feature extraction\n    backbone = resnet18(weights=None)\n    state_dict = torch.load(WEIGHTS_DICT['resnet18'])\n    backbone.load_state_dict(state_dict)\n    feature_extractor = create_feature_extractor(model=backbone, return_nodes=['avgpool'])\n    feature_extractor.to('cuda')\n    feature_extractor.eval()\n    feature_extractor.requires_grad_(False)\n    base_transform = ResNet18_Weights.IMAGENET1K_V1.transforms()\n    \n    # LR for TMA prediction\n    with open(str(WEIGHTS_DICT['tma_detector']), 'rb') as fb:\n        clf = joblib.load(fb)\n\n    # Dataset\n    _img_paths = [Path(get_path_thumbnail_image(slide_id)) for slide_id in slide_ids]\n    img_paths = list(filter(lambda fpath: fpath.is_file(), _img_paths))\n    dataset = ThumbnailsDataset(img_paths=img_paths, transform=base_transform)\n    \n    # DataLoader\n    loader = DataLoader(dataset=dataset, \n                        batch_size=batch_size, \n                        num_workers=num_workers, \n                        pin_memory=pin_memory)\n    \n    # Run feature extraction with ResNet18\n    features, slide_ids_ = [], []\n    with torch.inference_mode(): \n        for batch in loader:\n            images, ids = batch\n            images = images.to('cuda', non_blocking=True)\n\n            batch_features = feature_extractor(images)['avgpool'].squeeze()\n\n            batch_size = images.shape[0]\n            batch_features_np = batch_features.detach().cpu().numpy()\n            if batch_size == 1:\n                batch_features_np = batch_features_np.reshape((1, -1))\n            features.append(batch_features_np)\n            slide_ids_.append(ids.detach().cpu().numpy())\n    \n    features = np.concatenate(features, axis=0)\n    slide_ids_ = np.concatenate(slide_ids_, axis=0)\n    \n    y_pred = clf.predict(features)\n    \n    return pd.DataFrame({'image_id': slide_ids_, 'is_tma': y_pred.astype(bool)})","metadata":{"papermill":{"duration":0.043975,"end_time":"2023-12-20T23:12:41.959588","exception":false,"start_time":"2023-12-20T23:12:41.915613","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.700579Z","iopub.execute_input":"2024-06-09T08:24:50.701369Z","iopub.status.idle":"2024-06-09T08:24:50.713312Z","shell.execute_reply.started":"2024-06-09T08:24:50.701344Z","shell.execute_reply":"2024-06-09T08:24:50.712286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4 - Tiling and feature extraction (run)","metadata":{"papermill":{"duration":0.02725,"end_time":"2023-12-20T23:12:42.52774","exception":false,"start_time":"2023-12-20T23:12:42.50049","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_slide_ids_sorted_by_image_size(df: pd.DataFrame) -> List[int]:\n    img_size_px = df['image_width'].values * df['image_height'].values\n    sort_idxs = np.argsort(img_size_px)\n    return df['image_id'].values[sort_idxs].tolist()","metadata":{"papermill":{"duration":0.034637,"end_time":"2023-12-20T23:12:42.589488","exception":false,"start_time":"2023-12-20T23:12:42.554851","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.714374Z","iopub.execute_input":"2024-06-09T08:24:50.714656Z","iopub.status.idle":"2024-06-09T08:24:50.726931Z","shell.execute_reply.started":"2024-06-09T08:24:50.714632Z","shell.execute_reply":"2024-06-09T08:24:50.726106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_TRAINING:\n    info = pd.read_csv(Path(DATA_DIR).joinpath('train.csv'))\nelse:\n    info = pd.read_csv(Path(DATA_DIR).joinpath('test.csv'))\n\nimage_ids = get_slide_ids_sorted_by_image_size(info)","metadata":{"papermill":{"duration":0.050559,"end_time":"2023-12-20T23:12:42.667086","exception":false,"start_time":"2023-12-20T23:12:42.616527","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.728051Z","iopub.execute_input":"2024-06-09T08:24:50.728347Z","iopub.status.idle":"2024-06-09T08:24:50.754811Z","shell.execute_reply.started":"2024-06-09T08:24:50.728324Z","shell.execute_reply":"2024-06-09T08:24:50.753997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Reorder slide IDs to avoid processing all small/large (TMAs) images at the same time\nchunk1 = image_ids[::2]\nchunk2 = image_ids[1::2][::-1]\nreordered_image_ids = [chunk1[idx // 2] if idx % 2 == 0 else chunk2[idx // 2] for idx in range(len(image_ids))]\nassert set(reordered_image_ids) == set(image_ids)\nassert len(reordered_image_ids) == len(image_ids)","metadata":{"papermill":{"duration":0.037223,"end_time":"2023-12-20T23:12:42.732106","exception":false,"start_time":"2023-12-20T23:12:42.694883","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:50.755994Z","iopub.execute_input":"2024-06-09T08:24:50.756318Z","iopub.status.idle":"2024-06-09T08:24:50.762712Z","shell.execute_reply.started":"2024-06-09T08:24:50.756286Z","shell.execute_reply":"2024-06-09T08:24:50.761783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/features/","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2024-06-09T08:24:50.763722Z","iopub.execute_input":"2024-06-09T08:24:50.764073Z","iopub.status.idle":"2024-06-09T08:24:51.712248Z","shell.execute_reply.started":"2024-06-09T08:24:50.764042Z","shell.execute_reply":"2024-06-09T08:24:51.711069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/tmp","metadata":{"papermill":{"duration":0.981275,"end_time":"2023-12-20T23:12:46.997859","exception":false,"start_time":"2023-12-20T23:12:46.016584","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:51.713908Z","iopub.execute_input":"2024-06-09T08:24:51.714269Z","iopub.status.idle":"2024-06-09T08:24:52.729399Z","shell.execute_reply.started":"2024-06-09T08:24:51.714239Z","shell.execute_reply":"2024-06-09T08:24:52.727934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.1 - TMA detection","metadata":{"papermill":{"duration":0.028622,"end_time":"2023-12-20T23:12:47.056005","exception":false,"start_time":"2023-12-20T23:12:47.027383","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df_tma = run_tma_detector(slide_ids=image_ids,\n                          batch_size=256, \n                          num_workers=4, \n                          pin_memory=True)","metadata":{"papermill":{"duration":8.988047,"end_time":"2023-12-20T23:12:56.072854","exception":false,"start_time":"2023-12-20T23:12:47.084807","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:52.731135Z","iopub.execute_input":"2024-06-09T08:24:52.731487Z","iopub.status.idle":"2024-06-09T08:24:54.800166Z","shell.execute_reply.started":"2024-06-09T08:24:52.731455Z","shell.execute_reply":"2024-06-09T08:24:54.799058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" ## 4.2 - Feature extraction (with Ray)","metadata":{"papermill":{"duration":0.028621,"end_time":"2023-12-20T23:12:56.131298","exception":false,"start_time":"2023-12-20T23:12:56.102677","status":"completed"},"tags":[]}},{"cell_type":"code","source":"ray.init(num_cpus=4, num_gpus=1, _temp_dir='/kaggle/working/tmp', _plasma_directory='/kaggle/working/tmp')","metadata":{"papermill":{"duration":5.799808,"end_time":"2023-12-20T23:13:01.960434","exception":false,"start_time":"2023-12-20T23:12:56.160626","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:24:54.801635Z","iopub.execute_input":"2024-06-09T08:24:54.801959Z","iopub.status.idle":"2024-06-09T08:25:00.044344Z","shell.execute_reply.started":"2024-06-09T08:24:54.801916Z","shell.execute_reply":"2024-06-09T08:25:00.043182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mid = len(reordered_image_ids) // 2\nslide_ids_chunks = [reordered_image_ids[:mid], reordered_image_ids[mid:]]","metadata":{"papermill":{"duration":0.053655,"end_time":"2023-12-20T23:13:02.059464","exception":false,"start_time":"2023-12-20T23:13:02.005809","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:00.045774Z","iopub.execute_input":"2024-06-09T08:25:00.047473Z","iopub.status.idle":"2024-06-09T08:25:00.053708Z","shell.execute_reply.started":"2024-06-09T08:25:00.047432Z","shell.execute_reply":"2024-06-09T08:25:00.052527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ray.get([_remote_extract_features.remote(slide_ids=_chunk,\n                                         extractor_name='ibotvitbasepancancer',\n                                         extractor_weights_path=WEIGHTS_DICT['ibotvitbasepancancer'],\n                                         output_dir=OUTPUT_DIR,\n                                         tile_size=224,\n                                         max_tiles=200,\n                                         seed=123456,\n                                         batch_size=128,\n                                         num_workers=0,\n                                         device='cuda',\n                                         average_pooling=False, \n                                         tma_predictions=df_tma)\n         for _chunk in slide_ids_chunks])","metadata":{"papermill":{"duration":50.918443,"end_time":"2023-12-20T23:13:53.008148","exception":false,"start_time":"2023-12-20T23:13:02.089705","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:00.055111Z","iopub.execute_input":"2024-06-09T08:25:00.05607Z","iopub.status.idle":"2024-06-09T08:25:07.733529Z","shell.execute_reply.started":"2024-06-09T08:25:00.056036Z","shell.execute_reply":"2024-06-09T08:25:07.731406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ray.shutdown()","metadata":{"papermill":{"duration":1.433191,"end_time":"2023-12-20T23:13:54.474845","exception":false,"start_time":"2023-12-20T23:13:53.041654","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.734783Z","iopub.status.idle":"2024-06-09T08:25:07.735203Z","shell.execute_reply.started":"2024-06-09T08:25:07.73499Z","shell.execute_reply":"2024-06-09T08:25:07.735007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5 - Predictions & Submission","metadata":{"papermill":{"duration":0.029319,"end_time":"2023-12-20T23:13:54.534545","exception":false,"start_time":"2023-12-20T23:13:54.505226","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## 5.1 Predict tissue types (`CC`, `EC`, `HGSC`, `LGSC`, `MC`)","metadata":{"papermill":{"duration":0.029589,"end_time":"2023-12-20T23:13:54.661742","exception":false,"start_time":"2023-12-20T23:13:54.632153","status":"completed"},"tags":[]}},{"cell_type":"code","source":"extracted_features_fpaths = list(Path('/kaggle/working/features').rglob('features.npy'))\ntest_ds = FeaturesDataset(features_fpaths=extracted_features_fpaths,\n                          labels=['CC'] * len(extracted_features_fpaths),\n                          max_tiles=200, \n                          shuffle=False)","metadata":{"papermill":{"duration":0.038301,"end_time":"2023-12-20T23:13:54.729933","exception":false,"start_time":"2023-12-20T23:13:54.691632","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.736359Z","iopub.status.idle":"2024-06-09T08:25:07.736731Z","shell.execute_reply.started":"2024-06-09T08:25:07.736537Z","shell.execute_reply":"2024-06-09T08:25:07.736552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compute logits\ny_logits, slide_ids_in = predict_with_ensemble(dataset=test_ds, \n                                               models_dir='/kaggle/input/models-ensembles-final', \n                                               batch_size=32, \n                                               num_workers=4, \n                                               max_tiles=200, \n                                               device='cuda')","metadata":{"papermill":{"duration":11.166857,"end_time":"2023-12-20T23:14:05.926709","exception":false,"start_time":"2023-12-20T23:13:54.759852","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.73792Z","iopub.status.idle":"2024-06-09T08:25:07.738306Z","shell.execute_reply.started":"2024-06-09T08:25:07.738129Z","shell.execute_reply":"2024-06-09T08:25:07.738145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert y_logits.ndim == 4  # logits: (reps_folds, n_samples, n_classes, n_models)\nassert y_logits.shape[0] == 3\nprobas1 = calibrate_predictions(logits=y_logits[0], fname_suffix='rep_0_fold_0')\nprobas2 = calibrate_predictions(logits=y_logits[1], fname_suffix='rep_2_fold_3')\nprobas3 = calibrate_predictions(logits=y_logits[2], fname_suffix='rep_2_fold_4')","metadata":{"papermill":{"duration":0.045194,"end_time":"2023-12-20T23:14:06.002527","exception":false,"start_time":"2023-12-20T23:14:05.957333","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.739304Z","iopub.status.idle":"2024-06-09T08:25:07.739632Z","shell.execute_reply.started":"2024-06-09T08:25:07.739469Z","shell.execute_reply":"2024-06-09T08:25:07.739483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Filtering models\nprobas1_ = probas1[..., [42, 24, 10, 49, 33]].mean(2, keepdims=True)\nprobas2_ = probas2[..., [37, 33, 48,  3, 31, 43, 40,  2, 24, 36]].mean(2, keepdims=True)\nprobas3_ = probas3.mean(2, keepdims=True)\nprobas_ = np.concatenate([probas1_, probas2_, probas3_], axis=2)","metadata":{"execution":{"iopub.status.busy":"2024-06-09T08:25:07.740761Z","iopub.status.idle":"2024-06-09T08:25:07.741158Z","shell.execute_reply.started":"2024-06-09T08:25:07.740924Z","shell.execute_reply":"2024-06-09T08:25:07.740938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"yp = probas_.mean(2)","metadata":{"execution":{"iopub.status.busy":"2024-06-09T08:25:07.742699Z","iopub.status.idle":"2024-06-09T08:25:07.743067Z","shell.execute_reply.started":"2024-06-09T08:25:07.742872Z","shell.execute_reply":"2024-06-09T08:25:07.742886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.2 Predict `Other` class","metadata":{"papermill":{"duration":0.029243,"end_time":"2023-12-20T23:14:06.129616","exception":false,"start_time":"2023-12-20T23:14:06.100373","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Predict 'other'\ny_ent = stats.entropy(yp, axis=1)\npredictions = np.where(y_ent > ENTROPY_THRESHOLD, OUTLIERS_CLASS, yp.argmax(1))","metadata":{"papermill":{"duration":0.041245,"end_time":"2023-12-20T23:14:06.200859","exception":false,"start_time":"2023-12-20T23:14:06.159614","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.745781Z","iopub.status.idle":"2024-06-09T08:25:07.746584Z","shell.execute_reply.started":"2024-06-09T08:25:07.746331Z","shell.execute_reply":"2024-06-09T08:25:07.74635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove feature files (in case it prevents ``submission.csv`` from being found)\n!rm -rf /kaggle/working/features","metadata":{"papermill":{"duration":1.021855,"end_time":"2023-12-20T23:14:07.252494","exception":false,"start_time":"2023-12-20T23:14:06.230639","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.747918Z","iopub.status.idle":"2024-06-09T08:25:07.74829Z","shell.execute_reply.started":"2024-06-09T08:25:07.748125Z","shell.execute_reply":"2024-06-09T08:25:07.748139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Also remove tmp directory\n!rm -rf /kaggle/working/tmp","metadata":{"papermill":{"duration":1.03905,"end_time":"2023-12-20T23:14:08.321654","exception":false,"start_time":"2023-12-20T23:14:07.282604","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.749554Z","iopub.status.idle":"2024-06-09T08:25:07.749888Z","shell.execute_reply.started":"2024-06-09T08:25:07.749726Z","shell.execute_reply":"2024-06-09T08:25:07.74974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.3 - Create submission","metadata":{}},{"cell_type":"code","source":"def make_submission(df_preds: pd.DataFrame, df_other: Optional[pd.DataFrame] = None) -> None:\n    test_info = pd.read_csv(Path(DATA_DIR).joinpath('test.csv'))\n\n    out = pd.merge(left=test_info[['image_id']], \n                   right=df_preds, \n                   on=['image_id'], \n                   how='left', \n                   sort=False)\n\n    # Fill NaN with random labels\n    num_nans = out['label'].isna().sum()\n    if num_nans > 0:\n        random_labels = np.random.choice(['CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other'], size=num_nans)\n        out.loc[out['label'].isna(), 'label'] = random_labels\n    \n    # Use `df_other` to predict the `Other` class\n    if df_other is not None:\n        out = pd.merge(left=out, \n                       right=df_other, \n                       on=['image_id'], \n                       how='left', \n                       sort=False)\n        \n        # Set label 'Other' where `is_other=1`\n        out.loc[out['is_other'] == 1, 'label'] = 'Other'\n        \n        # Drop the `is_other` column\n        out.drop(columns=['is_other'], inplace=True)\n    \n    out.to_csv(\"/kaggle/working/submission.csv\", index=False)","metadata":{"papermill":{"duration":0.043835,"end_time":"2023-12-20T23:14:08.396641","exception":false,"start_time":"2023-12-20T23:14:08.352806","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.751024Z","iopub.status.idle":"2024-06-09T08:25:07.751358Z","shell.execute_reply.started":"2024-06-09T08:25:07.751195Z","shell.execute_reply":"2024-06-09T08:25:07.75121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the submission file\npredictions_as_str = [_CLS_INV_MAPPING.get(str(yy)) for yy in predictions]\ndf_predictions = pd.DataFrame({'image_id': slide_ids_in.tolist(), 'label': predictions_as_str})\n\nmake_submission(df_preds=df_predictions, df_other=None)\n\nif not Path('/kaggle/working/submission.csv').is_file():\n    raise FileNotFoundError('Submission CSV file not found!')","metadata":{"papermill":{"duration":0.072833,"end_time":"2023-12-20T23:14:08.499425","exception":false,"start_time":"2023-12-20T23:14:08.426592","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-09T08:25:07.753177Z","iopub.status.idle":"2024-06-09T08:25:07.753884Z","shell.execute_reply.started":"2024-06-09T08:25:07.753638Z","shell.execute_reply":"2024-06-09T08:25:07.753665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}