{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":8713513,"sourceType":"datasetVersion","datasetId":5227521}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import time\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\nfrom transformers import get_cosine_schedule_with_warmup\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:00.141215Z","iopub.execute_input":"2024-10-03T13:11:00.141745Z","iopub.status.idle":"2024-10-03T13:11:02.822073Z","shell.execute_reply.started":"2024-10-03T13:11:00.141700Z","shell.execute_reply":"2024-10-03T13:11:02.821291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_interactive():\n    return 'runtime' in get_ipython().config.IPKernelApp.connection_file\n\nprint('Interactive?', is_interactive())","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.823521Z","iopub.execute_input":"2024-10-03T13:11:02.823942Z","iopub.status.idle":"2024-10-03T13:11:02.829401Z","shell.execute_reply.started":"2024-10-03T13:11:02.823909Z","shell.execute_reply":"2024-10-03T13:11:02.828295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    GRAD_ACC = 2\n    if is_interactive():\n        EPOCHS = 5\n    else:\n        EPOCHS = 20\n    TARGET_BATCH_SIZE = 32\n    BATCH_SIZE = TARGET_BATCH_SIZE // GRAD_ACC\n    IMAGE_SIZE_TRAIN = 640\n    IMAGE_SIZE_TEST = 640\n    S_FOLD = [0]\n    N_MODEL = 3\n    MODEL = \"hrnet_w18\"\n    INTERPOLATION = cv2.INTER_CUBIC\n    DROPOUT = 0.4\n    WD = 2e-2\n    PRE_LR = 1e-4\n    SEED = 2007\n    AUG_PROB = 0.75\n    N_FOLD = 5\n    IN_CHANNELS = 30\n    MAX_GRAD_NORMS = 1e9\n    IS_PROCESS = False\n    EARLY_STOPPING = 5\n    USE_AMP = True\n    AUG = True\n    N_CLASSES = 75","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.830546Z","iopub.execute_input":"2024-10-03T13:11:02.830844Z","iopub.status.idle":"2024-10-03T13:11:02.850902Z","shell.execute_reply.started":"2024-10-03T13:11:02.830811Z","shell.execute_reply":"2024-10-03T13:11:02.850197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, random\nimport numpy as np\ndef set_seeds(seed):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)  # type: ignore\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = False\nset_seeds(seed=CFG.SEED)\nrng = np.random.default_rng(seed=CFG.SEED)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.852894Z","iopub.execute_input":"2024-10-03T13:11:02.853199Z","iopub.status.idle":"2024-10-03T13:11:02.876333Z","shell.execute_reply.started":"2024-10-03T13:11:02.853167Z","shell.execute_reply":"2024-10-03T13:11:02.875601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.877381Z","iopub.execute_input":"2024-10-03T13:11:02.877682Z","iopub.status.idle":"2024-10-03T13:11:02.886363Z","shell.execute_reply.started":"2024-10-03T13:11:02.877651Z","shell.execute_reply":"2024-10-03T13:11:02.885662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\").fillna(-100)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.887541Z","iopub.execute_input":"2024-10-03T13:11:02.887823Z","iopub.status.idle":"2024-10-03T13:11:02.922652Z","shell.execute_reply.started":"2024-10-03T13:11:02.887793Z","shell.execute_reply":"2024-10-03T13:11:02.921985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\n# value = np.array(\n#     sorted(glob.glob(\"/kaggle/input/rsna-image-preprocessing1/*.npz\"), key=(lambda x: int(os.path.basename(x).replace(\".npz\", \"\")))) +\n#     sorted(glob.glob(\"/kaggle/input/rsna-image-processing/*.npz\"), key=(lambda x: int(os.path.basename(x).replace(\".npz\", \"\"))))\n# )\n# value[0:5], value[1785:1790]","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.923762Z","iopub.execute_input":"2024-10-03T13:11:02.924143Z","iopub.status.idle":"2024-10-03T13:11:02.928524Z","shell.execute_reply.started":"2024-10-03T13:11:02.924101Z","shell.execute_reply":"2024-10-03T13:11:02.927604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target = target.replace({\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2})\n# target.loc[:, \"study_id\"] = value","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.929666Z","iopub.execute_input":"2024-10-03T13:11:02.929968Z","iopub.status.idle":"2024-10-03T13:11:02.973312Z","shell.execute_reply.started":"2024-10-03T13:11:02.929936Z","shell.execute_reply":"2024-10-03T13:11:02.972443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nkf = KFold(n_splits=CFG.N_FOLD, shuffle=True, random_state=CFG.SEED)\nkf.get_n_splits(target)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:02.974447Z","iopub.execute_input":"2024-10-03T13:11:02.974774Z","iopub.status.idle":"2024-10-03T13:11:03.438231Z","shell.execute_reply.started":"2024-10-03T13:11:02.974741Z","shell.execute_reply":"2024-10-03T13:11:03.437272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nimport albumentations.pytorch as Ap","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:03.442713Z","iopub.execute_input":"2024-10-03T13:11:03.443169Z","iopub.status.idle":"2024-10-03T13:11:04.097530Z","shell.execute_reply.started":"2024-10-03T13:11:03.443135Z","shell.execute_reply":"2024-10-03T13:11:04.096526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_train = A.Compose([\n    A.Resize(CFG.IMAGE_SIZE_TRAIN, CFG.IMAGE_SIZE_TRAIN, interpolation=CFG.INTERPOLATION),\n    \n    A.OneOf([\n        A.Sharpen(alpha=(0.0, 0.5), lightness=(0.0, 1.5), always_apply=True, p=1.0),\n        A.Posterize(num_bits=(3, 8), always_apply=True, p=1.0),\n        A.RandomBrightnessContrast(brightness_limit=(-0.3, 0.3), contrast_limit=(-0.3, 0.3), always_apply=True, p=1.0),\n    ], p=CFG.AUG_PROB),\n    \n    A.OneOf([\n        A.MotionBlur(blur_limit=5, always_apply=True, p=1.0),\n        A.MedianBlur(blur_limit=5, always_apply=True, p=1.0),\n        A.GaussianBlur(blur_limit=5, always_apply=True, p=1.0),\n    ], p=CFG.AUG_PROB),\n    \n    A.OneOf([\n        A.OpticalDistortion(distort_limit=(-1.0, 1.0), always_apply=True, p=1.0),\n        A.GridDistortion(num_steps=5, distort_limit=(-1.0, 1.0), always_apply=True, p=1.0),\n        A.ElasticTransform(alpha=12, sigma=3, always_apply=True, p=1.0),\n    ], p=CFG.AUG_PROB),\n    \n    A.CoarseDropout(num_holes_range=(1, 16), hole_height_range=(10, 40), hole_width_range=(10, 40), p=CFG.AUG_PROB),\n    \n    A.Affine(rotate=(-30, 30), shear=(0, 0.1), translate_percent = (0, 0.1), p=CFG.AUG_PROB),\n    \n    A.Normalize(mean=0.5, std=0.5),\n    Ap.ToTensorV2(),\n])\ntransforms_test = A.Compose([\n    A.Resize(CFG.IMAGE_SIZE_TEST, CFG.IMAGE_SIZE_TEST, interpolation=CFG.INTERPOLATION),\n    A.Normalize(mean=0.5, std=0.5),\n    Ap.ToTensorV2(),\n])","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:04.098805Z","iopub.execute_input":"2024-10-03T13:11:04.099386Z","iopub.status.idle":"2024-10-03T13:11:04.117477Z","shell.execute_reply.started":"2024-10-03T13:11:04.099341Z","shell.execute_reply":"2024-10-03T13:11:04.116527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip --quiet install ultralytics","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:04.118553Z","iopub.execute_input":"2024-10-03T13:11:04.118865Z","iopub.status.idle":"2024-10-03T13:11:15.720190Z","shell.execute_reply.started":"2024-10-03T13:11:04.118833Z","shell.execute_reply":"2024-10-03T13:11:15.718922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ultralytics import YOLO","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:15.721763Z","iopub.execute_input":"2024-10-03T13:11:15.722128Z","iopub.status.idle":"2024-10-03T13:11:16.164107Z","shell.execute_reply.started":"2024-10-03T13:11:15.722090Z","shell.execute_reply":"2024-10-03T13:11:16.163307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:07.072108Z","iopub.execute_input":"2024-10-03T13:12:07.073023Z","iopub.status.idle":"2024-10-03T13:12:07.081063Z","shell.execute_reply.started":"2024-10-03T13:12:07.072982Z","shell.execute_reply":"2024-10-03T13:12:07.080143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nclass RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=CFG.IN_CHANNELS, n_model=CFG.N_MODEL, n_classes=75, pretrained=False, features_only=False):\n        super().__init__()\n        self.model = YOLO(model='yolo11s-cls.yaml', verbose=False)\n        self.model.model.model[0].conv = nn.Conv2d(in_c, self.model.model.model[0].conv.out_channels, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n        self.model.model.model[-1].pool = GeM(eps=1e-5)\n        self.model.model.model[-1].linear = nn.Linear(1280, 75)\n        self.model = self.model.model\n    def forward(self, x):\n        return self.model(x)\n    def count_params(self):\n        return sum(p.numel() for p in self.model.parameters() if p.requires_grad)","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:07.316857Z","iopub.execute_input":"2024-10-03T13:12:07.317188Z","iopub.status.idle":"2024-10-03T13:12:08.588572Z","shell.execute_reply.started":"2024-10-03T13:12:07.317154Z","shell.execute_reply":"2024-10-03T13:12:08.587781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input = torch.rand(1, CFG.IN_CHANNELS, CFG.IMAGE_SIZE_TRAIN, CFG.IMAGE_SIZE_TRAIN)\nmodel = RSNA24Model(CFG.MODEL).to('cpu')\nf'Model Output Size: {model(input).shape}, Params(M): {model.count_params() / 1e6: .2f}'","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:11.306421Z","iopub.execute_input":"2024-10-03T13:12:11.306800Z","iopub.status.idle":"2024-10-03T13:12:11.847693Z","shell.execute_reply.started":"2024-10-03T13:12:11.306764Z","shell.execute_reply":"2024-10-03T13:12:11.846776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del input, model","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:14.110102Z","iopub.execute_input":"2024-10-03T13:12:14.110489Z","iopub.status.idle":"2024-10-03T13:12:14.119526Z","shell.execute_reply.started":"2024-10-03T13:12:14.110451Z","shell.execute_reply":"2024-10-03T13:12:14.118721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:14.307766Z","iopub.execute_input":"2024-10-03T13:12:14.308088Z","iopub.status.idle":"2024-10-03T13:12:14.370388Z","shell.execute_reply.started":"2024-10-03T13:12:14.308024Z","shell.execute_reply":"2024-10-03T13:12:14.369306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def custom_metrics(y_pred, y, device):\n    loss = 0\n    for i in range(25):\n        loss += F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], weight=torch.tensor([1.0, 2.0, 4.0]).to(y_pred.dtype).to(device))/25\n#     y_pred = y_pred.to(torch.float32)\n#     slices = [slice(0, 5), slice(5, 15), slice(15, 25)] \n#     w = 2 ** y  \n#     loss = None\n#     for i in range(25):\n#         if loss is None:\n#             loss = F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], reduction='none').reshape(-1, 1)\n#         else:\n#             loss = torch.cat((loss, F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], reduction='none').reshape(-1, 1)), dim=1)\n#     wloss_sums = []\n#     for k, idx in enumerate(slices):\n#         wloss_sums.append((w[:, idx] * loss[:, idx]).sum())\n\n#     y_spinal_prob = y_pred.reshape(-1, 25, 3)[:, :5, :].softmax(dim=2)   \n#     w_max = torch.amax(w[:, :5], dim=1)                        \n#     y_max = torch.amax(y[:, :5] == 2, dim=1).to(y_pred.dtype)   \n#     y_pred_max = y_spinal_prob[:, :, 2].amax(dim=1)\n\n#     loss_max = F.binary_cross_entropy(y_pred_max, y_max, reduction='none')\n#     wloss_sums.append((w_max * loss_max).sum())\n\n#     loss = (wloss_sums[0] / 6.084050632911392 +\n#             wloss_sums[1] / 12.962531645569621 + \n#             wloss_sums[2] / 14.38632911392405 +\n#             wloss_sums[3] / 1.729113924050633) / (4 * y.size(0))\n\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:14.866297Z","iopub.execute_input":"2024-10-03T13:12:14.866615Z","iopub.status.idle":"2024-10-03T13:12:14.873360Z","shell.execute_reply.started":"2024-10-03T13:12:14.866581Z","shell.execute_reply":"2024-10-03T13:12:14.872445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluation_metrics(y_pred, y, device):\n    loss = 0\n    for i in range(25):\n        loss += F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], weight=torch.tensor([1.0, 2.0, 4.0]).to(y_pred.dtype).to(device))/25\n#     y_pred = y_pred.to(torch.float32)\n#     loss = None\n#     slices = [slice(0, 5), slice(5, 15), slice(15, 25)] \n#     w = 2 ** y  \n#     loss4_sum  = torch.zeros(4)\n#     w_sum = torch.zeros(4)\n#     for i in range(25):\n#         if loss is None:\n#             loss = F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], reduction='none').reshape(-1, 1)\n#         else:\n#             loss = torch.cat((loss, F.cross_entropy(y_pred[:, i*3:i*3+3], y[:, i], reduction='none').reshape(-1, 1)), dim=1)\n            \n#     for k, idx in enumerate(slices):\n#         w_sum[k] += w[:, idx].sum()\n#         loss4_sum[k] += (w[:, idx] * loss[:, idx]).sum()\n\n#     y_spinal_prob = y_pred.reshape(-1, 25, 3)[:, :5, :].softmax(dim=2)          \n#     w_max = torch.amax(w[:, :5], dim=1)                        \n#     y_max = torch.amax(y[:, :5] == 2, dim=1).to(y_pred.dtype)   \n#     y_pred_max = y_spinal_prob[:, :, 2].amax(dim=1)          \n\n#     loss_max = F.binary_cross_entropy(y_pred_max, y_max, reduction='none')\n#     loss4_sum[3] += (w_max * loss_max).sum()\n#     w_sum[3] += w_max.sum()\n\n#     score = (loss4_sum / w_sum).sum() / 4\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:15.069517Z","iopub.execute_input":"2024-10-03T13:12:15.069835Z","iopub.status.idle":"2024-10-03T13:12:15.077527Z","shell.execute_reply.started":"2024-10-03T13:12:15.069802Z","shell.execute_reply":"2024-10-03T13:12:15.076571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = torch.rand(2, 75).to(torch.half)\nb = torch.randint(0, 3, (2, 25))\ncustom_metrics(a, b, 'cpu'), evaluation_metrics(a, b, 'cpu')","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:15.316055Z","iopub.execute_input":"2024-10-03T13:12:15.316351Z","iopub.status.idle":"2024-10-03T13:12:15.332474Z","shell.execute_reply.started":"2024-10-03T13:12:15.316319Z","shell.execute_reply":"2024-10-03T13:12:15.331653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate(model, loader_test):\n    print(f\"-----> Start evaluating epochs {epoch+1}\")\n    loss = 0\n    out = None\n    target = None\n    if is_interactive():\n        with tqdm(loader_test, leave=True) as pbar:\n            for i, (a, b) in enumerate(pbar):\n                a = a.to(device)\n                b = b.to(device)\n                with torch.no_grad():\n                    with autocast:\n                        if out is None:\n                            out = model(a)\n                            target = b\n                        else:\n                            out = torch.cat((out, model(a)))\n                            target = torch.cat((target, b))\n                        del a, b\n    else:\n        for i, (a, b) in enumerate(loader_test):\n            a = a.to(device)\n            b = b.to(device)\n            with torch.no_grad():\n                with autocast:                       \n                    if out is None:\n                        out = model(a)\n                        target = b\n                    else:\n                        out = torch.cat((out, model(a)))\n                        target = torch.cat((target, b))\n                    del a, b\n    loss = evaluation_metrics(out.detach().cpu(), target.detach().cpu(), 'cpu')\n    del out, target\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:15.810399Z","iopub.execute_input":"2024-10-03T13:12:15.810708Z","iopub.status.idle":"2024-10-03T13:12:15.820269Z","shell.execute_reply.started":"2024-10-03T13:12:15.810674Z","shell.execute_reply":"2024-10-03T13:12:15.819477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nclass RSNA24Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        x = np.zeros((512, 512, 30)).astype(np.uint8)\n        t = self.df.iloc[idx]\n        st_id = int(t['study_id'])\n        label = t[1:].values.astype(np.int64)\n        \n        # Sagittal T1\n        for i in range(0, 10, 1):\n            try:\n                p = f'/kaggle/input/lsdcgcs/cvt_png/{st_id}/Sagittal T1/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T1')\n                pass\n            \n        # Sagittal T2/STIR\n        for i in range(0, 10, 1):\n            try:\n                p = f'/kaggle/input/lsdcgcs/cvt_png/{st_id}/Sagittal T2_STIR/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+10] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass\n            \n        # Axial T2\n        axt2 = glob.glob(f'/kaggle/input/lsdcgcs/cvt_png/{st_id}/Axial T2/*.png')\n        axt2 = sorted(axt2)\n    \n        step = len(axt2) / 10.0\n        st = len(axt2)/2.0 - 4.0*step\n        end = len(axt2)+0.0001\n                \n        for i, j in enumerate(np.arange(st, end, step)):\n            try:\n                p = axt2[max(0, int((j-0.5001).round()))]\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+20] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass  \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n        return x, label\n# class RSNA24Dataset(torch.utils.data.Dataset):\n#     def __init__(self, df, transform=None):\n#         self.df = df\n#         self.transform = transform    \n#     def __len__(self):\n#         return len(self.df)\n\n#     def __getitem__(self, idx):\n#         t = self.df.iloc[idx]\n#         st_id = t['study_id']\n#         label = torch.from_numpy(t[1:].values.astype(np.int64))\n#         x = np.load(st_id)[\"arr_0\"].astype(np.uint8)\n#         if self.transform is not None:\n#             x = self.transform(image=x)['image']\n#         return x, label","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:16.245986Z","iopub.execute_input":"2024-10-03T13:12:16.246280Z","iopub.status.idle":"2024-10-03T13:12:16.261769Z","shell.execute_reply.started":"2024-10-03T13:12:16.246250Z","shell.execute_reply":"2024-10-03T13:12:16.260921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = RSNA24Dataset(\n    target,\n    transforms_train)\nloader_train = torch.utils.data.DataLoader(\n        ds_train, \n        batch_size = 1, \n        shuffle = False, \n        drop_last = True, \n        num_workers=2, \n        pin_memory=True)\nimport matplotlib.pyplot as plt\nfor i, (a, b) in enumerate(loader_train):\n    y = ((a.numpy().transpose(0, 2, 3, 1)[0,..., 0:3] + 1)/2)\n    print(y.max(), y.min(), y.mean(), y.std(), y.shape)\n    plt.imshow(y)\n    plt.show()\n    y = ((a.numpy().transpose(0, 2, 3, 1)[0,..., 17:20] + 1)/2)\n    print(y.max(), y.min(), y.mean(), y.std(), y.shape)\n    plt.imshow(y)\n    plt.show()\n    y = ((a.numpy().transpose(0, 2, 3, 1)[0,..., 27:30] + 1)/2)\n    print(y.max(), y.min(), y.mean(), y.std(), y.shape)\n    plt.imshow(y)\n    plt.show()\n    break\nplt.close()\ndel y, a, b, loader_train, ds_train","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:16.559671Z","iopub.execute_input":"2024-10-03T13:12:16.559947Z","iopub.status.idle":"2024-10-03T13:12:19.276390Z","shell.execute_reply.started":"2024-10-03T13:12:16.559918Z","shell.execute_reply":"2024-10-03T13:12:19.275374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:19.278635Z","iopub.execute_input":"2024-10-03T13:12:19.278953Z","iopub.status.idle":"2024-10-03T13:12:19.283639Z","shell.execute_reply.started":"2024-10-03T13:12:19.278920Z","shell.execute_reply":"2024-10-03T13:12:19.282517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport matplotlib.pyplot as plt\nfrom collections import OrderedDict\nautocast = torch.autocast(device_type=\"cuda\", enabled=CFG.USE_AMP, dtype=torch.half)\nscaler = torch.GradScaler(device='cuda', enabled=CFG.USE_AMP, init_scale=4096)\nfor fold, (a, b) in enumerate(kf.split(target)):\n    if fold not in CFG.S_FOLD:\n        continue\n    else:\n        print(f\"--> Start fold: {fold}\")\n    # Data\n    ds_train = RSNA24Dataset(\n        target.loc[a, :],\n        transforms_train,\n    )\n    loader_train = torch.utils.data.DataLoader(\n        ds_train, \n        batch_size = CFG.BATCH_SIZE, \n        shuffle = True, \n        drop_last = True,\n        num_workers=os.cpu_count(),\n        prefetch_factor=1,\n        pin_memory=True)\n    \n    ds_test =   RSNA24Dataset(\n        target.loc[b, :],\n        transforms_test, )\n    loader_test = torch.utils.data.DataLoader(\n        ds_test, \n        batch_size = CFG.BATCH_SIZE, \n        shuffle = False, \n        drop_last = False, \n        prefetch_factor=1,\n        num_workers=os.cpu_count(), \n        pin_memory=True)\n\n\n    # Model\n    model = RSNA24Model(CFG.MODEL) # Just use pretrained=True with internet\n    \n    if torch.cuda.device_count() > 1:\n        print(\"Let's use\", torch.cuda.device_count(), \"GPUs!\")\n        model = torch.nn.DataParallel(model)\n    model.to(device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.PRE_LR,\n                                  weight_decay=CFG.WD)\n    \n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps=len(loader_train) // CFG.GRAD_ACC,\n        num_training_steps=CFG.EPOCHS * len(loader_train) // CFG.GRAD_ACC,\n        num_cycles=0.475\n    )\n\n    weight=torch.tensor([1.0, 2.0, 4.0]).to(device)\n    criterion = torch.nn.CrossEntropyLoss(weight=weight)\n    best_loss = None\n    val_loss = []\n    total_loss = []\n    for epoch in range(CFG.EPOCHS):\n        print(f\"-----> Start epoch: [{epoch+1} / {CFG.EPOCHS}]\")\n        loss_sum = 0\n        count = 0\n        model.train()\n        if is_interactive():\n            optimizer.zero_grad()\n            with tqdm(loader_train, leave=True) as pbar:\n                for idx, (x, y) in enumerate(pbar):\n                    x = x.to(device)  \n                    y = y.to(device)  \n                    \n                    with autocast:\n                        out = model(x) \n                    loss = custom_metrics(out, y, device)\n                    del x, y, out\n                        \n                    if CFG.GRAD_ACC > 1:\n                        loss /= CFG.GRAD_ACC\n                        \n                    scaler.scale(loss).backward()\n                    count += 1\n                    loss_sum +=  loss.item()\n                    \n                    pbar.set_postfix(\n                        OrderedDict(\n                            current_loss=f'{(loss*CFG.GRAD_ACC).item():.6f}',\n                            total_loss=f'{loss_sum*CFG.GRAD_ACC/count:.6f}',\n                            lr=f'{optimizer.param_groups[0][\"lr\"]:.3e}'\n                        )\n                    )\n                    \n                    del loss\n                    if CFG.MAX_GRAD_NORMS is not None:\n                        nn.utils.clip_grad_norm_(model.parameters(), CFG.MAX_GRAD_NORMS)\n                        \n                    if (idx + CFG.GRAD_ACC - 1) % CFG.GRAD_ACC == 0 or idx == len(loader_train) - 1:\n                        scaler.step(optimizer)\n                        scaler.update()\n                        scheduler.step() \n                        optimizer.zero_grad()\n        else:\n            for idx, (x, y) in enumerate(loader_train):\n                x = x.to(device)  \n                y = y.to(device)  \n                \n                with autocast:\n                    out = model(x)\n                loss = custom_metrics(out, y, device)\n\n                del x, y, out\n                \n                if CFG.GRAD_ACC > 1:\n                    loss /= CFG.GRAD_ACC\n                    \n                scaler.scale(loss).backward()\n                count += 1\n                loss_sum += loss.item()\n                del loss\n                \n                if CFG.MAX_GRAD_NORMS is not None:\n                    nn.utils.clip_grad_norm_(model.parameters(), CFG.MAX_GRAD_NORMS)\n                    \n                if (idx + CFG.GRAD_ACC - 1) % CFG.GRAD_ACC == 0 or idx == len(loader_train) - 1:\n                    scaler.step(optimizer)\n                    scaler.update()\n                    scheduler.step()  \n                    optimizer.zero_grad()\n                    \n        model.eval()\n\n        val = evaluate(model, loader_test)\n        lr = optimizer.param_groups[0]['lr']\n        print(f\"rsna_metrics: {loss_sum*CFG.GRAD_ACC/count:.6f}, val_rsna_metrics: {val.item():.6f}, lr: {lr:.3e}\")\n        print()\n        if best_loss is not None:\n            if val < best_loss:\n                early_stopping = 0\n                print(f\"best loss improve from {best_loss:.6f} to {val:.6f}. Saving models\")\n                best_loss = val\n                torch.save(model.state_dict(), f\"RSNA_model_{fold}.pth\")\n            else:\n                print(f\"best loss is not a improvement from {best_loss:.6f}, current val_loss: {val:.6f}. Not Saving models\")\n                early_stopping += 1\n                if early_stopping == CFG.EARLY_STOPPING:\n                    print(f\"early stopping since best_loss: {best_loss:.6f} has not improve for {CFG.EARLY_STOPPING} round\")\n                    break\n        else:\n            early_stopping = 0\n            print(f\"best loss improve from -inf to {val:.6f}. Saving models\")\n            best_loss = val\n            torch.save(model.state_dict(), f\"RSNA_model_{fold}.pth\")\n        del val","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:12:19.285253Z","iopub.execute_input":"2024-10-03T13:12:19.285547Z","iopub.status.idle":"2024-10-03T13:21:37.929754Z","shell.execute_reply.started":"2024-10-03T13:12:19.285516Z","shell.execute_reply":"2024-10-03T13:21:37.928193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\n!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-10-03T13:11:16.199554Z","iopub.status.idle":"2024-10-03T13:11:16.199886Z","shell.execute_reply.started":"2024-10-03T13:11:16.199717Z","shell.execute_reply":"2024-10-03T13:11:16.199734Z"},"trusted":true},"execution_count":null,"outputs":[]}]}