{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6661702,"sourceType":"datasetVersion","datasetId":3844162},{"sourceId":146934283,"sourceType":"kernelVersion"}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Baseline","metadata":{}},{"cell_type":"code","source":"# !yes | sudo dpkg -i /kaggle/input/libvips-pyvips-installation-and-getting-started/libvips/*.deb\n# !pip install /kaggle/input/libvips-pyvips-installation-and-getting-started/pyvips/pyvips-2.2.1-py2.py3-none-any.whl --no-index --find-links /kaggle/input/libvips-pyvips-installation-and-getting-started/pyvips\n# !pip install git+https://github.com/ildoonet/pytorch-gradual-warmup-lr.git","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:30:20.823801Z","iopub.execute_input":"2023-12-18T10:30:20.824415Z","iopub.status.idle":"2023-12-18T10:30:20.852334Z","shell.execute_reply.started":"2023-12-18T10:30:20.824383Z","shell.execute_reply":"2023-12-18T10:30:20.851458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, gc, time, copy\nimport h5py\nimport pickle\nos.environ[\"OPENCV_IO_MAX_IMAGE_PIXELS\"] = pow(2,40).__str__()\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\ntqdm.pandas()\nfrom collections import defaultdict\n\nimport math\nimport random\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn import model_selection\nfrom sklearn import metrics\nfrom sklearn import preprocessing\n\nimport tensorflow as tf\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data.sampler import SubsetRandomSampler, RandomSampler, SequentialSampler\nimport torchvision\n# from warmup_scheduler import GradualWarmupScheduler\n# from torchvision.transforms import v2\n\nimport timm\nfrom timm.data import resolve_data_config\nfrom timm.data.transforms_factory import create_transform\n\nimport IPython.display as display\n\nfrom PIL import Image\nimport cv2\n# import pyvips\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:30:20.862067Z","iopub.execute_input":"2023-12-18T10:30:20.862774Z","iopub.status.idle":"2023-12-18T10:30:39.251805Z","shell.execute_reply.started":"2023-12-18T10:30:20.862731Z","shell.execute_reply":"2023-12-18T10:30:39.250433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = dict(\n    seed = 42,\n    folds = 5,\n    img_size = [1024, 1024],\n    learning_rate = 3e-3, # 2e-5, 3e-4\n    eta_min = 1e-5,\n    epochs = 15,\n    batch_size = 8,\n    num_workers = 2,\n    num_tiles = 20,\n    warmup_epoch = 1,\n    warmup_factor = 10,\n)\n\ndef seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    torch.manual_seed(SEED)\n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(SEED)\n        torch.cuda.manual_seed_all(SEED)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n#     os.environ['TF_CUDNN_DETERMINISTIC'] = str(SEED)\n#     tf.random.set_seed(SEED)\n#     keras.utils.set_random_seed(seed=SEED)\n    print('seeding done!!!')\n\ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n    \nseeding(config['seed'])","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:31:51.643346Z","iopub.execute_input":"2023-12-18T10:31:51.644216Z","iopub.status.idle":"2023-12-18T10:31:51.661216Z","shell.execute_reply.started":"2023-12-18T10:31:51.644178Z","shell.execute_reply":"2023-12-18T10:31:51.660111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = Path(\"../input/UBC-OCEAN/\")\nos.listdir(DATA_PATH)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:31:53.505904Z","iopub.execute_input":"2023-12-18T10:31:53.507381Z","iopub.status.idle":"2023-12-18T10:31:53.519537Z","shell.execute_reply.started":"2023-12-18T10:31:53.507322Z","shell.execute_reply":"2023-12-18T10:31:53.518365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(DATA_PATH/'train.csv')\ntest_df = pd.read_csv(DATA_PATH/'test.csv')\nsample_df = pd.read_csv(DATA_PATH/'sample_submission.csv')\n#  '/kaggle/input/ubc-reduced-png-2964x2964'\nget_train_images = lambda x: \"/kaggle/input/ubc-reduced-png-2964x2964/\" + str(x) + \".png\"\nget_test_images = lambda x: \"/kaggle/input/UBC-OCEAN/test_images/\" + str(x) + \".png\"\n\ncheck_path = lambda path: tf.io.gfile.exists(path)\n\ntrain_df['image_path'] = train_df.loc[:, 'image_id'].progress_apply(get_train_images)\ntrain_df['exists'] = train_df.loc[:, 'image_path'].map(check_path)\n\nprint(\"Checking training data ...\")\ndisplay.display(train_df['exists'].value_counts())\ntrain_df = train_df[train_df['exists'] == True]\ntrain_df.reset_index(drop=True, inplace=True)\n\ntest_df['image_path'] = test_df.loc[:, 'image_id'].progress_apply(get_test_images)\ntest_df['exists'] = test_df.loc[:, 'image_path'].map(check_path)\n\nprint(\"Checking test data ...\")\ndisplay.display(test_df['exists'].value_counts())\ntest_df = test_df[test_df['exists'] == True]\ntest_df.reset_index(drop=True, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:32:28.156731Z","iopub.execute_input":"2023-12-18T10:32:28.157192Z","iopub.status.idle":"2023-12-18T10:32:30.283659Z","shell.execute_reply.started":"2023-12-18T10:32:28.157157Z","shell.execute_reply":"2023-12-18T10:32:30.282527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = train_df['label'].unique().tolist()\nid2label = {l:i for i, l in enumerate(labels)}\nlabel2id = {i:l for i, l in enumerate(labels)}\n\ntrain_df['target'] = train_df['label'].map(id2label)\ntrain_df['target'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:32:34.768899Z","iopub.execute_input":"2023-12-18T10:32:34.769375Z","iopub.status.idle":"2023-12-18T10:32:34.790514Z","shell.execute_reply.started":"2023-12-18T10:32:34.769339Z","shell.execute_reply":"2023-12-18T10:32:34.789251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MIN_SIZE = 750 # 500\nN = 6\nM = 2\nMARGIN = 5\nTHRESHOLD_DUP = 0.1\nSIZE = 256\nBIN_THRESH = 240\nSIGMA = 2\nMIN_PROP = 0.02\n\nIMG_SIZE = 1024  # 512\n\ndef get_img_prop(tile, min_saturation=20):\n    hsv = cv2.cvtColor(tile, cv2.COLOR_RGB2HSV)\n    _, s, _ = cv2.split(hsv)\n\n    low_sat = (s < min_saturation).mean()\n\n    counts = np.bincount((tile.mean(-1).flatten().astype(int)))\n    background_prop = counts[np.argsort(counts)[::-1][:3]].sum() / counts.sum()\n\n    background_prop = max(background_prop, low_sat)\n\n    return 1 - background_prop\n\ndef get_grid(orig_size, tile_size, overlap_factor=1):\n    top_x = np.arange(\n        orig_size[0] % tile_size // 2,  # shift to center grid\n        orig_size[0],\n        int(tile_size / overlap_factor),\n    )[:-1]\n    top_y = np.arange(\n        orig_size[1] % tile_size // 2,  # shift to center grid\n        orig_size[1],\n        int(tile_size / overlap_factor),\n    )[:-1]\n    grid = []\n    for x in top_x:\n        right_space = orig_size[0] - (x + tile_size)\n        if right_space > 0:\n            boundaries_x = (x, x + tile_size)\n        else:\n            boundaries_x = (x + right_space, x + right_space + tile_size)\n\n        for y in top_y:\n            down_space = orig_size[1] - (y + tile_size)\n            if down_space > 0:\n                boundaries_y = (y, y + tile_size)\n            else:\n                boundaries_y = (y + down_space, y + down_space + tile_size)\n            grid.append((boundaries_x, boundaries_y))\n\n    return grid\n\ndef remove_chunks(image, size=256, min_prop=0.01):\n    h, w, _ = image.shape\n\n    top_x = np.arange(0, h, size)\n    top_y = np.arange(0, w, size)\n\n    kept = []\n    for x in top_x:\n        prop = get_img_prop(image[x: x + size])\n        if prop > min_prop:\n            kept.append(image[x: x + size])\n\n    if len(kept):\n        image = np.concatenate(kept, 0)\n\n    kept = []\n    for y in top_y:\n        prop = get_img_prop(image[:, y: y + size])\n        if prop > min_prop:\n            kept.append(image[:, y: y + size])\n\n    if len(kept):\n        image = np.concatenate(kept, 1)\n\n    return image\n\n\ndef get_normalization_ratio(img, tile_size=256, min_saturation=20):\n    h, w, _ = img.shape\n    grid = get_grid((h, w), tile_size)\n\n    for x, y in grid:\n        tile = img[x[0]: x[1], y[0]: y[1]]\n        img_prop = get_img_prop(tile, min_saturation=min_saturation)\n\n        if img_prop < 0.25:\n            ratio = np.mean(tile, (0, 1)) / np.array([255.0, 255.0, 255.0])\n            return ratio\n\n    return None\n\n\ndef normalize(image, ratio=None):\n    if ratio is None:\n        return image\n    else:\n        return np.clip((image / ratio), 0, 255).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:32:38.239779Z","iopub.execute_input":"2023-12-18T10:32:38.241692Z","iopub.status.idle":"2023-12-18T10:32:38.269929Z","shell.execute_reply.started":"2023-12-18T10:32:38.241633Z","shell.execute_reply":"2023-12-18T10:32:38.26887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_scale_factor(data, idx):\n    min_width = data.loc[idx, 'image_width'].min()\n    min_height = data.loc[idx, 'image_height'].min()\n    mean_width = data.loc[idx, 'image_width'].mean()\n    mean_height = data.loc[idx, 'image_height'].mean()\n    \n    SF = 2\n    \n    if (min_width and min_height) < 7_000: \n        SF = 4\n    elif (min_width and min_height) >= 7_000 and (min_width and min_height) < 15_000: \n        SF = 8\n    elif (min_width and min_height) >= 15_000 and (min_width and min_height) <= 30_000: \n        SF = 16\n    else: \n        SF = 25\n    \n    return SF\n\n\ndef read_pyvips(image_path: str, scale_factor: int) -> np.array:\n    image = pyvips.Image.new_from_file(image_path, access='sequential')\n    image = image.resize(1.0 / scale_factor).numpy()\n# #     image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n#     gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n#     mask = cv2.compare(gray, 5, cv2.CMP_LT)\n#     image[mask > 0] = 255\n    return image\n","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:32:41.255Z","iopub.execute_input":"2023-12-18T10:32:41.255439Z","iopub.status.idle":"2023-12-18T10:32:41.266019Z","shell.execute_reply.started":"2023-12-18T10:32:41.255405Z","shell.execute_reply":"2023-12-18T10:32:41.264778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.columns","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:16:04.42101Z","iopub.execute_input":"2023-12-18T10:16:04.421994Z","iopub.status.idle":"2023-12-18T10:16:04.429137Z","shell.execute_reply.started":"2023-12-18T10:16:04.421951Z","shell.execute_reply":"2023-12-18T10:16:04.427986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_image(data, idx, tile_size=128, img_size=1024, white_bg=True):\n#     scale_factor = get_scale_factor(data, idx)\n#     is_tma = data.loc[idx, 'is_tma']\n#     print(f\"Image {data.loc[idx, 'image_id']}, is tma: {is_tma}, scale_factor: {scale_factor}\")\n#     image = read_pyvips(data.loc[idx, 'image_path'], scale_factor=scale_factor)\n    image = cv2.imread(data.loc[idx, 'image_path'], cv2.IMREAD_COLOR)\n#     ratio = get_normalization_ratio(image)\n    sz = np.array(image.shape[:2]) \n    sz = (sz * 0.8).astype(int)\n    img = cv2.resize(image, sz[::-1])\n    img = remove_chunks(img, tile_size, min_prop=MIN_PROP)\n    img = cv2.resize(img, (img_size, img_size))\n    \n    if white_bg:\n        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        mask = cv2.compare(gray, 10, cv2.CMP_LT)\n        img[mask > 0] = 255\n        return img\n    gc.collect()\n    \n    return img\n    \nnum_samples = train_df.shape[0]\nimages = [Image.fromarray(read_image(data=train_df, idx=i, white_bg=False)).convert('RGB') for i in tqdm(range(num_samples))]\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:33:42.272406Z","iopub.execute_input":"2023-12-18T10:33:42.272911Z","iopub.status.idle":"2023-12-18T10:44:11.106855Z","shell.execute_reply.started":"2023-12-18T10:33:42.272869Z","shell.execute_reply":"2023-12-18T10:44:11.105686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DUMPED_DATALOADER_PATH = '/kaggle/working/resized_ubc_images.pkl'\n\nwith open(DUMPED_DATALOADER_PATH, 'wb') as file:\n    pickle.dump(images, file)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:50:12.544856Z","iopub.execute_input":"2023-12-18T10:50:12.545344Z","iopub.status.idle":"2023-12-18T10:50:19.016234Z","shell.execute_reply.started":"2023-12-18T10:50:12.545306Z","shell.execute_reply":"2023-12-18T10:50:19.011843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('/kaggle/working/resized_ubc_images.pkl', 'rb') as file:\n     loaded_images = pickle.load(file)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:50:19.020884Z","iopub.execute_input":"2023-12-18T10:50:19.022052Z","iopub.status.idle":"2023-12-18T10:50:38.518686Z","shell.execute_reply.started":"2023-12-18T10:50:19.021954Z","shell.execute_reply":"2023-12-18T10:50:38.514571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(loaded_images), train_df.shape[0]","metadata":{"execution":{"iopub.status.busy":"2023-12-18T10:50:38.524162Z","iopub.execute_input":"2023-12-18T10:50:38.524614Z","iopub.status.idle":"2023-12-18T10:50:38.541242Z","shell.execute_reply.started":"2023-12-18T10:50:38.52458Z","shell.execute_reply":"2023-12-18T10:50:38.539627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def get_kernel(ks=3):\n#     kernels = {\n#         \"rect\": cv2.getStructuringElement(cv2.MORPH_RECT, (ks,ks)),\n#         \"ellipse\": cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (ks, ks)),\n#         \"cross\": cv2.getStructuringElement(cv2.MORPH_CROSS, (ks, ks))\n#     }\n    \n#     return kernels\n\n\n# def get_erosion(image, kernel, iterations=1):\n#     erosion = cv2.erode(image, kernel, iterations=iterations)\n#     return erosion\n\n\n# def get_dilation(image, kernel, iterations=1):\n#     dilation = cv2.dilate(image, kernel, iterations=iterations)\n#     return dilation\n\n# def get_erosion_and_dilation(image, kernel):\n#     opening = cv2.morphologyEx(image, cv2.MORPH_OPEN, kernel)\n#     return opening\n\n# def get_dilation_and_erosion(image, kernel):\n#     closing = cv2.morphologyEx(image, cv2.MORPH_CLOSE, kernel)\n#     return closing\n\n# def get_gradient_morph(image, kernel):\n#     grad = cv2.morphologyEx(image, cv2.MORPH_GRADIENT, kernel)\n#     return grad\n\n# k = get_kernel(ks=5)['cross']\n# # i2 = get_erosion(img, k, iterations=1)\n# # i2 = get_dilation(img, k, iterations=1)\n# # i2 = get_erosion_and_dilation(img, k)\n# # i2 = get_dilation_and_erosion(img, k)\n# i2 = get_gradient_morph(img, k)\n# plt.imshow(i2);","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.232901Z","iopub.status.idle":"2023-12-18T09:56:28.233325Z","shell.execute_reply.started":"2023-12-18T09:56:28.233107Z","shell.execute_reply":"2023-12-18T09:56:28.233127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class UBCDataset(Dataset):\n    \n#     def __init__(self, data, transform):\n#         self.data = data\n#         self.transform = transform\n# #         self.n_images, _, _, _ = data_path['images'].shape\n#         self.n_images = len(data)\n    \n#     def __len__(self):\n#         return self.n_images\n    \n    \n#     def __getitem__(self, idx):\n#         image_path = self.data.loc[idx, 'image_path']\n# #         image = read_image(image_path)\n#         image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n#         image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n# #         k = get_kernel(ks=5)['rect']\n# #         if self.data.loc[idx, 'is_tma'] == False:\n# #             image = get_erosion(image, kernel=k, iterations=2)\n# #             image = get_dilation(image, kernel=k, iterations=1)\n#         image = self.transform(image=image)['image']\n#         image = self.normalize(image)\n#         image = torch.tensor(image).float().permute(2,0,1)\n#         label = self.data.loc[idx, 'target']\n        \n#         return {\"image\": image, \"target\": torch.tensor(label, dtype=torch.long)}\n    \n#     def normalize(self, image):\n#         image = image.astype(np.float32)\n#         return image / 255","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.236194Z","iopub.status.idle":"2023-12-18T09:56:28.23656Z","shell.execute_reply.started":"2023-12-18T09:56:28.236371Z","shell.execute_reply":"2023-12-18T09:56:28.236387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tsfm = A.Compose([\n#     A.ColorJitter(brightness=0.2, contrast=0.5, saturation=0.5, hue=0.5, always_apply=False, p=0.5),\n# #     A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n# #     ToTensorV2(),\n# ])\n\n\n# ds = UBCDataset(train_df, transform=tsfm)\n# dls = DataLoader(ds, batch_size=8, sampler=RandomSampler(ds), num_workers=2, pin_memory=True)\n# b = next(iter(dls))\n# # b['image'].size(), b['target']\n# b.keys()\n\n# # test_df.loc[0, 'image_path']","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.238343Z","iopub.status.idle":"2023-12-18T09:56:28.238684Z","shell.execute_reply.started":"2023-12-18T09:56:28.238523Z","shell.execute_reply":"2023-12-18T09:56:28.238539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataloaders","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:16:25.304136Z","iopub.execute_input":"2023-10-15T06:16:25.304554Z","iopub.status.idle":"2023-10-15T06:16:25.309045Z","shell.execute_reply.started":"2023-10-15T06:16:25.304523Z","shell.execute_reply":"2023-10-15T06:16:25.308233Z"}}},{"cell_type":"code","source":"# def get_transforms(img_size):\n#     train_tsfm = A.Compose(\n#         [\n#             A.VerticalFlip(p=0.5),\n#             A.HorizontalFlip(p=0.5),\n# #             A.Rotate(limit=(-30, 30), p=0.5),\n# #             A.ImageCompression(quality_lower=10, quality_upper=100, p=0.5),\n# #             A.ColorJitter(brightness=0.2, contrast=0.5, saturation=0.5, hue=0.2, p=0.5),\n# #             A.Cutout(max_h_size=int(img_size[0]*0.2), max_w_size=int(img_size[0]*0.2), num_holes=1, p=0.3),\n#             A.Resize(height=img_size[0], width=img_size[1]),\n#         ]\n#     )\n    \n#     valid_tsfm = A.Compose(\n#         [\n# #             A.ColorJitter(brightness=0.2, contrast=0.5, saturation=0.5, hue=0.2, p=0.5),\n#             A.Resize(height=img_size[0], width=img_size[1]),\n# #             ToTensorV2(),\n#         ]\n#     )\n    \n#     return {\"train\": train_tsfm, \"valid\": valid_tsfm}\n\n\n# def get_dataloaders(data, cfg, split='train'):\n#     tsfm = get_transforms(img_size=cfg['img_size'])\n#     idx = data.index.tolist()\n    \n#     if split.lower() == 'train':\n#         sampler = SubsetRandomSampler(idx)\n#         ds = UBCDataset(data, tsfm[split])\n#         dls = DataLoader(ds, batch_size=cfg['batch_size'], shuffle=True, num_workers=cfg['num_workers'], \n#                          pin_memory=True, drop_last=True)\n        \n#     elif split.lower() == 'valid':\n#         ds = UBCDataset(data, tsfm[split])\n#         dls = DataLoader(ds, batch_size=cfg['batch_size']*2, shuffle=False, num_workers=cfg['num_workers'], \n#                          pin_memory=True, drop_last=False)\n        \n#     else:\n#         raise ValueError('Invalid split choose either train or valid')\n#     return dls","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.239883Z","iopub.status.idle":"2023-12-18T09:56:28.240213Z","shell.execute_reply.started":"2023-12-18T09:56:28.240049Z","shell.execute_reply":"2023-12-18T09:56:28.240064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tma = train_df[train_df['is_tma']==True].reset_index(drop=True)\n# wsi = train_df[train_df['is_tma']==False].sample(8).reset_index(drop=True)\n# df = pd.concat([wsi, tma], axis=0).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.241703Z","iopub.status.idle":"2023-12-18T09:56:28.242053Z","shell.execute_reply.started":"2023-12-18T09:56:28.24189Z","shell.execute_reply":"2023-12-18T09:56:28.241907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dl = get_dataloaders(train_df, config, split='train')\n\n# b = next(iter(dl))\n# b_size = b['image'].size()[0]\n# row = 2\n# col = b_size // row\n\n# plt.figure(figsize=(16, 6))\n\n# for i in tqdm(range(b_size)):\n#     image, target = b['image'][i], b['target'][i]\n#     image = image.permute(1,2,0).cpu().numpy()\n    \n#     plt.subplot(row, col, i + 1)\n#     plt.xticks([])\n#     plt.yticks([])\n#     plt.title(target.cpu().numpy(), fontsize=8)\n#     plt.imshow(image);","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.244009Z","iopub.status.idle":"2023-12-18T09:56:28.244346Z","shell.execute_reply.started":"2023-12-18T09:56:28.244179Z","shell.execute_reply":"2023-12-18T09:56:28.244201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Gem layer","metadata":{}},{"cell_type":"code","source":"# def gem(x, p=3, eps=1e-6):\n#     \"\"\"\n#     Apply Generalized Mean Pooling (GeM) to a tensor.\n\n#     Args:\n#         x (torch.Tensor): Input tensor of shape (batch_size, channels, height, width).\n#         p (float): The p-value for the generalized mean. Default is 3.\n#         eps (float): A small constant added to the denominator to prevent division by zero. Default is 1e-6.\n\n#     Returns:\n#         torch.Tensor: GeM-pooled representation of the input tensor.\n#     \"\"\"\n#     return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1.0 / p)\n\n\n# class GeM(nn.Module):\n#     \"\"\"\n#     Generalized Mean Pooling (GeM) layer for global average pooling.\n\n#     Attributes:\n#         p (float or torch.Tensor): The p-value for the generalized mean.\n#         eps (float): A small constant added to the denominator to prevent division by zero.\n#     \"\"\"\n#     def __init__(self, p=3, eps=1e-6, p_trainable=False):\n#         \"\"\"\n#         Initialize the GeM layer.\n\n#         Args:\n#             p (float or torch.Tensor): The p-value for the generalized mean.\n#             eps (float, optional): Eps to prevent division by zero. Defaults to 1e-6.\n#             p_trainable (bool, optional): Whether p is trainable. Defaults to False.\n#         \"\"\"\n#         super(GeM, self).__init__()\n#         if p_trainable:\n#             self.p = Parameter(torch.ones(1) * p)\n#         else:\n#             self.p = p\n#         self.eps = eps\n\n#     def forward(self, x):\n#         \"\"\"\n#         Perform the GeM pooling operation on the input tensor.\n\n#         Args:\n#             x (torch.Tensor): Input tensor of shape (batch_size, channels, height, width).\n\n#         Returns:\n#             torch.Tensor: GeM-pooled representation of the input tensor.\n#         \"\"\"\n#         ret = gem(x, p=self.p, eps=self.eps)\n#         return ret\n\n\n# class Attention(nn.Module):\n#     \"\"\"\n#     Attention module for sequence data.\n\n#     Attributes:\n#         hidden_dim (int): The dimension of the input sequence.\n#         attention_dim (int): The dimension of the attention layer.\n#     \"\"\"\n#     def __init__(self, hidden_dim, attention_dim=None):\n#         \"\"\"\n#         Constructor\n\n#         Args:\n#             hidden_dim (int): The dimension of the input sequence.\n#             attention_dim (int, optional): The dimension of the attention layer.\n#                 Defaults to None, in which case it's set to `hidden_dim`.\n#         \"\"\"\n#         super().__init__()\n\n#         self.hidden_dim = hidden_dim\n#         self.attention_dim = attention_dim\n#         if self.attention_dim is None:\n#             self.attention_dim = self.hidden_dim\n#         # W * x + b\n#         self.proj_w = nn.Linear(self.hidden_dim, self.attention_dim, bias=True)\n#         # v.T\n#         self.proj_v = nn.Linear(self.attention_dim, 1, bias=False)\n\n#     def forward(self, x):\n#         \"\"\"\n#         Perform the forward pass of the attention mechanism.\n\n#         Args:\n#             x (torch.Tensor): Input sequence data of shape (batch_size, seq_len, input_dim).\n\n#         Returns:\n#             torch.Tensor: Attention-weighted representation of the input sequence.\n#         \"\"\"\n#         batch_size, seq_len, _ = x.size()\n#         H = torch.tanh(self.proj_w(x))\n#         att_scores = torch.softmax(self.proj_v(H), axis=1)\n#         attn_x = (x * att_scores).sum(1)\n#         return attn_x","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.245349Z","iopub.status.idle":"2023-12-18T09:56:28.245694Z","shell.execute_reply.started":"2023-12-18T09:56:28.245533Z","shell.execute_reply":"2023-12-18T09:56:28.245549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load pretrained models\n\n`\"timm/efficientnet_b3.ra2_in1k\"` (1024x1024) `lb=0.27 (fold 1)`\n\n`timm/tiny_vit_21m_512.dist_in22k_ft_in1k` (512x512)\n\n`timm/maxvit_tiny_tf_512.in1k` (512x512)\n\n`timm/tf_efficientnetv2_b3.in21k_ft_in1k` (1), `lb=0.29`\n\n`timm/coatnet_2_rw_224.sw_in12k_ft_in1k`\n\n`timm/tf_efficientnetv2_s.in21k_ft_in1k`\n\n`timm/tf_efficientnetv2_m.in21k_ft_in1k`\n\n`timm/convnext_tiny.fb_in22k_ft_in1k`","metadata":{}},{"cell_type":"code","source":"# timm.list_pretrained(\"convnext_pico*\")","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.24704Z","iopub.status.idle":"2023-12-18T09:56:28.24737Z","shell.execute_reply.started":"2023-12-18T09:56:28.24721Z","shell.execute_reply":"2023-12-18T09:56:28.247226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model_config = dict()\n# model_config['num_classes'] = 5\n# model_config['backbone'] = \"timm/efficientnet_b3.ra2_in1k\"\n# model_config['pretrained'] = False\n\n# class UBCModel(nn.Module):\n#     def __init__(self, cfg: dict):\n#         super(UBCModel, self).__init__()\n#         self.cfg = cfg\n#         self.num_classes = cfg['num_classes']\n#         self.backbone = timm.create_model(cfg['backbone'], pretrained=cfg['pretrained'], num_classes=0)\n#         if 'vit' in cfg['backbone'].split('/')[-1]:\n#             self.num_features = self.backbone.num_features\n            \n#         elif 'convnext' in cfg['backbone'].split('/')[-1]:\n#             self.num_features = self.backbone.num_features\n            \n#         elif 'coatnet' in cfg['backbone'].split('/')[-1]:\n#             self.num_features = self.backbone.num_features\n            \n#         elif 'resnet' in cfg['backbone'].split('/')[-1]:\n#             self.num_features = self.backbone.num_features\n            \n#         elif 'efficient' in cfg['backbone'].split('/')[-1]:\n#             self.num_features = self.backbone.num_features\n            \n# #         elif 'efficientnet' in cfg['backbone'].split('/')[-1]:\n# #             self.num_features = self.backbone.classifier.in_features\n#         else:\n#             raise Exception('Please, use the correct backbone')\n            \n#         self.backbone.classifier = nn.Identity()\n#         self.backbone.global_pool = nn.Identity()\n#         self.global_pool = GeM(p_trainable=False)\n#         self.head = nn.Sequential(\n#             nn.AdaptiveAvgPool2d(output_size=1),\n#             nn.Flatten(),\n#             nn.Linear(self.num_features, self.num_classes),\n#         )\n# #         self.loss_fn = torch.nn.CrossEntropyLoss()        \n        \n#     def forward(self, x):\n#         x = self.backbone.forward_features(x)\n#         x = self.global_pool(x)\n#         x = self.head(x)\n#         return x\n    \n#     def freeze_encoder(self, flag):\n#         for param in self.backbone.parameters():\n#             param.requires_grad = not flag\n    \n# net = UBCModel(model_config)\n# net.freeze_encoder(True)\n# net.eval()\n# net(b['image']).softmax(dim=-1)","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.249213Z","iopub.status.idle":"2023-12-18T09:56:28.249573Z","shell.execute_reply.started":"2023-12-18T09:56:28.249382Z","shell.execute_reply":"2023-12-18T09:56:28.249398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Focal loss","metadata":{}},{"cell_type":"code","source":"# class FocalLoss(nn.Module):\n#     def __init__(self, alpha=1, gamma=2, reduction='mean'):\n#         super(FocalLoss, self).__init__()\n#         self.alpha = alpha\n#         self.gamma = gamma\n#         self.reduction = reduction\n\n#     def forward(self, inputs, targets):\n#         ce_loss = nn.CrossEntropyLoss(reduction='none')(inputs, targets)\n#         pt = torch.exp(-ce_loss)\n#         ft = self.alpha * (1 - pt) ** self.gamma\n#         focal_loss = ft * ce_loss\n        \n#         if self.reduction == 'mean':\n#             return torch.mean(focal_loss)\n#         elif self.reduction == 'sum':\n#             return torch.sum(focal_loss)\n#         else:\n#             return focal_loss","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.250909Z","iopub.status.idle":"2023-12-18T09:56:28.25124Z","shell.execute_reply.started":"2023-12-18T09:56:28.251076Z","shell.execute_reply":"2023-12-18T09:56:28.251091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class MetricMonitor:\n#     def __init__(self, float_precision=4):\n#         self.float_precision = float_precision\n#         self.reset()\n\n#     def reset(self):\n#         self.metrics = defaultdict(lambda: {\"val\": 0, \"count\": 0, \"avg\": 0})\n\n#     def update(self, metric_name, val):\n#         metric = self.metrics[metric_name]\n\n#         metric[\"val\"] += val\n#         metric[\"count\"] += 1\n#         metric[\"avg\"] = metric[\"val\"] / metric[\"count\"]\n\n#     def __str__(self):\n#         return \" | \".join(\n#             [\n#                 \"{metric_name}: {avg:.{float_precision}f}\".format(\n#                     metric_name=metric_name, avg=metric[\"avg\"], float_precision=self.float_precision\n#                 )\n#                 for (metric_name, metric) in self.metrics.items()\n#             ]\n#         )\n\n\n# def ACC(y_true, y_preds):\n#     y_true = y_true.detach().cpu().numpy()\n#     y_preds = torch.argmax(y_preds, dim=1)\n#     y_preds = y_preds.detach().cpu().numpy()\n#     return metrics.balanced_accuracy_score(y_true, y_preds)\n\n# def flush():\n#     gc.collect()\n#     torch.cuda.empty_cache()\n#     torch.cuda.reset_peak_memory_stats()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.252242Z","iopub.status.idle":"2023-12-18T09:56:28.252593Z","shell.execute_reply.started":"2023-12-18T09:56:28.252404Z","shell.execute_reply":"2023-12-18T09:56:28.252419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def mixup(inputs, truth, clip=[0, 1]):\n#     indices = torch.randperm(inputs.size(0))\n#     shuffled_input = inputs[indices]\n#     shuffled_labels = truth[indices]\n    \n#     lam = np.random.uniform(clip[0], clip[1])\n#     inputs = inputs * lam + shuffled_input * (1 - lam)\n#     return inputs, truth, shuffled_labels, lam\n\n\n# def train_with_mixup(model, optimizer, criterion, data_loader, scaler, device='cpu', epoch=1, n_iters=1000):\n#     model.train()\n#     train_loss = 0\n#     correct = 0\n#     n_total = 0\n#     example_ct = 0\n#     step_ct = 0\n#     metric_monitor = MetricMonitor()\n#     stream = tqdm(data_loader)\n#     for i, batch in enumerate(stream, start=1):\n#         xb, yb = batch['image'], batch['target']\n#         xb, yb = xb.to(device, non_blocking=True), yb.to(device, non_blocking=True)\n        \n#         do_mixup = False\n#         if random.random() < 0.5:\n#             do_mixup = True\n#             xb, yb, yb_mix, lam = mixup(xb, yb)\n        \n#         with torch.autocast(device_type='cuda', dtype=torch.float16):\n#             outputs = model(xb)\n#             outputs = outputs.softmax(dim=-1)\n#             loss = criterion(outputs, yb)\n#             if do_mixup:\n#                 loss11 = criterion(outputs, yb_mix)\n#                 loss = loss * lam + loss11 * (1 - lam)\n            \n#         train_loss += loss.detach().float()\n#         scaler.scale(loss).backward()\n        \n#         acc = ACC(yb, outputs)\n#         metric_monitor.update('Loss', loss.item())\n#         metric_monitor.update('Balanced Accuracy', acc)\n        \n#         example_ct += len(xb)\n#         METRICS = {\n#             \"train/train_loss\": train_loss,\n#             \"train/epoch\": (i + 1 + (n_iters * epoch)) / n_iters,\n#             \"train/example_ct\": example_ct,\n#             \"train/balanced_acc\": acc,\n#         }\n        \n# #         if (i + 1) < n_iters:\n# #             # log train metrics to wandb\n# #             wandb.log(METRICS)\n            \n#         step_ct += 1\n        \n# #         if (i+1) % n_iters == 0:\n# #             scaler.step(optimizer)\n# #             scaler.update()\n# #             optimizer.zero_grad(set_to_none=True)\n            \n#         scaler.step(optimizer)\n#         scaler.update()\n#         optimizer.zero_grad(set_to_none=True)\n            \n#         stream.set_description(\n#         \"Epoch: {epoch}. Train.      {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor))\n    \n#     train_loss_total = (train_loss / len(data_loader)).item()\n#     flush()\n    \n#     return train_loss_total, METRICS\n\n\n# def train_one_loop(model, optimizer, criterion, data_loader, scaler, device='cpu', epoch=1, n_iters=1000):\n#     model.train()\n#     train_loss = 0\n#     correct = 0\n#     n_total = 0\n#     example_ct = 0\n#     step_ct = 0\n#     metric_monitor = MetricMonitor()\n#     stream = tqdm(data_loader)\n#     for i, batch in enumerate(stream, start=1):\n#         xb, yb = batch['image'], batch['target']\n#         xb, yb = xb.to(device, non_blocking=True), yb.to(device, non_blocking=True)\n        \n#         with torch.autocast(device_type='cuda', dtype=torch.float16):\n#             outputs = model(xb)\n#             outputs = outputs.softmax(-1)\n#             loss = criterion(outputs, yb)\n            \n#         train_loss += loss.detach().float()\n#         scaler.scale(loss).backward()\n        \n#         acc = ACC(yb, outputs)\n#         metric_monitor.update('Loss', loss.item())\n#         metric_monitor.update('Balanced Accuracy', acc)\n        \n#         example_ct += len(xb)\n#         METRICS = {\n#             \"train/train_loss\": train_loss,\n#             \"train/epoch\": (i + 1 + (n_iters * epoch)) / n_iters,\n#             \"train/example_ct\": example_ct,\n#             \"train/balanced_acc\": acc,\n#         }\n        \n# #         if (i + 1) < n_iters:\n# #             # log train metrics to wandb\n# #             wandb.log(METRICS)\n            \n#         step_ct += 1\n        \n# #         if (i+1) % n_iters == 0:\n# #             scaler.step(optimizer)\n# #             scaler.update()\n# #             optimizer.zero_grad(set_to_none=True)\n            \n#         scaler.step(optimizer)\n#         scaler.update()\n#         optimizer.zero_grad(set_to_none=True)\n            \n#         stream.set_description(\n#         \"Epoch: {epoch}. Train.      {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor))\n    \n#     train_loss_total = (train_loss / len(data_loader)).item()\n#     flush()\n    \n#     return train_loss_total, METRICS\n        \n    \n# def valid_one_loop(model, criterion, data_loader, device='cpu', epoch=1):\n#     model.eval()\n#     valid_loss = 0\n#     correct = 0\n#     n_total = 0\n#     metric_monitor = MetricMonitor()\n#     stream = tqdm(data_loader)\n#     for i, batch in enumerate(stream, start=1):\n#         xb, yb = batch['image'], batch['target']\n#         xb, yb = xb.to(device, non_blocking=True), yb.to(device, non_blocking=True)\n        \n#         with torch.autocast(device_type='cuda', dtype=torch.float16):\n#             with torch.no_grad():\n#                 outputs = model(xb)\n#                 outputs = outputs.softmax(dim=-1)\n# #                 PREDS.append(torch.argmax(outputs, dim=1).detach().cpu().numpy())\n#             loss = criterion(outputs, yb)\n            \n#         valid_loss += loss.detach().float()\n#         acc = ACC(yb, outputs)\n#         metric_monitor.update('Loss', loss.item())\n#         metric_monitor.update('Balanced Accuracy', acc)\n#         stream.set_description(\n#         \"Epoch: {epoch}. Valid.      {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor))\n#         val_metrics = {\n#             \"valid/val_loss\": valid_loss,\n#             \"valid/balanced_acc\": acc,\n#         }\n# #         wandb.log(val_metrics)\n    \n#     valid_loss_total = (valid_loss / len(data_loader)).item()\n#     flush()\n#     return valid_loss_total, val_metrics","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.25412Z","iopub.status.idle":"2023-12-18T09:56:28.254455Z","shell.execute_reply.started":"2023-12-18T09:56:28.254292Z","shell.execute_reply":"2023-12-18T09:56:28.254308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# name = model_config['backbone'].split('/')[-1]\n\n# def train(model, optimizer, criterion, train_loader, valid_loader, epochs, scheduler, lr_reduce, \n#           scaler, device='cpu', device_ids=[0,1], fold=0, n_iters=100):\n    \n#     model = nn.DataParallel(model, device_ids=device_ids)\n#     model.to(device)\n#     best_metric = np.inf\n#     loss_min = np.inf\n    \n#     for epoch in tqdm(range(1, epochs+1)):\n        \n#         train_loss, train_metrics = train_one_loop(model, optimizer, criterion, train_loader, \n#                                     scaler=scaler, device=device, epoch=epoch, n_iters=n_iters)\n        \n# #         train_loss, train_metrics = train_with_mixup(model, optimizer, criterion, train_loader, \n# #                                     scaler=scaler, device=device, epoch=epoch, n_iters=n_iters)\n        \n#         valid_loss, valid_metrics = valid_one_loop(model, criterion, valid_loader, device=device, epoch=epoch)\n#         scheduler.step(epoch-1)\n#         lr_reduce.step(valid_loss)\n        \n#         train_metrics[\"train/rmse\"] = train_loss\n#         valid_metrics[\"valid/rmse\"] = valid_loss\n# #         wandb.log({**train_metrics, **valid_metrics})\n        \n#         metric = valid_loss\n#         if metric < best_metric:\n#             print(f\"Best metric: ({best_metric:.6f} --> {metric:.6f}). Saving model ...\")\n#             torch.save(model.module.state_dict(), f\"{name}_fold_{fold}.pth\")\n#             best_metric = metric\n            \n# def train_swa(model, optimizer, criterion, train_loader, valid_loader, epochs, scheduler, lr_reduce, \n#           scaler, device='cpu', device_ids=[0, 1], fold=0, n_iters=100):\n    \n#     model = nn.DataParallel(model, device_ids=device_ids)\n#     model.to(device)\n#     swa_model = torch.optim.swa_utils.AveragedModel(model)\n#     swa_optimizer = torch.optim.swa_utils.SWALR(optimizer, swa_lr=0.05)\n#     swa_start = 1\n#     best_metric = np.inf\n#     loss_min = np.inf\n#     for epoch in tqdm(range(1, epochs+1)):\n#         train_loss, train_metrics = train_one_loop(model, optimizer, criterion, train_loader, \n#                                     scaler=scaler, device=device, epoch=epoch, n_iters=n_iters)\n#         valid_loss, valid_metrics = valid_one_loop(model, criterion, valid_loader, device=device, epoch=epoch)\n#         if epoch > swa_start: \n#             swa_model.update_parameters(model)\n#             swa_optimizer.step()\n#         else:\n#             scheduler.step(epoch-1)\n#             lr_reduce.step(valid_loss)\n        \n#         metric = valid_loss\n#         if metric < best_metric:\n#             print(f\"Best metric: ({best_metric:.6f} --> {metric:.6f}). Saving model ...\")\n#             torch.save(model.state_dict(), f\"{name}_fold_{fold}.pth\")\n#             best_metric = metric\n#     batch = next(iter(valid_loader))\n#     torch.optim.swa_utils.update_bn(batch['image'], swa_model)\n\n# def predict(model, data_loader, loss_config, device='cpu'):\n#     model.to(device)\n#     model.eval()\n#     preds = np.empty((0, model.num_classes))\n    \n#     for i, batch in enumerate(tqdm(data_loader, total=len(data_loader))):\n#         xb = batch['image']\n#         xb = xb.to(device, non_blocking=True)\n        \n#         with torch.autocast(device_type='cuda', dtype=torch.float16):\n#             with torch.inference_mode():\n#                 outputs = model(xb)\n                \n#                 # get probabilities\n#                 if loss_config['activation'] == 'sigmoid':\n#                     outputs = ouputs.sigmoid()\n#                 elif loss_config['activation'] == 'softmax':\n#                     outputs = outputs.softmax(-1)\n#                 preds = np.concatenate([preds, outputs.detach().cpu().numpy()])\n                \n#     return preds","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.255734Z","iopub.status.idle":"2023-12-18T09:56:28.25606Z","shell.execute_reply.started":"2023-12-18T09:56:28.255901Z","shell.execute_reply":"2023-12-18T09:56:28.255916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loss_config = dict()\n# loss_config['activation'] = 'softmax'","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.25745Z","iopub.status.idle":"2023-12-18T09:56:28.257801Z","shell.execute_reply.started":"2023-12-18T09:56:28.257642Z","shell.execute_reply":"2023-12-18T09:56:28.257658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n# print(f\"Using {device} ...\")\n\n# train_df = train_df.sample(frac=1)\n\n# gc.collect()\n\n# kfold = model_selection.StratifiedKFold(n_splits=config['folds'], shuffle=True, random_state=config['seed'])\n\n# x = train_df.index\n# y = train_df['target'].values\n# OOF_PREDS = np.zeros(len(train_df))\n\n# for fold, (tr_idx, val_idx) in enumerate(kfold.split(x, y)):\n# #     run = wandb.init(\n# #         project=\"cgair-pytorch-baseline\"\n# #     )\n# #     artifact = wandb.Artifact(f'fold_{fold}_weights', type='model')\n    \n#     if fold > 1:\n#         break\n    \n#     if fold in [0,1,4]:\n#         epochs = 10+2*fold\n#     else:\n#         epochs = config['epochs']\n#     print(f\"\\n===> Fold {fold} ...\")\n    \n#     train_ds = train_df.iloc[tr_idx]\n#     valid_ds = train_df.iloc[val_idx]\n#     train_ds = train_ds.reset_index(drop=True)\n#     valid_ds = valid_ds.reset_index(drop=True)\n    \n#     train_loader = get_dataloaders(train_ds, config, split='train')\n#     valid_loader = get_dataloaders(valid_ds, config, split='valid')\n    \n#     model = UBCModel(model_config)\n#     model.freeze_encoder(True)\n# #     optimizer = torch.optim.AdamW(model.parameters(), lr=config['learning_rate'])\n#     optimizer = torch.optim.Adam(model.parameters(), lr=config['learning_rate'])\n# #     optimizer = torch.optim.SGD(model.parameters(), lr=config['learning_rate'], momentum=0.9, weight_decay=1e-5)\n#     criterion = torch.nn.CrossEntropyLoss().to(device)\n# #     criterion = FocalLoss().to(device)\n#     scaler = torch.cuda.amp.GradScaler()\n# #     scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.9)\n#     lr_reduce = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=2, verbose=True)\n#     scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, config['epochs']-config['warmup_factor'])\n# #     scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, config['epochs'], eta_min=config['eta_min'])\n#     scheduler = GradualWarmupScheduler(optimizer, multiplier=config['warmup_factor'], total_epoch=config['warmup_factor'],\n#                                       after_scheduler=scheduler_cosine)\n#     num_training_steps = math.ceil(len(train_loader)/config['batch_size'])\n# #     wandb.config = config\n\n#     train(model, optimizer, criterion, train_loader, valid_loader, epochs=config['epochs'], scheduler=scheduler, \n#         lr_reduce=lr_reduce, scaler=scaler, device=device, fold=fold, n_iters=num_training_steps)\n    \n# #     train_swa(model, optimizer, criterion, train_loader, valid_loader, epochs=config['epochs'], scheduler=scheduler, \n# #         lr_reduce=lr_reduce, scaler=scaler, device=device, fold=fold, n_iters=num_training_steps)\n    \n#     PREDS = predict(model, valid_loader, loss_config, device=device)\n#     print(f\"\\nFold {fold} performance: {metrics.balanced_accuracy_score(valid_ds['target'].values, np.argmax(PREDS, axis=1)):.4f}\\n\")\n# #     OOF_PREDS[val_idx] = np.argmax(PREDS, axis=1)\n\n#     del model, train_loader, valid_loader, train_ds, valid_ds\n# #     artifact.add_file(f\"{name}_fold_{fold}.pth\")\n# #     run.log_artifact(artifact)\n# # wandb.finish()\n#     flush()","metadata":{"execution":{"iopub.status.busy":"2023-12-18T09:56:28.26068Z","iopub.status.idle":"2023-12-18T09:56:28.261021Z","shell.execute_reply.started":"2023-12-18T09:56:28.26086Z","shell.execute_reply":"2023-12-18T09:56:28.260876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Current performance:\n\n## `Distortion+cutout+coarsedropout-fold2 = 0.27, lr=3e-4, 2-layer-bi-lstm` \n\n# Augmentation performance: \n\n## Fold 3: `Random Brightness`\n\n`loss: 0.971777, bal_acc: 0.9111`\n\n## Fold 3: `Optical distortion`\n\n`loss: 0.980489, bal_acc: 0.9`\n\n## Fold 3: `GeometricTransformation`\n\n`loss: 1.04539, bal_acc: 0.75`\n\n## Fold 3: `Cutout`\n\n`loss: 1.0824141, bal_acc: 0.7694`\n\n## Fold 3: `CoarseDropout`\n\n`loss: 1.090829, bal_acc: 0.7103`\n\n## Fold 3: `Shitfscale`\n\n`loss: 1.294702, bal_acc: 0.5610`","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}