{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":189030930,"sourceType":"kernelVersion"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\n\nimport numpy as np\nimport pandas as pd\npd.set_option(\"future.no_silent_downcasting\", True)\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport timm\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import EarlyStopping\n\nfrom sklearn.model_selection import KFold\n\nfrom tqdm import tqdm\nimport pickle\nfrom copy import deepcopy","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:40:09.129733Z","iopub.execute_input":"2024-07-23T19:40:09.130136Z","iopub.status.idle":"2024-07-23T19:40:09.136682Z","shell.execute_reply.started":"2024-07-23T19:40:09.130108Z","shell.execute_reply":"2024-07-23T19:40:09.135556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wandb","metadata":{}},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\napi_key = user_secrets.get_secret(\"wandb_api_key\")\n\nwandb.login(key = api_key)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:32:23.862030Z","iopub.execute_input":"2024-07-23T19:32:23.862317Z","iopub.status.idle":"2024-07-23T19:32:25.672376Z","shell.execute_reply.started":"2024-07-23T19:32:23.862293Z","shell.execute_reply":"2024-07-23T19:32:25.671464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config():\n    \n    DEBUG = False\n    SEED = 454\n    \n    n_slices_per_series = 10\n    image_size = (256, 256)\n    \n    root = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n    volumes_root = '/kaggle/input/rsna-1-trial-dataset/volumes'\n    \n    N_FOLDS = 5\n    MAX_EPOCHS = 20\n    EARLY_STOPPING_ROUNDS = 3\n    \n    MODEL_NAME = 'tf_efficientnet_b0.ns_jft_in1k'\n    PRETRAINED = True\n    IN_CHANNELS = 3 * n_slices_per_series\n    NUM_CLASSES = 75\n    NUM_LABELS = NUM_CLASSES // 3\n    GLOBAL_POOL = 'avg'\n    \n    LR = 1e-4\n    BATCH_SIZE = 32\n    \n    # Must be negative\n    # For any_spinal_severe\n    # Max logic\n    NAN_IDX = -1","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:32:25.673412Z","iopub.execute_input":"2024-07-23T19:32:25.673884Z","iopub.status.idle":"2024-07-23T19:32:25.680657Z","shell.execute_reply.started":"2024-07-23T19:32:25.673858Z","shell.execute_reply":"2024-07-23T19:32:25.679661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\npl.seed_everything(Config.SEED)\nnp.random.seed(Config.SEED)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T19:32:25.683027Z","iopub.execute_input":"2024-07-23T19:32:25.683822Z","iopub.status.idle":"2024-07-23T19:32:25.729334Z","shell.execute_reply.started":"2024-07-23T19:32:25.683781Z","shell.execute_reply":"2024-07-23T19:32:25.728544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Volume Example","metadata":{}},{"cell_type":"code","source":"volumes_files = os.listdir(Config.volumes_root)\n\nvolumes_path = [os.path.join(Config.volumes_root, file) for file in volumes_files]\npresent_studies = [int(file[:-4]) for file in volumes_files]\nvolumes_path = pd.Series(volumes_path, index=present_studies, name='path')","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:09.696774Z","iopub.execute_input":"2024-07-23T04:52:09.697018Z","iopub.status.idle":"2024-07-23T04:52:09.961607Z","shell.execute_reply.started":"2024-07-23T04:52:09.696996Z","shell.execute_reply":"2024-07-23T04:52:09.960737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"volume = np.load(volumes_path.iloc[0])","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:09.962822Z","iopub.execute_input":"2024-07-23T04:52:09.963531Z","iopub.status.idle":"2024-07-23T04:52:09.987259Z","shell.execute_reply.started":"2024-07-23T04:52:09.963500Z","shell.execute_reply":"2024-07-23T04:52:09.986414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N = volume.shape[0]\nnrows = N // 3\n\nfig, axarr = plt.subplots(ncols=3, nrows=nrows, figsize=(3*3, 3*10))\n\nfor i in range(N):\n    \n    ax = axarr[i % nrows][i // nrows]\n    \n    ax.set_xticks([])\n    ax.set_yticks([])\n    ax.imshow(volume[i], cmap='bone')","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:09.988358Z","iopub.execute_input":"2024-07-23T04:52:09.988624Z","iopub.status.idle":"2024-07-23T04:52:12.123517Z","shell.execute_reply.started":"2024-07-23T04:52:09.988596Z","shell.execute_reply":"2024-07-23T04:52:12.122569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mean and variance\n\n- Not enough RAM to load all volumes at once and calculate mean and std.\n- Do these formulas give the same result as loading all images and calculating mean and std\n- Need to do it in KFold?","metadata":{}},{"cell_type":"code","source":"def get_sample_mean(df):\n    \n    means = []\n\n    for path in df.path:\n        \n        volume = np.load(path)\n\n        means.append(volume.mean())\n\n    mean = np.mean(means)\n    \n    return mean","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.124744Z","iopub.execute_input":"2024-07-23T04:52:12.125038Z","iopub.status.idle":"2024-07-23T04:52:12.129993Z","shell.execute_reply.started":"2024-07-23T04:52:12.125013Z","shell.execute_reply":"2024-07-23T04:52:12.129208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_sample_std(df, mean):\n    \n    variances = []\n\n    for path in df.path:\n\n        volume = np.load(path)\n\n        variance = ((volume - mean) ** 2).mean()\n        variances.append(variance)\n\n    std = np.sqrt(np.mean(variances))\n    \n    return std","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.131073Z","iopub.execute_input":"2024-07-23T04:52:12.131398Z","iopub.status.idle":"2024-07-23T04:52:12.142196Z","shell.execute_reply.started":"2024-07-23T04:52:12.131372Z","shell.execute_reply":"2024-07-23T04:52:12.141088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_sample_stats(df):\n    \n    mean = get_sample_mean(df)\n    \n    return mean, get_sample_std(df, mean)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.145768Z","iopub.execute_input":"2024-07-23T04:52:12.146063Z","iopub.status.idle":"2024-07-23T04:52:12.151098Z","shell.execute_reply.started":"2024-07-23T04:52:12.146041Z","shell.execute_reply":"2024-07-23T04:52:12.150201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Labels Preparation","metadata":{}},{"cell_type":"code","source":"labels = pd.read_csv(os.path.join(Config.root, 'train.csv'))\nlabels = labels.set_index('study_id')\n\nprint(f'Number of unique studies - {labels.shape}')\n\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.152265Z","iopub.execute_input":"2024-07-23T04:52:12.152634Z","iopub.status.idle":"2024-07-23T04:52:12.412647Z","shell.execute_reply.started":"2024-07-23T04:52:12.152604Z","shell.execute_reply":"2024-07-23T04:52:12.411773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = labels.replace({'Normal/Mild':0, 'Moderate':1, 'Severe':2})\nlabels = labels.fillna(Config.NAN_IDX)\nlabels = labels.astype('int8')\n\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.413764Z","iopub.execute_input":"2024-07-23T04:52:12.414038Z","iopub.status.idle":"2024-07-23T04:52:12.454875Z","shell.execute_reply.started":"2024-07-23T04:52:12.414015Z","shell.execute_reply":"2024-07-23T04:52:12.453997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.concat([labels, volumes_path], axis=1)\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.455797Z","iopub.execute_input":"2024-07-23T04:52:12.456028Z","iopub.status.idle":"2024-07-23T04:52:12.479523Z","shell.execute_reply.started":"2024-07-23T04:52:12.456007Z","shell.execute_reply":"2024-07-23T04:52:12.478588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df[df.path.notna()]\n\ndf.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.480579Z","iopub.execute_input":"2024-07-23T04:52:12.480839Z","iopub.status.idle":"2024-07-23T04:52:12.489709Z","shell.execute_reply.started":"2024-07-23T04:52:12.480817Z","shell.execute_reply":"2024-07-23T04:52:12.488711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_sample_stats(df)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.490954Z","iopub.execute_input":"2024-07-23T04:52:12.491357Z","iopub.status.idle":"2024-07-23T04:53:03.933188Z","shell.execute_reply.started":"2024-07-23T04:52:12.491314Z","shell.execute_reply":"2024-07-23T04:53:03.932191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    \n    def __init__(self, df, mean, std):\n        \n        self.path = df['path']\n        self.labels = df.drop(columns='path')\n        \n        self.mean = mean\n        self.std = std\n        \n    def __len__(self):\n        \n        return self.path.shape[0]\n    \n    def __getitem__(self, idx):\n        \n        volume = np.load(self.path.iloc[idx])\n        volume = volume.astype('float32')\n        volume = (volume - self.mean) / self.std\n        \n        return volume, self.labels.iloc[idx].values.astype('long')","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:03.934419Z","iopub.execute_input":"2024-07-23T04:53:03.934711Z","iopub.status.idle":"2024-07-23T04:53:03.941656Z","shell.execute_reply.started":"2024-07-23T04:53:03.934687Z","shell.execute_reply":"2024-07-23T04:53:03.940724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_loaders(train_idx, valid_idx, mean, std):\n\n    train_set = RSNADataset(df.iloc[train_idx], mean, std)\n    valid_set = RSNADataset(df.iloc[valid_idx], mean, std)\n    \n    train_loader = DataLoader(train_set, batch_size=Config.BATCH_SIZE, \n                              drop_last=True, shuffle=True, num_workers=3)\n    valid_loader = DataLoader(valid_set, batch_size=Config.BATCH_SIZE, \n                              drop_last=False, shuffle=False, num_workers=3)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:03.942963Z","iopub.execute_input":"2024-07-23T04:53:03.943389Z","iopub.status.idle":"2024-07-23T04:53:03.957486Z","shell.execute_reply.started":"2024-07-23T04:53:03.943358Z","shell.execute_reply":"2024-07-23T04:53:03.956707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Competition Metric (torch implementation)","metadata":{}},{"cell_type":"code","source":"metric_fn = nn.CrossEntropyLoss(weight=torch.Tensor([1, 2, 4]).to(device),\n                                ignore_index=Config.NAN_IDX)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:03.958560Z","iopub.execute_input":"2024-07-23T04:53:03.958845Z","iopub.status.idle":"2024-07-23T04:53:04.157892Z","shell.execute_reply.started":"2024-07-23T04:53:03.958823Z","shell.execute_reply":"2024-07-23T04:53:04.157130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_any_severe_spinal_loss(y_pred, y):\n    \n    severe_max_label = y[:, :5].max(axis=1)[0]\n    \n    nan_filter = (severe_max_label != Config.NAN_IDX)\n    severe_max_label = severe_max_label[nan_filter]\n    \n    any_severe_spinal_pred = []\n    for i in range(5):\n        any_severe_spinal_pred.append(\n            F.softmax(y_pred[:, i*3:(i+1)*3], dim=1)[:, -1]\n        )\n    any_severe_spinal_pred = torch.stack(\n        any_severe_spinal_pred,\n        dim=1\n    ).max(axis=1)[0][nan_filter]\n    \n    any_severe_spinal_true = (severe_max_label == 2) * 1\n    \n    # TODO: More efficient way to replace\n    # Does need copy?\n    any_severe_spinal_weights = severe_max_label.clone().float()\n    any_severe_spinal_weights[severe_max_label == 0] = 1\n    any_severe_spinal_weights[severe_max_label == 1] = 2\n    any_severe_spinal_weights[severe_max_label == 2] = 4\n    \n    any_severe_spinal_loss = F.binary_cross_entropy(\n        any_severe_spinal_pred,\n        any_severe_spinal_true.float(),\n        weight=any_severe_spinal_weights,\n        reduction='sum'\n    )\n\n    return any_severe_spinal_loss / any_severe_spinal_weights.sum()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.158884Z","iopub.execute_input":"2024-07-23T04:53:04.159136Z","iopub.status.idle":"2024-07-23T04:53:04.168034Z","shell.execute_reply.started":"2024-07-23T04:53:04.159114Z","shell.execute_reply":"2024-07-23T04:53:04.167128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Very strange any_severe_spinal weights logic - [discussion](https://www.kaggle.com/competitions/rsna-2024-lumbar-spine-degenerative-classification/discussion/520479#2930546)","metadata":{}},{"cell_type":"code","source":"def get_competition_metric(y_pred, y):\n    \n    spinal_loss = metric_fn(y_pred[:, :15].reshape((-1, 3)), \n                            y[:, :5].flatten())\n    \n    foraminal_loss = metric_fn(y_pred[:, 15:45].reshape((-1, 3)),\n                               y[:, 5:15].flatten())\n    \n    subarticular_loss = metric_fn(y_pred[:, 45:75].reshape((-1, 3)),\n                                  y[:, 15:25].flatten())\n    \n    any_severe_spinal_loss = get_any_severe_spinal_loss(y_pred, y)\n    \n    return (spinal_loss + foraminal_loss + subarticular_loss + any_severe_spinal_loss) / 4","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.169243Z","iopub.execute_input":"2024-07-23T04:53:04.169575Z","iopub.status.idle":"2024-07-23T04:53:04.176932Z","shell.execute_reply.started":"2024-07-23T04:53:04.169545Z","shell.execute_reply":"2024-07-23T04:53:04.176127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Host Competition metric implementation","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport pandas.api.types\nimport sklearn.metrics\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef get_condition(full_location: str) -> str:\n    # Given an input like spinal_canal_stenosis_l1_l2 extracts 'spinal'\n    for injury_condition in ['spinal', 'foraminal', 'subarticular']:\n        if injury_condition in full_location:\n            return injury_condition\n    raise ValueError(f'condition not found in {full_location}')\n\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        any_severe_scalar: float\n    ) -> float:\n    '''\n    Pseudocode:\n    1. Calculate the sample weighted log loss for each medical condition:\n    2. Derive a new any_severe label.\n    3. Calculate the sample weighted log loss for the new any_severe label.\n    4. Return the average of all of the label group log losses as the final score, normalized for the number of columns in each group.\n       This mitigates the impact of spinal stenosis having only half as many columns as the other two conditions.\n    '''\n\n    target_levels = ['normal_mild', 'moderate', 'severe']\n\n    # Run basic QC checks on the inputs\n    if not pandas.api.types.is_numeric_dtype(submission[target_levels].values):\n        raise ParticipantVisibleError('All submission values must be numeric')\n\n    if not np.isfinite(submission[target_levels].values).all():\n        raise ParticipantVisibleError('All submission values must be finite')\n\n    if solution[target_levels].min().min() < 0:\n        raise ParticipantVisibleError('All labels must be at least zero')\n    if submission[target_levels].min().min() < 0:\n        raise ParticipantVisibleError('All predictions must be at least zero')\n\n    solution['study_id'] = solution['row_id'].apply(lambda x: x.split('_')[0])\n    solution['location'] = solution['row_id'].apply(lambda x: '_'.join(x.split('_')[1:]))\n    solution['condition'] = solution['row_id'].apply(get_condition)\n\n    del solution[row_id_column_name]\n    del submission[row_id_column_name]\n    assert sorted(submission.columns) == sorted(target_levels)\n\n    submission['study_id'] = solution['study_id']\n    submission['location'] = solution['location']\n    submission['condition'] = solution['condition']\n\n    condition_losses = []\n    condition_weights = []\n    for condition in ['spinal', 'foraminal', 'subarticular']:\n        condition_indices = solution.loc[solution['condition'] == condition].index.values\n        condition_loss = sklearn.metrics.log_loss(\n            y_true=solution.loc[condition_indices, target_levels].values,\n            y_pred=submission.loc[condition_indices, target_levels].values,\n            sample_weight=solution.loc[condition_indices, 'sample_weight'].values\n        )\n        condition_losses.append(condition_loss)\n        condition_weights.append(1)\n\n    any_severe_spinal_labels = pd.Series(solution.loc[solution['condition'] == 'spinal'].groupby('study_id', sort=False)['severe'].max())\n    any_severe_spinal_weights = pd.Series(solution.loc[solution['condition'] == 'spinal'].groupby('study_id', sort=False)['sample_weight'].max())\n    any_severe_spinal_predictions = pd.Series(submission.loc[submission['condition'] == 'spinal'].groupby('study_id', sort=False)['severe'].max())\n    any_severe_spinal_loss = sklearn.metrics.log_loss(\n        y_true=any_severe_spinal_labels,\n        y_pred=any_severe_spinal_predictions,\n        sample_weight=any_severe_spinal_weights\n    )\n    condition_losses.append(any_severe_spinal_loss)\n    condition_weights.append(any_severe_scalar)\n    return np.average(condition_losses, weights=condition_weights)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-07-23T04:53:04.178079Z","iopub.execute_input":"2024-07-23T04:53:04.178387Z","iopub.status.idle":"2024-07-23T04:53:04.195367Z","shell.execute_reply.started":"2024-07-23T04:53:04.178349Z","shell.execute_reply":"2024-07-23T04:53:04.194453Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Comparison of metric implementations","metadata":{}},{"cell_type":"code","source":"logits= torch.rand((labels.shape[0], 75))\npred = logits.clone()\nlogits = logits.to(device)\n\nfor i in range(25):\n    pred[:, i*3:(i+1)*3] = F.softmax(logits[:, i*3:(i+1)*3], dim=1)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.196617Z","iopub.execute_input":"2024-07-23T04:53:04.196958Z","iopub.status.idle":"2024-07-23T04:53:04.273723Z","shell.execute_reply.started":"2024-07-23T04:53:04.196929Z","shell.execute_reply":"2024-07-23T04:53:04.272842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_tensor = torch.Tensor(labels.values).long().to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.274850Z","iopub.execute_input":"2024-07-23T04:53:04.275123Z","iopub.status.idle":"2024-07-23T04:53:04.281966Z","shell.execute_reply.started":"2024-07-23T04:53:04.275099Z","shell.execute_reply":"2024-07-23T04:53:04.281026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\n\nfor i in range(25):\n    t = pd.DataFrame(pred[:, i*3:(i+1)*3], columns=['normal_mild', 'moderate', 'severe'])\n    t['row_id'] = labels.reset_index()['study_id'].astype(str) + '_' + labels.columns[i]\n    preds.append(t)\n    \npreds = pd.concat(preds, axis=0).reset_index(drop=True)\npreds.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.282963Z","iopub.execute_input":"2024-07-23T04:53:04.283276Z","iopub.status.idle":"2024-07-23T04:53:04.373858Z","shell.execute_reply.started":"2024-07-23T04:53:04.283253Z","shell.execute_reply":"2024-07-23T04:53:04.373072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y = labels.reset_index().melt(id_vars='study_id')\ny['row_id'] = y['study_id'].astype(str) + '_' + y['variable']\ny = y.drop(columns=['study_id', 'variable'])\ny[['nan', 'normal_mild', 'moderate', 'severe']] = pd.get_dummies(y['value']) * 1\ny = y[y['nan'] == 0]\ny = y.drop(columns=['value', 'nan'])\ny['sample_weight'] = y['normal_mild'] * 1 + y['moderate'] * 2 + y['severe'] * 4\ny.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.374868Z","iopub.execute_input":"2024-07-23T04:53:04.375129Z","iopub.status.idle":"2024-07-23T04:53:04.451777Z","shell.execute_reply.started":"2024-07-23T04:53:04.375106Z","shell.execute_reply":"2024-07-23T04:53:04.450932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = preds[preds['row_id'].isin(y.row_id)]\npreds.shape, y.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.462491Z","iopub.execute_input":"2024-07-23T04:53:04.462754Z","iopub.status.idle":"2024-07-23T04:53:04.486903Z","shell.execute_reply.started":"2024-07-23T04:53:04.462731Z","shell.execute_reply":"2024-07-23T04:53:04.485918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_competition_metric(logits, y_tensor)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.488386Z","iopub.execute_input":"2024-07-23T04:53:04.488685Z","iopub.status.idle":"2024-07-23T04:53:04.857738Z","shell.execute_reply.started":"2024-07-23T04:53:04.488660Z","shell.execute_reply":"2024-07-23T04:53:04.856271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"score(y, preds, 'row_id', 1)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:04.859286Z","iopub.execute_input":"2024-07-23T04:53:04.859817Z","iopub.status.idle":"2024-07-23T04:53:05.080672Z","shell.execute_reply.started":"2024-07-23T04:53:04.859753Z","shell.execute_reply":"2024-07-23T04:53:05.079788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"def get_classic_log_loss(y_pred, y):\n    \n    loss = 0\n    # 75 logits are mapped to 25 multiclass problems\n    for i in range(Config.NUM_CLASSES // 3):\n        label_slice = slice(i*3, (i+1)*3)\n        loss += loss_fn(y_pred[:, label_slice], y[:, i])\n        \n    return loss / Config.NUM_LABELS","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:05.081698Z","iopub.execute_input":"2024-07-23T04:53:05.081995Z","iopub.status.idle":"2024-07-23T04:53:05.087158Z","shell.execute_reply.started":"2024-07-23T04:53:05.081969Z","shell.execute_reply":"2024-07-23T04:53:05.086266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss(ignore_index=Config.NAN_IDX)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:05.088329Z","iopub.execute_input":"2024-07-23T04:53:05.088668Z","iopub.status.idle":"2024-07-23T04:53:05.095788Z","shell.execute_reply.started":"2024-07-23T04:53:05.088636Z","shell.execute_reply":"2024-07-23T04:53:05.094940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class Model(pl.LightningModule):\n    \n    def __init__(self, fold):\n        \n        super(Model, self).__init__()\n        \n        self.fold = fold\n        \n        self.model =  timm.create_model(\n            model_name = Config.MODEL_NAME,\n            pretrained = Config.PRETRAINED, \n            features_only = False,\n            in_chans = Config.IN_CHANNELS,\n            num_classes = Config.NUM_CLASSES,\n            global_pool = Config.GLOBAL_POOL\n        )\n        \n        self._zero_metrics()\n        \n        self.run = wandb.init(\n            group='Гипотеза 1',\n            name=f'Fold #{self.fold}',\n            project='RSNA Lumbar'\n        )\n        \n        self.best_loss = 100\n        self.best_state = None\n        self.best_epoch = 0\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=Config.LR)\n        return optimizer\n    \n    def training_step(self, batch, batch_idx):\n        return self._step(batch, 'train')\n            \n    def validation_step(self, batch, batch_idx):\n        return self._step(batch, 'valid')\n        \n    def on_train_epoch_end(self):\n        self._log_epoch('train')\n        self._log_epoch('valid')\n        self._zero_metrics()\n    \n    def on_fit_end(self):\n        self.run.finish()\n        torch.save({\n            'model':self.best_state,\n            'best_epoch':self.best_epoch,\n            'best_loss':self.best_loss\n        }, f'best_model_{self.fold}')\n\n    def _step(self, batch, set_label):\n        \n        X, y = batch\n        y_pred = self.model(X)\n        \n        loss = get_classic_log_loss(y_pred, y)\n        metric = get_competition_metric(y_pred, y)\n        \n        self.metrics[set_label]['n'] += X.shape[0]\n        self.metrics[set_label]['loss'] += loss * X.shape[0]\n        self.metrics[set_label]['metric'] += metric * X.shape[0]\n\n        return loss\n    \n    def _log_epoch(self, set_label):\n        \n        n = self.metrics[set_label]['n']\n        loss = self.metrics[set_label]['loss'] / n\n        metric = self.metrics[set_label]['metric'] / n\n        \n        # For EarlyStopCallback\n        if set_label == 'valid':\n            self.log('val_loss', loss)\n            \n            if loss < self.best_loss:\n                self.best_loss = loss\n                \n                model.to('cpu')\n                self.best_state = deepcopy(model.state_dict())\n                self.best_epoch = self.current_epoch\n                model.to(device)\n        \n        self.run.log(\n            {\n                f'{set_label}_loss':loss,\n                f'{set_label}_metric':metric\n            },\n            step=self.current_epoch\n        )\n    \n    def _zero_metrics(self):\n        self.metrics = {\n            'train':{\n                'n':0,\n                'loss':0,\n                'metric':0\n            },\n            'valid':{\n                'n':0,\n                'loss':0,\n                'metric':0\n            }\n        }","metadata":{"execution":{"iopub.status.busy":"2024-07-23T05:12:13.573427Z","iopub.execute_input":"2024-07-23T05:12:13.574180Z","iopub.status.idle":"2024-07-23T05:12:13.591211Z","shell.execute_reply.started":"2024-07-23T05:12:13.574131Z","shell.execute_reply":"2024-07-23T05:12:13.590389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cross Validation","metadata":{}},{"cell_type":"code","source":"kf = KFold(Config.N_FOLDS, shuffle=True, random_state=Config.SEED)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T05:12:13.906071Z","iopub.execute_input":"2024-07-23T05:12:13.906721Z","iopub.status.idle":"2024-07-23T05:12:13.911450Z","shell.execute_reply.started":"2024-07-23T05:12:13.906689Z","shell.execute_reply":"2024-07-23T05:12:13.910384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (train_idx, valid_idx) in enumerate(kf.split(df)):\n    \n    model = Model(i)\n    \n    mean, std = get_sample_stats(df.iloc[train_idx])\n    \n    train_loader, valid_loader = get_loaders(train_idx, valid_idx, mean, std)\n    \n    early_stop_callback = EarlyStopping(\n        monitor='val_loss',\n        patience=Config.EARLY_STOPPING_ROUNDS,\n        mode='min',\n        strict=True\n    )\n    \n    trainer = pl.Trainer(accelerator=device, \n                         logger=None, \n                         fast_dev_run=Config.DEBUG,\n                         num_sanity_val_steps = 0,\n                         max_epochs=Config.MAX_EPOCHS,\n                         callbacks=[early_stop_callback])\n    \n    trainer.fit(model=model, \n                train_dataloaders=train_loader,\n                val_dataloaders=valid_loader)\n    \n    pickle.dump((mean, std), open(f'scaler_params_{i}.pkl', 'wb'))","metadata":{"execution":{"iopub.status.busy":"2024-07-23T05:12:56.929907Z","iopub.execute_input":"2024-07-23T05:12:56.930266Z","iopub.status.idle":"2024-07-23T05:14:06.035822Z","shell.execute_reply.started":"2024-07-23T05:12:56.930235Z","shell.execute_reply":"2024-07-23T05:14:06.034280Z"},"trusted":true},"execution_count":null,"outputs":[]}]}