{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9029634,"sourceType":"datasetVersion","datasetId":5442257}],"dockerImageVersionId":30839,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"80529c8e","cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nimport random\nimport os, sys\nimport timm\nimport torch.nn.functional as F\nfrom glob import glob\nfrom PIL import Image\nimport cv2\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport albumentations as A\nfrom torch.optim import AdamW\nfrom transformers import get_cosine_schedule_with_warmup\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nimport math\nfrom pathlib import Path\nfrom collections import OrderedDict\nfrom transformers.models.distilbert.modeling_distilbert import Transformer as T\nfrom torch.utils.tensorboard import SummaryWriter\nwriter = SummaryWriter()\ntorch.multiprocessing.set_sharing_strategy('file_descriptor')","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:13.513365Z","iopub.execute_input":"2025-01-23T18:26:13.513648Z","iopub.status.idle":"2025-01-23T18:26:48.204327Z","shell.execute_reply.started":"2025-01-23T18:26:13.513621Z","shell.execute_reply":"2025-01-23T18:26:48.203169Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"c7ab32a3","cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\nOUTPUT_DIR = '/kaggle/working/rsna-results-2.5d'\n\nif not Path(OUTPUT_DIR).exists():\n    os.mkdir(OUTPUT_DIR)\n    \nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nIMG_SIZE = [256, 256]\nN_FOLDS = 5\nEPOCHS = 100\nUSE_AMP = True\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\nAUG_PROB = 0.75\nSELECTED_FOLDS = [0, 1, 2, 3, 4]\nSEED = 69\nGRAD_ACC = 1\nTGT_BATCH_SIZE = 8\nIN_CHANS = 18\nBATCH_SIZE = TGT_BATCH_SIZE // GRAD_ACC // 2\nMAX_GRAD_NORM = None\nEARLY_STOPPING_EPOCH = 20\nLR = 2e-4 * TGT_BATCH_SIZE / 32\nWD = 1e-2\nAUG = True\nMODEL_NAME = 'convnext_pico.d1_in1k'\n# MODEL_NAME = 'edgenext_base.in21k_ft_in1k'\n\n# MODEL_NAME = 'convnextv2_pico.fcmae'\nNOT_DEBUG = True\nN_WORKERS = 4","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:51.872291Z","iopub.execute_input":"2025-01-23T18:26:51.872638Z","iopub.status.idle":"2025-01-23T18:26:51.883949Z","shell.execute_reply.started":"2025-01-23T18:26:51.872610Z","shell.execute_reply":"2025-01-23T18:26:51.882625Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"278a9af2","cell_type":"code","source":"os.makedirs(OUTPUT_DIR, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:52.081703Z","iopub.execute_input":"2025-01-23T18:26:52.082049Z","iopub.status.idle":"2025-01-23T18:26:52.086982Z","shell.execute_reply.started":"2025-01-23T18:26:52.082023Z","shell.execute_reply":"2025-01-23T18:26:52.085992Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"b88b8672","cell_type":"code","source":"def set_random_seed(seed: int = 2222, deterministic: bool = False):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = deterministic\n\nset_random_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:52.596120Z","iopub.execute_input":"2025-01-23T18:26:52.596457Z","iopub.status.idle":"2025-01-23T18:26:52.608605Z","shell.execute_reply.started":"2025-01-23T18:26:52.596431Z","shell.execute_reply":"2025-01-23T18:26:52.607538Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"61925079","cell_type":"code","source":"df = pd.read_csv(f'{rd}/train.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:52.921063Z","iopub.execute_input":"2025-01-23T18:26:52.921390Z","iopub.status.idle":"2025-01-23T18:26:52.998001Z","shell.execute_reply.started":"2025-01-23T18:26:52.921365Z","shell.execute_reply":"2025-01-23T18:26:52.996804Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f7b704f3","cell_type":"code","source":"df = df.fillna(-100)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:54.807462Z","iopub.execute_input":"2025-01-23T18:26:54.807831Z","iopub.status.idle":"2025-01-23T18:26:54.818211Z","shell.execute_reply.started":"2025-01-23T18:26:54.807804Z","shell.execute_reply":"2025-01-23T18:26:54.817111Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"c6f8a96a","cell_type":"code","source":"label2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2}\ndf = df.replace(label2id)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:56.389598Z","iopub.execute_input":"2025-01-23T18:26:56.389994Z","iopub.status.idle":"2025-01-23T18:26:56.435462Z","shell.execute_reply.started":"2025-01-23T18:26:56.389964Z","shell.execute_reply":"2025-01-23T18:26:56.434335Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"f129645b","cell_type":"code","source":"CONDITIONS = [\n    'Spinal Canal Stenosis', \n    'Left Neural Foraminal Narrowing', \n    'Right Neural Foraminal Narrowing',\n    'Left Subarticular Stenosis',\n    'Right Subarticular Stenosis'\n]\n\nLEVELS = [\n    'L1/L2',\n    'L2/L3',\n    'L3/L4',\n    'L4/L5',\n    'L5/S1',\n]\nmodel_names = list(df.columns)[1:]\nmodel_names","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:57.457689Z","iopub.execute_input":"2025-01-23T18:26:57.458034Z","iopub.status.idle":"2025-01-23T18:26:57.465963Z","shell.execute_reply.started":"2025-01-23T18:26:57.458005Z","shell.execute_reply":"2025-01-23T18:26:57.464617Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"a4696df6","cell_type":"code","source":"from pathlib import Path\nclass RSNA24Dataset(Dataset):\n    def __init__(self, df, phase='train', transform=None):\n        self.df = df\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], IN_CHANS, 3), dtype=np.float32)\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\n        sat1 = glob(f'/kaggle/input/cvt_png/{st_id}/Sagittal T1/*.png')\n        sat1 = sorted(sat1)\n    \n        step = len(sat1) / (IN_CHANS - 1)\n        st = 0\n        end = len(sat1)+0.0001\n        if len(sat1) != 0:\n            for i, j in enumerate(np.arange(st, end, step)):\n                try:\n                    p = sat1[max(0, int((j-0.5001).round()))]\n                    img = Image.open(p).convert('L')\n                    img = np.array(img)\n                    x[..., i, 0] = img.astype(np.float32)\n                except:\n#                     print(f'failed to load on {st_id}, Sagittal T1')\n                    pass\n            \n        #Sagittal T2/STIR\n        sat2 = glob(f'/kaggle/input/cvt_png/{st_id}/Sagittal T2_STIR/*.png')\n        sat2 = sorted(sat2)\n    \n        step = len(sat2) / (IN_CHANS - 1)\n        st = 0\n        end = len(sat2)+0.0001\n\n        if len(sat2) != 0:\n            for i, j in enumerate(np.arange(st, end, step)):\n                try:\n                    p = sat2[max(0, int((j-0.5001).round()))]\n                    img = Image.open(p).convert('L')\n                    img = np.array(img)\n                    x[..., i, 1] = img.astype(np.float32)\n                except:\n#                     print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                    pass\n            \n        # Axial T2\n        axt2 = glob(f'/kaggle/input/cvt_png/{st_id}/Axial T2/*.png')\n        axt2 = sorted(axt2)\n    \n        step = len(axt2) / (IN_CHANS - 1)\n        st = 0\n        end = len(axt2)+0.0001\n\n        if len(axt2) != 0:\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, 2] = img.astype(np.float32)\n                except:\n#                     print(f'failed to load on {st_id}, Axial T2')\n                    pass  \n            \n#         assert np.sum(x)>0\n        if self.transform is not None:\n            for i in range(x.shape[-1]):\n                x[..., i] = self.transform(image=x[..., i])['image']\n\n        x = x.transpose(2, 3, 0, 1)\n        \n                \n        return x, label","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:58.215218Z","iopub.execute_input":"2025-01-23T18:26:58.215601Z","iopub.status.idle":"2025-01-23T18:26:58.229339Z","shell.execute_reply.started":"2025-01-23T18:26:58.215570Z","shell.execute_reply":"2025-01-23T18:26:58.228205Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"9cc44227","cell_type":"code","source":"transforms_train = A.Compose([\n    A.OneOf([\n        A.MotionBlur(blur_limit=5),\n        A.MedianBlur(blur_limit=5),\n        A.GaussianBlur(blur_limit=5),\n        A.GaussNoise(var_limit=50),\n    ], p=AUG_PROB),\n\n    A.OneOf([\n        A.OpticalDistortion(distort_limit=1.0),\n        A.GridDistortion(num_steps=5, distort_limit=1.),\n        A.ElasticTransform(alpha=3),\n    ], p=AUG_PROB),\n\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=AUG_PROB),\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n#     A.CoarseDropout(max_holes=16, max_height=16, max_width=16, min_holes=1, min_height=2, min_width=2, p=AUG_PROB),    \n    A.Normalize(mean=0.5, std=0.5)\n])\n\ntransforms_val = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])\n\nif not AUG:\n    transforms_train = transforms_val","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:26:59.681293Z","iopub.execute_input":"2025-01-23T18:26:59.681684Z","iopub.status.idle":"2025-01-23T18:26:59.695853Z","shell.execute_reply.started":"2025-01-23T18:26:59.681652Z","shell.execute_reply":"2025-01-23T18:26:59.694523Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"a0ec4b66","cell_type":"code","source":"tmp_ds = RSNA24Dataset(df, phase='train', transform=transforms_train)\ntmp_dl = DataLoader(\n            tmp_ds,\n            batch_size=1,\n            shuffle=False,\n            pin_memory=False,\n            drop_last=False,\n            num_workers=0\n            )\n\nfor i, (x, t) in enumerate(tmp_dl):\n    if i==2:break\n    print('x stat:', x.shape, x.min(), x.max(),x.mean(), x.std())\n    print(t, t.shape)\n    y = x.numpy()[0,1,0,:,:]\n    plt.imshow(y)\n    plt.show()\n    print('y stat:', y.shape, y.min(), y.max(),y.mean(), y.std())\n    print()\nplt.close()\ndel tmp_ds, tmp_dl","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:02.569696Z","iopub.execute_input":"2025-01-23T18:27:02.570107Z","iopub.status.idle":"2025-01-23T18:27:03.970101Z","shell.execute_reply.started":"2025-01-23T18:27:02.570075Z","shell.execute_reply":"2025-01-23T18:27:03.969139Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"504e8f03","cell_type":"code","source":"class Attention(nn.Module):\n    def __init__(self, feature_dim, step_dim, bias=True, **kwargs):\n        super(Attention, self).__init__(**kwargs)\n        \n        self.supports_masking = True\n\n        self.bias = bias\n        self.feature_dim = feature_dim\n        self.step_dim = step_dim\n        self.features_dim = 0\n        \n        weight = torch.zeros(feature_dim, 1)\n#         nn.init.kaiming_uniform_(weight)\n        self.weight = nn.Parameter(weight)\n        \n        if bias:\n            self.b = nn.Parameter(torch.zeros(step_dim))\n        \n    def forward(self, x, mask=None):\n        feature_dim = self.feature_dim \n        step_dim = self.step_dim\n\n        eij = torch.mm(\n            x.contiguous().view(-1, feature_dim), \n            self.weight\n        ).view(-1, step_dim)\n        \n        if self.bias:\n            eij = eij + self.b\n            \n        eij = torch.tanh(eij)\n        a = torch.exp(eij)\n        \n        if mask is not None:\n            a = a * mask\n\n        a = a / (torch.sum(a, 1, keepdim=True) + 1e-10)\n\n        weighted_input = x * torch.unsqueeze(a, -1)\n        return torch.sum(weighted_input, 1)\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:04.658776Z","iopub.execute_input":"2025-01-23T18:27:04.659175Z","iopub.status.idle":"2025-01-23T18:27:04.667688Z","shell.execute_reply.started":"2025-01-23T18:27:04.659143Z","shell.execute_reply":"2025-01-23T18:27:04.666225Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d22cb30c","cell_type":"code","source":"class TimmModelCombo(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super(TimmModelCombo, self).__init__()\n\n        self.encoder_sagittal = timm.create_model(\n            backbone,\n            in_chans=2,\n            num_classes=1,\n            features_only=False,\n            drop_rate=0.4,\n            pretrained=pretrained\n        )\n        \n        self.encoder_axial = timm.create_model(\n            backbone,\n            in_chans=1,\n            num_classes=1,\n            features_only=False,\n            drop_rate=0.4,\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder_sagittal.conv_head.out_channels\n            self.encoder_sagittal.classifier = nn.Identity()\n            self.encoder_axial.classifier = nn.Identity()\n            \n        elif 'convnext' in backbone:\n            hdim = self.encoder_sagittal.head.fc.in_features\n            self.encoder_sagittal.head.fc = nn.Identity()\n            self.encoder_axial.head.fc = nn.Identity()\n            \n        if 'densenet121' in backbone:\n            hdim = 1024\n            self.encoder_sagittal.classifier = nn.Identity()\n            self.encoder_axial.classifier = nn.Identity()\n            \n        if 'densenet161' in backbone:\n            hdim = 2208\n            self.encoder_sagittal.classifier = nn.Identity()\n            self.encoder_axial.classifier = nn.Identity()\n            \n        if 'densenet201' in backbone:\n            hdim = 1920\n            self.encoder_sagittal.classifier = nn.Identity()\n            self.encoder_axial.classifier = nn.Identity()\n\n\n#         self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=0., bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(256, 128),\n            nn.Dropout(0.4),\n            nn.LeakyReLU(0.1),\n            nn.Linear(128, 75),\n        )\n        self.attention_layer_sagittal = Attention(512, IN_CHANS)\n        self.attention_layer_axial = Attention(512, IN_CHANS)\n        \n        self.fc_axial = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.Dropout(0.3),\n            nn.SiLU()\n        )\n        self.fc_sagittal = nn.Sequential(\n            nn.Linear(512, 256), \n            nn.Dropout(0.3),\n            nn.SiLU()\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        img_size = x.shape[3]\n        \n        x_sagittal = x[:, :, :2, :, :]\n        x_axial = x[:, :, 2:3, :, :]\n        \n        x_sagittal = x_sagittal.view(bs * IN_CHANS, x_sagittal.shape[2], img_size, img_size)\n        feat_sagittal = self.encoder_sagittal(x_sagittal)\n        feat_sagittal = feat_sagittal.view(bs, IN_CHANS, -1)\n        \n        x_axial = x_axial.view(bs * IN_CHANS, x_axial.shape[2], img_size, img_size)\n        feat_axial = self.encoder_axial(x_axial)\n        feat_axial = feat_axial.view(bs, IN_CHANS, -1)\n#         feat_lstm, _ = self.lstm(feat)\n#         feat_lstm = feat_lstm.contiguous().view(bs * 12, -1)\n#         feat_lstm = self.head(feat_lstm)\n#         feat_lstm = feat_lstm.view(bs, 12, 75).contiguous()\n        atten_sagittal = self.attention_layer_sagittal(feat_sagittal)\n        atten_axial = self.attention_layer_axial(feat_axial)\n#         atten = torch.cat((atten_sagittal, atten_axial), dim=1)\n        atten_sagittal = self.fc_sagittal(atten_sagittal)\n        atten_axial = self.fc_axial(atten_axial)\n        atten = (atten_sagittal + atten_axial) / 2\n        out = self.head(atten)\n        return out","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:06.368002Z","iopub.execute_input":"2025-01-23T18:27:06.368380Z","iopub.status.idle":"2025-01-23T18:27:06.382384Z","shell.execute_reply.started":"2025-01-23T18:27:06.368339Z","shell.execute_reply":"2025-01-23T18:27:06.380946Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"0ec737ae","cell_type":"code","source":"\nclass TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super(TimmModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=2,\n            num_classes=1,\n            features_only=False,\n            drop_rate=0.4,\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder.conv_head.out_channels\n            self.encoder.classifier = nn.Identity()\n        elif 'convnext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n            \n        if 'densenet121' in backbone:\n            hdim = 1024\n            self.encoder.classifier = nn.Identity()\n            \n        if 'densenet161' in backbone:\n            hdim = 2208\n            self.encoder.classifier = nn.Identity()\n        if 'densenet201' in backbone:\n            hdim = 1920\n            self.encoder.classifier = nn.Identity()\n            \n        if 'edgenext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n\n        self.lstm = nn.LSTM(hdim, 256, num_layers=1, dropout=0., bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(hdim, 256),\n            nn.Dropout(0.4),\n            nn.SiLU(),\n            nn.Linear(256, 75),\n        )\n        self.attention_layer = Attention(hdim, IN_CHANS)\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        x = x[:, :, 0:2, :, :]\n        bs = x.shape[0]\n        img_size = x.shape[3]\n        x = x.view(bs * IN_CHANS, 2, img_size, img_size)\n\n        feat = self.encoder(x)\n        feat = feat.view(bs, IN_CHANS, -1)\n        \n        \n#         feat_lstm, _ = self.lstm(feat)\n        \n#         feat_lstm = feat_lstm.contiguous().view(bs * 12, -1)\n#         feat_lstm = self.head(feat_lstm)\n#         feat_lstm = feat_lstm.view(bs, 12, 75).contiguous()\n        atten = self.attention_layer(feat)\n        \n        out = self.head(atten)\n        return out\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T18:27:08.237347Z","iopub.execute_input":"2025-01-23T18:27:08.237753Z","iopub.status.idle":"2025-01-23T18:27:08.247663Z","shell.execute_reply.started":"2025-01-23T18:27:08.237723Z","shell.execute_reply":"2025-01-23T18:27:08.246450Z"}},"outputs":[],"execution_count":null},{"id":"b445fcea","cell_type":"code","source":"m = TimmModel(MODEL_NAME)\nm = m.to(DEVICE)\ni = torch.randn(8, IN_CHANS, 3, 224, 224).to(DEVICE)\nwith torch.no_grad():\n    out = m(i)\nfor o in out:\n    print(o.shape, o.min(), o.max())","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:11.030843Z","iopub.execute_input":"2025-01-23T18:27:11.031181Z","iopub.status.idle":"2025-01-23T18:27:20.277342Z","shell.execute_reply.started":"2025-01-23T18:27:11.031157Z","shell.execute_reply":"2025-01-23T18:27:20.276237Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"83c877af","cell_type":"code","source":"del m, i, out\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:22.699280Z","iopub.execute_input":"2025-01-23T18:27:22.699691Z","iopub.status.idle":"2025-01-23T18:27:22.721760Z","shell.execute_reply.started":"2025-01-23T18:27:22.699649Z","shell.execute_reply":"2025-01-23T18:27:22.720392Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"cdf77829","cell_type":"code","source":"# m = RSNA24Model('efficientnet_b0', in_c=1, n_classes=512, pretrained=False)\n# m = m.to(DEVICE)\n# i = torch.randn(2, IN_CHANS // 3, 256, 256).to(DEVICE)\n# out = m(i)\n# for o in out:\n#     print(o.shape, o.min(), o.max())","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:23.754290Z","iopub.execute_input":"2025-01-23T18:27:23.754667Z","iopub.status.idle":"2025-01-23T18:27:23.758865Z","shell.execute_reply.started":"2025-01-23T18:27:23.754635Z","shell.execute_reply":"2025-01-23T18:27:23.757636Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"16e54fc8","cell_type":"code","source":"# del m, i, out","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:24.804251Z","iopub.execute_input":"2025-01-23T18:27:24.804633Z","iopub.status.idle":"2025-01-23T18:27:24.808983Z","shell.execute_reply.started":"2025-01-23T18:27:24.804601Z","shell.execute_reply":"2025-01-23T18:27:24.807945Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"17185852","cell_type":"code","source":"%time\n#autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.bfloat16) # if your gpu is newer Ampere, you can use this, lesser appearance of nan than half\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half) # you can use with T4 gpu. or newer\nscaler = torch.cuda.amp.GradScaler(enabled=USE_AMP, init_scale=2048)\n\nval_losses = []\ntrain_losses = []\ndf_tr, df_test = train_test_split(df, test_size=2/7, random_state=SEED)\nskf = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\ndevice = DEVICE\n\nfor fold, (trn_idx, val_idx) in enumerate(skf.split(range(len(df)))):\n    loss_scale = 1\n    if NOT_DEBUG == False:\n        if fold == 1: break;\n    if fold not in SELECTED_FOLDS: \n        print(f\"Jump fold {fold}\")\n        continue;\n    else:\n        print('#'*30)\n        print(f'Start fold {fold}')\n        print('#'*30)\n        print(len(trn_idx), len(val_idx))\n        df_train = df.iloc[trn_idx]\n        df_valid = df.iloc[val_idx]\n\n        train_ds = RSNA24Dataset(df_train, phase='train', transform=transforms_train)\n        train_dl = DataLoader(\n                        train_ds,\n                        batch_size=BATCH_SIZE,\n                        shuffle=True,\n                        pin_memory=False,\n                        drop_last=True,\n                        num_workers=N_WORKERS\n                        )\n\n        valid_ds = RSNA24Dataset(df_valid, phase='valid', transform=transforms_val)\n        valid_dl = DataLoader(\n                        valid_ds,\n                        batch_size=BATCH_SIZE*2,\n                        shuffle=False,\n                        pin_memory=False,\n                        drop_last=False,\n                        num_workers=N_WORKERS\n                        )\n\n    #         model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=True)\n        model = TimmModel(MODEL_NAME, pretrained=True)\n            \n        fname = f'{OUTPUT_DIR}/best_wll_model_fold-{fold}.pt'\n    #         if os.path.exists(fname):\n    #             model = TimmModel(MODEL_NAME, pretrained=False)\n    #             model.load_state_dict(torch.load(fname))\n        model.to(device)\n\n        optimizer = AdamW(model.parameters(), lr=LR*2, weight_decay=WD)\n    #         optimizer = torch.optim.SGD(model.parameters(), lr=LR*2, weight_decay=WD, nesterov=True, momentum=0.9)\n\n        warmup_steps = EPOCHS/10 * len(train_dl) // GRAD_ACC\n        num_total_steps = EPOCHS * len(train_dl) // GRAD_ACC\n        num_cycles = 0.475\n        scheduler = get_cosine_schedule_with_warmup(optimizer,\n                                                        num_warmup_steps=warmup_steps,\n                                                        num_training_steps=num_total_steps,\n                                                        num_cycles=num_cycles)\n    #         scheduler = get_linear_schedule_with_warmup(optimizer,\n    #                                                     num_warmup_steps=warmup_steps,\n    #                                                     num_training_steps=num_total_steps)\n\n        weights = torch.tensor([1.0, 2.0, 4.0])\n        criterion = nn.CrossEntropyLoss(weight=weights.to(device))\n        criterion_cpu = nn.CrossEntropyLoss(weight=weights)\n        best_loss = 1.2\n        es_step = 0\n\n        for epoch in range(1, EPOCHS+1):\n            print(f'start epoch {epoch}')\n            model.train()\n            total_loss = 0\n            with tqdm(train_dl, leave=True) as pbar:\n                optimizer.zero_grad()\n                for idx, (x, t) in enumerate(pbar):  \n                    op = ['nothing', 'nothing', 'nothing', 'nothing', 'nothing']\n                    x = x.to(device)\n                    t = t.to(device)\n    #                     t = torch.tensor(np.array(one_h(list(t.detach().cpu().numpy())))).to(device)\n                    rc = random.sample(op, 1)\n                    if rc[0] == 'mixup':\n                        x = x.detach().cpu().numpy()\n                        t = t.detach().cpu().numpy()\n                        reference_data = [{'image':x[i], 'proba': t[i]} \n                                            for i in range(len(x))]\n                        tr = A.Compose([A.MixUp(reference_data=reference_data,\n                                                  read_fn=read_fn, p=0.5)])\n                        for i in range(len(x)):\n                            transformed = tr(image=x[i], global_label=t[i])\n                            x[i] = transformed['image']\n                            t[i] = transformed['global_label']\n\n                        x = torch.tensor(x).to(device)\n                        t = torch.tensor(t).to(device)\n\n                    with autocast:\n                        loss = 0\n                        y = model(x)\n                        for col in range(N_LABELS):\n                            pred = y[:,col*3:col*3+3]\n                            gt = t[:,col]\n                            loss = loss + loss_scale * criterion(pred, gt) / N_LABELS\n\n                        if not math.isfinite(loss):\n                            loss = torch.tensor(1.2 * loss_scale * GRAD_ACC, requires_grad=True)\n                        total_loss += loss.item()\n                        if GRAD_ACC > 1:\n                            loss = loss / GRAD_ACC\n\n                    pbar.set_postfix(\n                            OrderedDict(\n                                loss=f'{loss.item()*GRAD_ACC:.6f}',\n                                lr=f'{optimizer.param_groups[0][\"lr\"]:.3e}'\n                            )\n                    )\n    #                     scaler.scale(loss).backward()\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM or 1e9)\n\n                    if (idx + 1) % GRAD_ACC == 0:\n    #                         scaler.step(optimizer)\n    #                         scaler.update()\n                        optimizer.step()\n                        optimizer.zero_grad()\n                        if scheduler is not None:\n                            scheduler.step()                    \n\n            train_loss = total_loss/len(train_dl)\n            print(f'train_loss:{train_loss/loss_scale:.6f}')\n            train_losses.append(train_loss)\n            total_loss = 0\n\n            model.eval()\n            y_preds, labels = [], []\n            with tqdm(valid_dl, leave=True) as pbar:\n                with torch.no_grad():\n                    for idx, (x, t) in enumerate(pbar):\n\n                        x = x.to(device)\n                        t = t.to(device)\n\n                        with autocast:\n                            loss = 0\n                            loss_ema = 0\n                            y = model(x)\n                            for col in range(N_LABELS):\n                                pred = y[:,col*3:col*3+3]\n                                gt = t[:,col]\n\n                                loss = loss + criterion(pred, gt) / N_LABELS\n                                y_pred = pred.float()\n                                y_preds.append(y_pred.cpu())\n                                labels.append(gt.cpu())\n\n                            if not math.isfinite(loss):\n                                loss = torch.tensor(1.2 * loss_scale * GRAD_ACC, requires_grad=True)\n\n                            total_loss += loss.item()   \n\n            val_loss = total_loss/len(valid_dl)\n            y_preds = torch.cat(y_preds, dim=0)\n            print(y_preds.shape)\n            labels = torch.cat(labels)\n\n            val_weighted_loss = criterion_cpu(y_preds, labels)\n            writer.add_scalar('val_wll', val_weighted_loss, epoch)\n            writer.flush()\n            print(f'val_loss:{val_loss:.6f}')\n            val_losses.append(val_loss)\n            if val_weighted_loss < best_loss:\n\n                if device!='cuda:0':\n                        model.to('cuda:0')                \n\n                print(f'epoch:{epoch}, best weighted_logloss updated from {best_loss:.6f} to {val_weighted_loss:.6f}')\n                best_loss = val_weighted_loss\n                fname = f'{OUTPUT_DIR}/best_wll_model_fold-{fold}.pt'\n                torch.save(model.state_dict(), fname)\n                print(f'{fname} is saved')\n                es_step = 0\n\n                if device!='cuda:0':\n                    model.to(device)\n\n            else:\n                es_step += 1\n                if es_step >= EARLY_STOPPING_EPOCH:\n                    print('early stopping')\n                    break  \n                                ","metadata":{"execution":{"iopub.status.busy":"2025-01-23T18:27:26.407304Z","iopub.execute_input":"2025-01-23T18:27:26.407657Z","iopub.status.idle":"2025-01-23T18:27:52.923637Z","shell.execute_reply.started":"2025-01-23T18:27:26.407631Z","shell.execute_reply":"2025-01-23T18:27:52.921254Z"},"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"id":"6c53934c","cell_type":"code","source":"cv = 0\ny_preds = []\nlabels = []\nweights = torch.tensor([1.0, 2.0, 4.0])\ncriterion2 = nn.CrossEntropyLoss(weight=weights)\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half) # you can use with T4 gpu. or newer\n\n\n## TODO: Modify EXIST_FOLDS by how many fold you've trained\nEXIST_FOLDS = [0, 1, 2, 3, 4]\nskf = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfor fold, (trn_idx, val_idx) in enumerate(skf.split(range(len(df)))):\n    #     if NOT_DEBUG == False:\n    #         if fold == 1: break;\n    if fold not in EXIST_FOLDS: \n        print(f\"Jump fold {fold}\")\n        continue;\n    else:\n        print('#'*30)\n        print(f'Start fold {fold}')\n        print('#'*30)\n        df_valid = df.iloc[val_idx]\n        valid_ds = RSNA24Dataset(df_valid, phase='valid', transform=transforms_val)\n        valid_dl = DataLoader(\n                        valid_ds,\n                        batch_size=16,\n                        shuffle=False,\n                        pin_memory=False,\n                        drop_last=False,\n                        num_workers=N_WORKERS\n                        )\n            \n\n        model = TimmModelCombo(MODEL_NAME)\n                \n            # print(\"No internet read\")\n        fname = f'{OUTPUT_DIR}/best_wll_model_fold-{fold}.pt'\n        model.load_state_dict(torch.load(fname))\n        model.to(device)   \n\n        model.eval()\n        with tqdm(valid_dl, leave=True) as pbar:\n            with torch.no_grad():\n                for idx, (x, t) in enumerate(pbar):\n\n                    x = x.to(device)\n                    t = t.to(device)\n\n                    with autocast:\n                        y = model(x)\n                        for col in range(N_LABELS):\n                            pred = y[:,col*3:col*3+3]\n                            gt = t[:,col] \n                            y_pred = pred.float()\n                            y_preds.append(y_pred.cpu())\n                            labels.append(gt.cpu())\n\ny_preds = torch.cat(y_preds)\nlabels = torch.cat(labels)","metadata":{"execution":{"iopub.status.busy":"2024-07-17T08:54:29.628236Z","iopub.status.idle":"2024-07-17T08:54:29.628659Z","shell.execute_reply":"2024-07-17T08:54:29.628490Z","shell.execute_reply.started":"2024-07-17T08:54:29.628473Z"}},"outputs":[],"execution_count":null},{"id":"e2facd86","cell_type":"code","source":"cv = criterion2(y_preds, labels)\nprint('cv score:', cv.item())","metadata":{"execution":{"iopub.status.busy":"2024-07-17T08:54:29.630786Z","iopub.status.idle":"2024-07-17T08:54:29.631268Z","shell.execute_reply":"2024-07-17T08:54:29.631055Z","shell.execute_reply.started":"2024-07-17T08:54:29.631035Z"}},"outputs":[],"execution_count":null},{"id":"68b00191","cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"id":"a0482ba3","cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"id":"dbf76a49","cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}