{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9069868,"sourceType":"datasetVersion","datasetId":5459726}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! cp -r /kaggle/input/rsna2024-public/src .\n! ls","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:50.416208Z","iopub.execute_input":"2025-01-11T23:20:50.416439Z","iopub.status.idle":"2025-01-11T23:20:50.690477Z","shell.execute_reply.started":"2025-01-11T23:20:50.416406Z","shell.execute_reply":"2025-01-11T23:20:50.689639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport polars as pl\nimport sys\nimport time\nimport yaml\nfrom sklearn.model_selection import KFold\n\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\n\ndi = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ndevice = torch.device('cuda')\n\ntb_global = time.time()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:50.691456Z","iopub.execute_input":"2025-01-11T23:20:50.691681Z","iopub.status.idle":"2025-01-11T23:20:55.404169Z","shell.execute_reply.started":"2025-01-11T23:20:50.691663Z","shell.execute_reply":"2025-01-11T23:20:55.403521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# y_pred (Tensor[float32]): logit with shape (batch_size, 3, 25)\n# y      (Tensor[int]):     ground truth index (batch_size, 25)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.404954Z","iopub.execute_input":"2025-01-11T23:20:55.405252Z","iopub.status.idle":"2025-01-11T23:20:55.408362Z","shell.execute_reply.started":"2025-01-11T23:20:55.405233Z","shell.execute_reply":"2025-01-11T23:20:55.407657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pl.read_csv(di + '/train.csv')\ncolumns = train.columns[1:]\n\nprint('spinal:      ', columns[:5][:2], '...');     assert all(['spinal' in c for c in columns[:5]])\nprint('foraminal:   ', columns[5:15][:2], '...');   assert all(['foraminal' in c for c in columns[5:15]])\nprint('subarticular:', columns[15:25][:2], '...');  assert all(['subarticular' in c for c in columns[15:25]])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.410406Z","iopub.execute_input":"2025-01-11T23:20:55.410602Z","iopub.status.idle":"2025-01-11T23:20:55.537555Z","shell.execute_reply.started":"2025-01-11T23:20:55.410584Z","shell.execute_reply":"2025-01-11T23:20:55.536467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Custom Loss for this competition\n\nclass SevereLoss(_Loss):\n    \"\"\"\n    For RSNA 2024\n    criterion = SevereLoss()     # you can replace nn.CrossEntropyLoss\n    loss = criterion(y_pred, y)\n    \"\"\"\n    def __init__(self, temperature=1.0):\n        \"\"\"\n        Use max if temperature = 0\n        \"\"\"\n        super().__init__()\n        self.t = temperature\n        assert self.t >= 0\n    \n    def __repr__(self):\n        return 'SevereLoss(t=%.1f)' % self.t\n\n    def forward(self, y_pred: torch.Tensor, y: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n          y_pred (Tensor[float]): logit             (batch_size, 3, 25)\n          y      (Tensor[int]):   true label index  (batch_size, 25)\n        \"\"\"\n        assert y_pred.size(0) == y.size(0)\n        assert y_pred.size(1) == 3 and y_pred.size(2) == 25\n        assert y.size(1) == 25\n        assert y.size(0) > 0\n        \n        slices = [slice(0, 5), slice(5, 15), slice(15, 25)] \n        w = 2 ** y  # sample_weight w = (1, 2, 4) for y = 0, 1, 2 (batch_size, 25)\n\n        loss = F.cross_entropy(y_pred, y, reduction='none')  # (batch_size, 25)\n\n        # Weighted sum of losses for spinal (:5), foraminal (5:15), and subarticular (15:25)\n        wloss_sums = []\n        for k, idx in enumerate(slices):\n            wloss_sums.append((w[:, idx] * loss[:, idx]).sum())\n\n        # Spinal max\n        y_spinal_prob = y_pred[:, :, :5].softmax(dim=1)             # (batch_size, 3,  5)\n        w_max = torch.amax(w[:, :5], dim=1)                         # batch_size\n        y_max = torch.amax(y[:, :5] == 2, dim=1).to(torch.float32)  # 0 or 1\n\n        if self.t > 0:\n            # Attention for the maximum value\n            attn = F.softmax(y_spinal_prob[:, 2, :] / self.t, dim=1)     # (batch_size, 5)\n\n            # Approximately the max among 5 severe=2 y_spinal_probs\n            y_pred_max = (attn * y_spinal_prob[:, 2, :]).sum(dim=1)     # weighted average among 5 spinal columns, \n        else:\n            # Exact max; this works too\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        # See below about these numbers\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.539101Z","iopub.execute_input":"2025-01-11T23:20:55.539438Z","iopub.status.idle":"2025-01-11T23:20:55.550548Z","shell.execute_reply.started":"2025-01-11T23:20:55.539406Z","shell.execute_reply":"2025-01-11T23:20:55.549023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute the weight global average\nweight_map = {'Normal/Mild': 1,\n              'Moderate': 2,\n              'Severe': 4,\n              None: 0}\n\nw_sum = [0, ] * 4\n\nfor r in train.iter_rows():\n    w = np.array([weight_map[x] for x in r[1:]])  # array[int] (25, )\n    assert len(w) == 25\n\n    w_sum[0] += w[:5].sum()    # spinal\n    w_sum[1] += w[5:15].sum()  # foraminal\n    w_sum[2] += w[15:25].sum() # subarticular \n    w_sum[3] += w[:5].max()    # any_severe_spinal\n\nfor k in range(4):\n    w_sum[k] /= len(train)\n\nw_sum\n# (6.084050632911392, 12.962531645569621, 14.38632911392405, 1.729113924050633)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.552278Z","iopub.execute_input":"2025-01-11T23:20:55.5526Z","iopub.status.idle":"2025-01-11T23:20:55.618978Z","shell.execute_reply.started":"2025-01-11T23:20:55.552568Z","shell.execute_reply":"2025-01-11T23:20:55.618028Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"debug = False  # Train only 2 epochs if True\n\ncfg = yaml.safe_load(\"\"\"\ndata:\n  image_size_in: 224\n\nmodel:\n  encoder: convnext_tiny.in12k_ft_in1k  # 224\n\nkfold:\n  k: 5\n  folds: [0, ]\n\ntrain:\n  lr: 1e-4\n  epochs: 20\n  weight_decay: 1e-4\n  max_grad_norm: 1000  # these two are negligibly weak\n\nvalidate:\n  every_n_epoch: 1\n\nloader:\n  batch_size: 8\n  num_workers: 2\n\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.619967Z","iopub.execute_input":"2025-01-11T23:20:55.620287Z","iopub.status.idle":"2025-01-11T23:20:55.626064Z","shell.execute_reply.started":"2025-01-11T23:20:55.620253Z","shell.execute_reply":"2025-01-11T23:20:55.625283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Hyperparameters\nweight_decay = float(cfg['train']['weight_decay'])\nmax_grad_norm = cfg['train']['max_grad_norm']\nval_every = cfg['validate']['every_n_epoch']\n\n# Criterion\ncriterion = SevereLoss(temperature=1.0)\n# criterion = nn.CrossEntropyLoss()       # for standard unweighted cross entropy\n\nprint('Criterion:', criterion)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.626979Z","iopub.execute_input":"2025-01-11T23:20:55.627309Z","iopub.status.idle":"2025-01-11T23:20:55.642491Z","shell.execute_reply.started":"2025-01-11T23:20:55.627276Z","shell.execute_reply":"2025-01-11T23:20:55.641756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load from src/data.py image.py model.py\nsys.path.append('src')\nfrom data import load_series_descriptions, load_labels, Dataset\nfrom model import Model\n\ndatad = load_series_descriptions('train')\nload_labels(datad)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:20:55.643276Z","iopub.execute_input":"2025-01-11T23:20:55.643501Z","iopub.status.idle":"2025-01-11T23:21:30.311247Z","shell.execute_reply.started":"2025-01-11T23:20:55.643482Z","shell.execute_reply":"2025-01-11T23:21:30.310499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, loader_val):\n    tb = time.time()\n\n    was_training = model.training\n    model.eval()\n\n    n_sum = 0\n    loss0_sum = 0.0  # loss for the criterion\n    \n    # 4 losses for the evaluation metric\n    loss4_sum = torch.zeros(4, device=device)\n    w_sum = torch.zeros(4, device=device)\n    slices = [slice(0, 5), slice(5, 15), slice(15, 25)]  # spinal, foraminal, subarticular\n\n    for d in loader_val:\n        x = d['x'].to(device)  # input image\n        y = d['y'].to(device)  # int (batch_size, 25)\n        batch_size = len(x)\n\n        # Predict\n        with torch.no_grad():\n            y_pred = model(x)  # (batch_size, 3, 25)\n\n        w = 2 ** y  # sample_weight w = (1, 2, 4) for y = 0, 1, 2 (batch_size, 25)\n\n        loss0 = criterion(y_pred, y)\n\n        n_sum += batch_size\n        loss0_sum += loss0.item() * batch_size\n\n        # Compute score\n        # - weighted loss for spinal, foraminal, subarticular\n        # - binary cross entropy for maximum spinal severe\n        ce_loss = F.cross_entropy(y_pred, y, reduction='none')  # (batch_size, 25)\n        for k, idx in enumerate(slices):\n            w_sum[k] += w[:, idx].sum()\n            loss4_sum[k] += (w[:, idx] * ce_loss[:, idx]).sum()\n\n        # Spinal max\n        y_spinal_prob = y_pred[:, :, :5].softmax(dim=1)            # (batch_size, 3,  5)\n        w_max = torch.amax(w[:, :5], dim=1)                        # (batch_size, )\n        y_max = torch.amax(y[:, :5] == 2, dim=1).to(torch.float)   # 0 or 1\n        y_pred_max = y_spinal_prob[:, 2, :].amax(dim=1)            # max in severe (class=2)\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    # Average over spinal, foraminal, subarticular, and any_severe_spinal\n    score = (loss_sum / w_sum).sum().item() / 4\n\n    model.train(was_training)\n\n    dt = time.time() - tb\n    ret = {'loss': loss0_sum / n_sum,\n           'score': score,\n           'dt': dt}\n    return ret","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:21:30.312099Z","iopub.execute_input":"2025-01-11T23:21:30.312413Z","iopub.status.idle":"2025-01-11T23:21:30.32031Z","shell.execute_reply.started":"2025-01-11T23:21:30.312382Z","shell.execute_reply":"2025-01-11T23:21:30.319335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# KFold\nstudy_ids = [sid for sid, d in datad.items() if d['filenames'] is not None]\ndf = pl.DataFrame({'study_id': study_ids})\n\nnfolds = cfg['kfold']['k']\nfolds = cfg['kfold']['folds']  # list[int]\nkfold = KFold(n_splits=nfolds, shuffle=True, random_state=42)\nprint('Folds', folds, '/', nfolds)\n\n\n#\n# Training loop\n#\nstudy_ids = [sid for sid, d in datad.items() if d['filenames'] is not None]\ndf = pl.DataFrame({'study_id': study_ids})\n\nfor ifold, (idx_train, idx_val) in enumerate(kfold.split(df)):\n    if ifold not in folds:\n        continue\n\n    # Data\n    ds_train = Dataset(df[idx_train], datad, cfg, pick='random', augment=True)\n    ds_val =   Dataset(df[idx_val],   datad, cfg, pick='middle')\n\n    loader_train = ds_train.loader(cfg, shuffle=True, drop_last=True)\n    loader_val = ds_val.loader(cfg)\n\n    nbatch = len(loader_train)\n\n    # Model: \n    model = Model(cfg, pretrained=False)  # Just use pretrained=True with internet\n    model_filename = '/kaggle/input/rsna2024-public/weights/model_pretrained.pytorch'\n    model.load_state_dict(torch.load(model_filename))  # pretrained model with internet=off\n    model.to(device)\n    model.train()\n\n    # Optimizer\n    lr = float(cfg['train']['lr'])\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr,\n                                  weight_decay=weight_decay)\n\n    epochs = 2 if debug else cfg['train']['epochs']\n    scheduler = get_cosine_schedule_with_warmup(optimizer,\n                    num_warmup_steps=nbatch,\n                    num_training_steps=(epochs * nbatch))\n\n    print('%d epochs' % epochs)\n\n    # n-epoch loop\n    tb = time.time()\n    dt_val = 0\n    loss_sum, n_sum = 0, 0\n    \n    print('Epoch  loss          score   lr      time')\n    for iepoch in range(epochs):\n        for ibatch, d in enumerate(loader_train):\n            x = d['x'].to(device)  # input image\n            y = d['y'].to(device)  # segmentation label\n            batch_size = len(x)\n\n            optimizer.zero_grad()\n\n            # Predict\n            y_pred = model(x)      # (batch_size, 3, 25)\n            loss = criterion(y_pred, y)\n\n            # Backpropagate\n            loss.backward()\n            n_sum += batch_size\n            loss_sum += batch_size * loss.item()\n\n            nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)\n            optimizer.step()\n\n            scheduler.step()\n        \n        # Validation\n        if (iepoch + 1) % val_every == 0:  # validation set\n            loss_train = loss_sum / n_sum\n\n            val = evaluate(model, loader_val)\n            dt_val += val['dt']\n            \n            lr = optimizer.param_groups[0]['lr']\n\n            dt = time.time() - tb\n            print('%3d %7.4f %7.4f  %.4f  %5.1e %.2f %.2f min' % (iepoch + 1,\n                  loss_train, val['loss'], val['score'],\n                  lr, dt_val / 60, dt / 60))\n\n            loss_sum, n_sum = 0, 0\n\n    # Save model\n    model.to('cpu')\n    model.eval()\n    ofilename = 'model%d.pytorch' % ifold\n    torch.save(model.state_dict(), ofilename)\n\n    print(ofilename, 'written')\n\n\nif debug:\n    print('\\nDebug %r: train only %d epochs' % (debug, epochs))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:27:15.860877Z","iopub.execute_input":"2025-01-11T23:27:15.861195Z","iopub.status.idle":"2025-01-11T23:49:37.601729Z","shell.execute_reply.started":"2025-01-11T23:27:15.861174Z","shell.execute_reply":"2025-01-11T23:49:37.600579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_missing_studies(datad, preds):\n    # I skip study_ids with no Sagittal T1\n    # filling such study_id with 1/3\n    y_pred = torch.ones((3, 25), dtype=torch.float32) / 3\n    for study_id, d in datad.items():\n        if d['filenames'] is None:\n            pred = {'study_id': study_id,\n                    'y_pred': y_pred}\n            preds.append(pred)\n\n    assert len(preds) == len(datad)\n\n\ndef predict(model, loader, device):\n    preds = []\n    for d in loader:\n        x = d['x'].to(device)  # input image\n        batch_size = len(x)\n\n        # Predict\n        with torch.no_grad():\n            y_pred = model(x)  # (batch_size, 3, 25)\n\n        study_ids = d['study_id']\n        y_pred = y_pred.cpu()  # Tensor[float32] (batch_size, 3, 25)\n        for i in range(batch_size):\n            pred = {'study_id': study_ids[i],\n                    'y_pred': y_pred[i].clone(),  # Tensor[float32] (3, 25)\n            }\n            preds.append(pred)\n\n    return preds\n\n\ndef check_submission(submit):\n    # Check row_ids exactly match those in sample_submission.csv\n    sample = pl.read_csv(di + '/sample_submission.csv')\n\n    assert submit.shape == sample.shape\n    assert submit.columns == sample.columns\n    assert (submit['row_id'] == sample['row_id']).all()\n\n\ndef create_submission(preds, columns):\n    \"\"\"\n    columns list[str]: name of 25 targets, train.columns[1:]\n    \"\"\"\n    assert len(columns) == 25\n\n    rows = []\n    for pred in preds:\n        study_id = pred['study_id']\n        y_pred = pred['y_pred'].to(torch.float64)  # Tensor[float] (3, 25)\n        y_pred = y_pred.softmax(dim=0).numpy()     # logit -> normalized probability\n        assert y_pred.shape == (3, 25)\n\n        for j, name in enumerate(columns):\n            row_id = '%s_%s' % (study_id, name)\n            row = (row_id, ) + tuple(y_pred[:, j])\n            assert len(row) == 4\n            rows.append(row)\n\n    # Sort like the sample submission.\n    rows.sort()\n\n    col_names = ('row_id', 'normal_mild', 'moderate', 'severe')\n    submit = pl.DataFrame(rows, schema=col_names, orient='row')\n\n    return submit","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:51:54.104982Z","iopub.execute_input":"2025-01-11T23:51:54.105311Z","iopub.status.idle":"2025-01-11T23:51:54.11446Z","shell.execute_reply.started":"2025-01-11T23:51:54.105289Z","shell.execute_reply":"2025-01-11T23:51:54.113518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test data\ndatad = load_series_descriptions('test')\nds_test = Dataset(None, datad, cfg, pick='middle')\nloader_test = ds_test.loader(cfg)\nprint('Data', len(ds_test))\n\n# Model trained locally for 20 epochs\nmodel_filename = '/kaggle/input/rsna2024-public/weights/model0.pytorch'\n\nmodel = Model(cfg)\nmodel.load_state_dict(torch.load(model_filename))\nmodel.to(device)\nmodel.eval()\n\n# Predict\npreds = predict(model, loader_test, device)\n\n# Write submission.csv\nofilename = 'submission.csv'\nsubmit = create_submission(preds, train.columns[1:])\nsubmit.write_csv(ofilename, float_precision=16)\ncheck_submission(submit)\n\nprint('%s written. %d rows' % (ofilename, len(submit)))\n\ndt = time.time() - tb_global\nprint('Total %.1f min' % (dt / 60))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:52:05.86694Z","iopub.execute_input":"2025-01-11T23:52:05.867236Z","iopub.status.idle":"2025-01-11T23:52:09.029057Z","shell.execute_reply.started":"2025-01-11T23:52:05.867213Z","shell.execute_reply":"2025-01-11T23:52:09.027952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! head -n 3 submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T23:52:21.084652Z","iopub.execute_input":"2025-01-11T23:52:21.085044Z","iopub.status.idle":"2025-01-11T23:52:21.24374Z","shell.execute_reply.started":"2025-01-11T23:52:21.085012Z","shell.execute_reply":"2025-01-11T23:52:21.242823Z"}},"outputs":[],"execution_count":null}]}