{"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-04T14:53:44.662640Z","iopub.execute_input":"2024-10-04T14:53:44.663044Z","iopub.status.idle":"2024-10-04T14:53:47.425422Z","shell.execute_reply.started":"2024-10-04T14:53:44.662992Z","shell.execute_reply":"2024-10-04T14:53:47.424632Z"},"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-04T14:53:47.427409Z","iopub.execute_input":"2024-10-04T14:53:47.428278Z","iopub.status.idle":"2024-10-04T14:53:47.434295Z","shell.execute_reply.started":"2024-10-04T14:53:47.428226Z","shell.execute_reply":"2024-10-04T14:53:47.432976Z"},"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 = \"densenet121\"\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-04T14:53:47.436250Z","iopub.execute_input":"2024-10-04T14:53:47.436671Z","iopub.status.idle":"2024-10-04T14:53:47.444381Z","shell.execute_reply.started":"2024-10-04T14:53:47.436583Z","shell.execute_reply":"2024-10-04T14:53:47.443460Z"},"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-04T14:53:47.445532Z","iopub.execute_input":"2024-10-04T14:53:47.445856Z","iopub.status.idle":"2024-10-04T14:53:47.457751Z","shell.execute_reply.started":"2024-10-04T14:53:47.445824Z","shell.execute_reply":"2024-10-04T14:53:47.456765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2024-10-04T14:53:47.460453Z","iopub.execute_input":"2024-10-04T14:53:47.460773Z","iopub.status.idle":"2024-10-04T14:53:47.465266Z","shell.execute_reply.started":"2024-10-04T14:53:47.460740Z","shell.execute_reply":"2024-10-04T14:53:47.464342Z"},"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-04T14:53:47.466472Z","iopub.execute_input":"2024-10-04T14:53:47.466817Z","iopub.status.idle":"2024-10-04T14:53:47.500989Z","shell.execute_reply.started":"2024-10-04T14:53:47.466784Z","shell.execute_reply":"2024-10-04T14:53:47.499819Z"},"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-04T14:53:47.502483Z","iopub.execute_input":"2024-10-04T14:53:47.502994Z","iopub.status.idle":"2024-10-04T14:53:47.508376Z","shell.execute_reply.started":"2024-10-04T14:53:47.502886Z","shell.execute_reply":"2024-10-04T14:53:47.507251Z"},"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-04T14:53:47.509761Z","iopub.execute_input":"2024-10-04T14:53:47.510125Z","iopub.status.idle":"2024-10-04T14:53:47.554911Z","shell.execute_reply.started":"2024-10-04T14:53:47.510089Z","shell.execute_reply":"2024-10-04T14:53:47.553851Z"},"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-04T14:53:47.556713Z","iopub.execute_input":"2024-10-04T14:53:47.557162Z","iopub.status.idle":"2024-10-04T14:53:48.287881Z","shell.execute_reply.started":"2024-10-04T14:53:47.557110Z","shell.execute_reply":"2024-10-04T14:53:48.286728Z"},"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-04T14:53:48.289303Z","iopub.execute_input":"2024-10-04T14:53:48.289898Z","iopub.status.idle":"2024-10-04T14:53:48.888249Z","shell.execute_reply.started":"2024-10-04T14:53:48.289850Z","shell.execute_reply":"2024-10-04T14:53:48.887189Z"},"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    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=(1, 8), always_apply=True, p=1.0),\n        A.RandomBrightnessContrast(brightness_limit=(-0.4, 0.4), contrast_limit=(-0.4, 0.4), always_apply=True, p=1.0),\n    ], p=CFG.AUG_PROB),\n    \n    A.OneOf([\n        A.MotionBlur(blur_limit=7, always_apply=True, p=1.0),\n        A.MedianBlur(blur_limit=7, always_apply=True, p=1.0),\n        A.GaussianBlur(blur_limit=7, 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=(-90, 90), shear=(0, 0.4), translate_percent = (0, 0.4), 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-04T14:53:48.889709Z","iopub.execute_input":"2024-10-04T14:53:48.890358Z","iopub.status.idle":"2024-10-04T14:53:48.909848Z","shell.execute_reply.started":"2024-10-04T14:53:48.890303Z","shell.execute_reply":"2024-10-04T14:53:48.908472Z"},"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 = timm.create_model(\n            model_name, \n            pretrained=False, \n            num_classes=75, \n            in_chans=30,\n            proj_drop_rate=0.3,\n        )\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-04T14:53:48.911386Z","iopub.execute_input":"2024-10-04T14:53:48.911815Z","iopub.status.idle":"2024-10-04T14:53:50.648051Z","shell.execute_reply.started":"2024-10-04T14:53:48.911769Z","shell.execute_reply":"2024-10-04T14:53:50.646951Z"},"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)\nf'Model Output Size: {model(input).shape}, Params(M): {model.count_params() / 1e6: .2f}'","metadata":{"execution":{"iopub.status.busy":"2024-10-04T14:53:50.649508Z","iopub.execute_input":"2024-10-04T14:53:50.649935Z","iopub.status.idle":"2024-10-04T14:53:52.241624Z","shell.execute_reply.started":"2024-10-04T14:53:50.649885Z","shell.execute_reply":"2024-10-04T14:53:52.239733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del input, model","metadata":{"execution":{"iopub.status.busy":"2024-10-04T14:53:52.245864Z","iopub.execute_input":"2024-10-04T14:53:52.246785Z","iopub.status.idle":"2024-10-04T14:53:52.260274Z","shell.execute_reply.started":"2024-10-04T14:53:52.246738Z","shell.execute_reply":"2024-10-04T14:53:52.259128Z"},"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-04T14:53:52.261967Z","iopub.execute_input":"2024-10-04T14:53:52.263059Z","iopub.status.idle":"2024-10-04T14:53:52.330004Z","shell.execute_reply.started":"2024-10-04T14:53:52.263010Z","shell.execute_reply":"2024-10-04T14:53:52.328672Z"},"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-04T14:53:52.331494Z","iopub.execute_input":"2024-10-04T14:53:52.332428Z","iopub.status.idle":"2024-10-04T14:53:52.340535Z","shell.execute_reply.started":"2024-10-04T14:53:52.332382Z","shell.execute_reply":"2024-10-04T14:53:52.339480Z"},"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-04T14:53:52.341882Z","iopub.execute_input":"2024-10-04T14:53:52.342760Z","iopub.status.idle":"2024-10-04T14:53:52.350815Z","shell.execute_reply.started":"2024-10-04T14:53:52.342717Z","shell.execute_reply":"2024-10-04T14:53:52.349791Z"},"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-04T14:53:52.352066Z","iopub.execute_input":"2024-10-04T14:53:52.352434Z","iopub.status.idle":"2024-10-04T14:53:52.373973Z","shell.execute_reply.started":"2024-10-04T14:53:52.352381Z","shell.execute_reply":"2024-10-04T14:53:52.372983Z"},"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-04T14:53:52.375496Z","iopub.execute_input":"2024-10-04T14:53:52.375845Z","iopub.status.idle":"2024-10-04T14:53:52.391805Z","shell.execute_reply.started":"2024-10-04T14:53:52.375807Z","shell.execute_reply":"2024-10-04T14:53:52.390661Z"},"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-04T14:53:52.393397Z","iopub.execute_input":"2024-10-04T14:53:52.394120Z","iopub.status.idle":"2024-10-04T14:53:52.411948Z","shell.execute_reply.started":"2024-10-04T14:53:52.394067Z","shell.execute_reply":"2024-10-04T14:53:52.410709Z"},"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-04T14:53:52.413065Z","iopub.execute_input":"2024-10-04T14:53:52.413441Z","iopub.status.idle":"2024-10-04T14:53:55.871204Z","shell.execute_reply.started":"2024-10-04T14:53:52.413398Z","shell.execute_reply":"2024-10-04T14:53:55.870049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-10-04T14:53:55.872611Z","iopub.execute_input":"2024-10-04T14:53:55.873043Z","iopub.status.idle":"2024-10-04T14:53:55.878141Z","shell.execute_reply.started":"2024-10-04T14:53:55.872984Z","shell.execute_reply":"2024-10-04T14:53:55.876962Z"},"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-04T14:53:55.879975Z","iopub.execute_input":"2024-10-04T14:53:55.880432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\n!nvidia-smi","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}