{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom pathlib import Path\nimport timm\nimport random\nfrom typing import List\nimport glob\n\nfrom fastai.basics           import *\nfrom fastai.medical.imaging  import *\n\nfrom dataclasses import dataclass, field, asdict\n\nimport logging\nfrom transformers import Trainer, TrainingArguments\n\nimport albumentations as A\n\nfrom sklearn.model_selection import GroupKFold\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-12T07:39:31.176489Z","iopub.execute_input":"2024-09-12T07:39:31.176905Z","iopub.status.idle":"2024-09-12T07:39:52.086236Z","shell.execute_reply.started":"2024-09-12T07:39:31.176864Z","shell.execute_reply":"2024-09-12T07:39:52.085413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False\n\n@dataclass\nclass DatasetConfig:\n    path: Path = Path('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/')\n    label_columns = ['spinal_canal_stenosis_l1_l2',\n       'spinal_canal_stenosis_l2_l3', 'spinal_canal_stenosis_l3_l4',\n       'spinal_canal_stenosis_l4_l5', 'spinal_canal_stenosis_l5_s1',\n       'left_neural_foraminal_narrowing_l1_l2',\n       'left_neural_foraminal_narrowing_l2_l3',\n       'left_neural_foraminal_narrowing_l3_l4',\n       'left_neural_foraminal_narrowing_l4_l5',\n       'left_neural_foraminal_narrowing_l5_s1',\n       'right_neural_foraminal_narrowing_l1_l2',\n       'right_neural_foraminal_narrowing_l2_l3',\n       'right_neural_foraminal_narrowing_l3_l4',\n       'right_neural_foraminal_narrowing_l4_l5',\n       'right_neural_foraminal_narrowing_l5_s1',\n       'left_subarticular_stenosis_l1_l2', 'left_subarticular_stenosis_l2_l3',\n       'left_subarticular_stenosis_l3_l4', 'left_subarticular_stenosis_l4_l5',\n       'left_subarticular_stenosis_l5_s1', 'right_subarticular_stenosis_l1_l2',\n       'right_subarticular_stenosis_l2_l3',\n       'right_subarticular_stenosis_l3_l4',\n       'right_subarticular_stenosis_l4_l5',\n       'right_subarticular_stenosis_l5_s1']\n    label_map = {\n        'Normal/Mild': 0,\n        'Moderate': 1,\n        'Severe': 2,\n        -100: -100 # for missing values\n    }\n    use_n_images: int = 18\n\n@dataclass\nclass ModelConfig:\n    name: str = 'efficientnet_b3.ra2_in1k'\n    train_backbone: bool = True\n    use_weights: bool = True\n\n@dataclass\nclass Config:\n    dataset: DatasetConfig = field(default_factory=DatasetConfig)\n    model: ModelConfig = field(default_factory=ModelConfig)\n    seed: int = 42\n    n_folds: int = 5\n        \nargs = TrainingArguments(\n    output_dir='/kaggle/working/',\n    do_eval=True,\n    report_to='none' if DEBUG else 'wandb',\n    remove_unused_columns=False,\n    per_device_train_batch_size=2,\n    per_device_eval_batch_size=4,\n#     torch_compile=True,\n    num_train_epochs=5,\n    warmup_ratio=0.05,\n    logging_strategy='steps',\n    logging_steps=10,\n    save_strategy='epoch',\n    load_best_model_at_end=True,\n    save_total_limit=2,\n    save_safetensors=False,\n#     fp16=True,\n    metric_for_best_model='log_loss',\n    eval_strategy='epoch',\n    dataloader_num_workers=4,\n    label_names=['labels'],\n#     optim=\"adafactor\",\n)\n\nconfig = Config()\n\nwandb_config = {\n    'config': asdict(config),\n    'training_args': asdict(args)\n}\n\nimport datetime\nNAME = datetime.datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\")","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:52.087763Z","iopub.execute_input":"2024-09-12T07:39:52.088432Z","iopub.status.idle":"2024-09-12T07:39:52.220328Z","shell.execute_reply.started":"2024-09-12T07:39:52.088394Z","shell.execute_reply":"2024-09-12T07:39:52.219278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    os.environ['WANDB_MODE'] = 'offline'\n    \nimport wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"wandb\")\n\n!wandb login $secret_value_0","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:52.221620Z","iopub.execute_input":"2024-09-12T07:39:52.221943Z","iopub.status.idle":"2024-09-12T07:39:55.294312Z","shell.execute_reply.started":"2024-09-12T07:39:52.221908Z","shell.execute_reply":"2024-09-12T07:39:55.293241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_data(path: Path):\n    sample_submission_df = pd.read_csv(path / 'sample_submission.csv')\n    test_series_descriptions_df = pd.read_csv(path / 'test_series_descriptions.csv')\n    train_df = pd.read_csv(path / 'train.csv')\n    train_label_coordinates_df = pd.read_csv(path / 'train_label_coordinates.csv')\n    train_series_descriptions_df = pd.read_csv(path / 'train_series_descriptions.csv')\n    \n    return sample_submission_df, test_series_descriptions_df, train_df, train_label_coordinates_df, train_series_descriptions_df\n\nsample_submission_df, test_series_descriptions_df, train_df, train_label_coordinates_df, train_series_descriptions_df = load_data(config.dataset.path)\n\nif DEBUG:\n    train_df = train_df.sample(100).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.296937Z","iopub.execute_input":"2024-09-12T07:39:55.297308Z","iopub.status.idle":"2024-09-12T07:39:55.451782Z","shell.execute_reply.started":"2024-09-12T07:39:55.297272Z","shell.execute_reply":"2024-09-12T07:39:55.450965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.452912Z","iopub.execute_input":"2024-09-12T07:39:55.453242Z","iopub.status.idle":"2024-09-12T07:39:55.474088Z","shell.execute_reply.started":"2024-09-12T07:39:55.453208Z","shell.execute_reply":"2024-09-12T07:39:55.473194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_series_descriptions_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.475447Z","iopub.execute_input":"2024-09-12T07:39:55.476122Z","iopub.status.idle":"2024-09-12T07:39:55.485039Z","shell.execute_reply.started":"2024-09-12T07:39:55.476076Z","shell.execute_reply":"2024-09-12T07:39:55.484005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_series_descriptions_df['series_description'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.486332Z","iopub.execute_input":"2024-09-12T07:39:55.487127Z","iopub.status.idle":"2024-09-12T07:39:55.503075Z","shell.execute_reply.started":"2024-09-12T07:39:55.487080Z","shell.execute_reply":"2024-09-12T07:39:55.502128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.504310Z","iopub.execute_input":"2024-09-12T07:39:55.505162Z","iopub.status.idle":"2024-09-12T07:39:55.525889Z","shell.execute_reply.started":"2024-09-12T07:39:55.505120Z","shell.execute_reply":"2024-09-12T07:39:55.524900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.526821Z","iopub.execute_input":"2024-09-12T07:39:55.527134Z","iopub.status.idle":"2024-09-12T07:39:55.535129Z","shell.execute_reply.started":"2024-09-12T07:39:55.527092Z","shell.execute_reply":"2024-09-12T07:39:55.534124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_label_coordinates_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.538179Z","iopub.execute_input":"2024-09-12T07:39:55.538571Z","iopub.status.idle":"2024-09-12T07:39:55.550899Z","shell.execute_reply.started":"2024-09-12T07:39:55.538523Z","shell.execute_reply":"2024-09-12T07:39:55.549907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_series_descriptions_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.552133Z","iopub.execute_input":"2024-09-12T07:39:55.552463Z","iopub.status.idle":"2024-09-12T07:39:55.562231Z","shell.execute_reply.started":"2024-09-12T07:39:55.552427Z","shell.execute_reply":"2024-09-12T07:39:55.561448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_series_descriptions_df['series_description'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.563147Z","iopub.execute_input":"2024-09-12T07:39:55.563427Z","iopub.status.idle":"2024-09-12T07:39:55.574614Z","shell.execute_reply.started":"2024-09-12T07:39:55.563396Z","shell.execute_reply":"2024-09-12T07:39:55.573699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"backbone = timm.create_model(config.model.name, pretrained=True)\n\nif not config.model.train_backbone:\n    for param in backbone.parameters():\n        param.require_grad = False\n\n# get model specific transforms (normalization, resize)\ndata_config = timm.data.resolve_model_data_config(backbone)\ntransforms = timm.data.create_transform(**data_config, is_training=False)\n\ntransforms.transforms[-1].mean = transforms.transforms[-1].mean.mean()\ntransforms.transforms[-1].std = transforms.transforms[-1].std.mean()\n\nclass Model2_5D(nn.Module):\n    def __init__(self, backbone):\n        super(Model2_5D, self).__init__()\n        self.backbone = backbone\n        self.pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.lstm = nn.LSTM(input_size=backbone.num_features * 3, hidden_size=512, num_layers=1, batch_first=True, bidirectional=True)\n        # we have 25 classification targets with 3 classes each\n        self.head = nn.Linear(1024, len(config.dataset.label_columns)*3)\n\n\n    def forward(self, **kwargs):\n        batch = dict(**kwargs)\n        # x: (B, N_IMAGES, C, H, W)\n        features_output = []\n        for series_description in ['Sagittal T2/STIR', 'Sagittal T1', 'Axial T2']:\n            if not series_description in batch:\n                # not used, we just forward zeros through backbone\n                features_output.append(None)\n            else:\n                images = batch[series_description]\n                B, N_IMAGES, C, H, W = images.shape\n                if config.model.train_backbone:\n                    features = self.backbone.forward_features(images.view(B * N_IMAGES, C, H, W))\n                else:\n                    with torch.no_grad():\n                        features = self.backbone.forward_features(images.view(B * N_IMAGES, C, H, W))\n                features = self.pool(features).squeeze(-1).squeeze(-1).view(B, N_IMAGES, -1)\n                # features: (B, N_IMAGES, backbone.num_features)\n                features_output.append(features)\n        for i in range(len(features_output)):\n            if features_output[i] is None:\n                features_output[i] = torch.zeros((B, N_IMAGES, self.backbone.num_features), device=images.device)\n        features_output = torch.cat(features_output, -1)\n        lstm_out, _ = self.lstm(features_output)\n        # lstm_out: (B, N_IMAGES, 1024)\n        lstm_out = lstm_out.mean(1)\n        # lstm_out: (B, 1024)\n        output = self.head(lstm_out).view(-1, 25, 3)\n            \n        loss = None\n        if 'labels' in batch:\n            if config.model.use_weights:\n                loss_fn = torch.nn.CrossEntropyLoss(torch.tensor([1., 2., 4.], device=output.device))\n            else:\n                loss_fn = torch.nn.CrossEntropyLoss()\n            loss = loss_fn(output.view(-1, 3), batch['labels'].view(-1, ))\n        \n        return loss, output","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:55.577641Z","iopub.execute_input":"2024-09-12T07:39:55.577909Z","iopub.status.idle":"2024-09-12T07:39:56.294670Z","shell.execute_reply.started":"2024-09-12T07:39:55.577880Z","shell.execute_reply":"2024-09-12T07:39:56.293702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_WIDTH, IMG_HEIGHT = transforms.transforms[1].size","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.296194Z","iopub.execute_input":"2024-09-12T07:39:56.296769Z","iopub.status.idle":"2024-09-12T07:39:56.302204Z","shell.execute_reply.started":"2024-09-12T07:39:56.296727Z","shell.execute_reply":"2024-09-12T07:39:56.301329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass ImagesDataset():\n    study_ids: List[int]\n    series_descriptions_df: pd.DataFrame\n    path: Path\n    is_train: bool\n    use_n_images: int\n    transforms: A.Compose | None = None\n    labels: pd.DataFrame | None = None\n        \n    def __post_init__(self):\n        if self.use_n_images % 3 != 0:\n            raise ValueError(\"use_n_images must be divisible by 3!\")\n        if self.transforms is None:\n            self.transforms = lambda x: x\n            \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def __getitem__(self, index: int):\n        study_id = self.study_ids[index]\n        \n        return_dict = {}\n\n        for series_description in ['Sagittal T2/STIR', 'Sagittal T1', 'Axial T2']:\n            rows = self.series_descriptions_df[\n                (self.series_descriptions_df.study_id == study_id)\n                & (self.series_descriptions_df.series_description == series_description)\n            ]\n\n            if len(rows) == 0:\n                return_dict[series_description] = torch.zeros(self.use_n_images // 3, 3, IMG_WIDTH, IMG_HEIGHT)\n            else:\n                row = rows.iloc[0]\n                series_id = row.series_id\n                use_path = self.path / (\"train_images\" if self.is_train else \"test_images\") / f\"{study_id}/{series_id}\"\n                image_paths = sorted(glob.glob(str(use_path / \"*.dcm\")))\n                image_paths = random.sample(image_paths, min(self.use_n_images, len(image_paths)))\n                images = [Path(image_path).dcmread() for image_path in image_paths]\n                images = [self.transforms(image.windowed(*dicom_windows.spine_bone).unsqueeze(0))[0] for image in images]\n                if len(images) < self.use_n_images:\n                    images.extend([torch.zeros(IMG_WIDTH, IMG_HEIGHT) for _ in range(self.use_n_images - len(images))])\n                images = torch.stack(images, 0)\n                images = images.reshape(self.use_n_images // 3, 3, images[0].shape[0], images[0].shape[1])\n                return_dict[series_description] = images\n\n        if self.labels is not None:\n            return_dict['labels'] = torch.tensor(self.labels.iloc[index].values, dtype=torch.long)\n\n        return return_dict","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.303522Z","iopub.execute_input":"2024-09-12T07:39:56.303837Z","iopub.status.idle":"2024-09-12T07:39:56.319738Z","shell.execute_reply.started":"2024-09-12T07:39:56.303802Z","shell.execute_reply":"2024-09-12T07:39:56.318938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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')['severe'].max())\n    any_severe_spinal_weights = pd.Series(solution.loc[solution['condition'] == 'spinal'].groupby('study_id')['sample_weight'].max())\n    any_severe_spinal_predictions = pd.Series(submission.loc[submission['condition'] == 'spinal'].groupby('study_id')['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":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.320965Z","iopub.execute_input":"2024-09-12T07:39:56.321290Z","iopub.status.idle":"2024-09-12T07:39:56.340753Z","shell.execute_reply.started":"2024-09-12T07:39:56.321255Z","shell.execute_reply":"2024-09-12T07:39:56.339885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_df, valid_df\ncv = GroupKFold(n_splits=config.n_folds)\ntrain_df['fold'] = -1\n\nfor i, (train_idx, valid_idx) in enumerate(cv.split(train_df, groups=train_df['study_id'])):\n    train_df.loc[valid_idx, 'fold'] = i\n    \ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.342092Z","iopub.execute_input":"2024-09-12T07:39:56.342427Z","iopub.status.idle":"2024-09-12T07:39:56.393342Z","shell.execute_reply.started":"2024-09-12T07:39:56.342391Z","shell.execute_reply":"2024-09-12T07:39:56.392449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import itertools\n\ndef build_gt_submission_df(df: pd.DataFrame) -> pd.DataFrame:\n \n    df = df.drop('fold', axis=1).set_index('study_id').T\n    rows = itertools.product(df.columns, df.index)\n    values = df.values.flatten()\n\n    result_df = pd.DataFrame({\n        'row_id': [f'{study_id}_{location}' for study_id, location in rows],\n        'normal_mild': [1 if value == 'Normal/Mild' else 0 for value in values],\n        'moderate': [1 if value == 'Moderate' else 0 for value in values],\n        'severe': [1 if value == 'Severe' else 0 for value in values],\n    })\n    \n    result_df.loc[\n        (result_df['normal_mild'] == 0)\n        & (result_df['moderate'] == 0)\n        & (result_df['severe'] == 0),\n        'normal_mild'\n    ] = 1\n    \n    return result_df\n    ","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.394629Z","iopub.execute_input":"2024-09-12T07:39:56.395472Z","iopub.status.idle":"2024-09-12T07:39:56.402368Z","shell.execute_reply.started":"2024-09-12T07:39:56.395438Z","shell.execute_reply":"2024-09-12T07:39:56.401436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = logging.getLogger(__name__)\n\ndef prepare_dataset(df, series_descriptions_df, config, is_train=True):\n    dataset = ImagesDataset(\n        study_ids=df.study_id.values.tolist(),\n        path=config.dataset.path,\n        is_train=is_train,\n        series_descriptions_df=series_descriptions_df,\n        use_n_images=config.dataset.use_n_images,\n        transforms=transforms,\n        labels=df[config.dataset.label_columns].fillna(-100).apply(lambda xx: [config.dataset.label_map[x] for x in xx], raw=True)\n    )\n    gt_submission = build_gt_submission_df(df)\n    gt_submission['sample_weight'] = gt_submission.apply(lambda row: row['normal_mild'] + row['moderate'] * 2 + row['severe'] * 4, axis=1)\n    return dataset, gt_submission\n\ndef softmax(x):\n    \"\"\"Compute softmax values for each sets of scores in x.\"\"\"\n    e_x = np.exp(x - np.max(x, 1, keepdims=True))\n    return e_x / e_x.sum(1, keepdims=True)\n\ndef compute_metrics(eval_pred, valid_gt_submission):\n    logits, _ = eval_pred\n    logits = softmax(logits)\n\n    submission_df = pd.DataFrame({\n        'row_id': valid_gt_submission['row_id'].values.tolist(),\n        'normal_mild': logits[:, :, 0].flatten(),\n        'moderate': logits[:, :, 1].flatten(),\n        'severe': logits[:, :, 2].flatten(),\n    })\n\n    return {\n        'log_loss': score(valid_gt_submission.copy(), submission_df.copy(), 'row_id', 1.0)\n    }\n\ndef train_fold(fold, train_df, train_series_descriptions_df, config, args):\n    logger.info(f\"Training Fold {fold}\")\n    \n    wandb_config['fold'] = fold\n    wandb.init(\n        project = 'RSNA 2024 Lumbar Spine Degenerative Classification',\n        config = wandb_config,\n        name = NAME,\n        tags = [],\n    )\n\n    fold_train_df = train_df[train_df.fold != fold].reset_index(drop=True)\n    fold_valid_df = train_df[train_df.fold == fold].reset_index(drop=True)\n\n    fold_train_series_descriptions_df = train_series_descriptions_df[train_series_descriptions_df.study_id.isin(fold_train_df.study_id)]\n    fold_valid_series_descriptions_df = train_series_descriptions_df[train_series_descriptions_df.study_id.isin(fold_valid_df.study_id)]\n\n    train_dataset, _ = prepare_dataset(fold_train_df, fold_train_series_descriptions_df, config)\n    valid_dataset, valid_gt_submission = prepare_dataset(fold_valid_df, fold_valid_series_descriptions_df, config)\n\n    model = Model2_5D(backbone)\n\n    trainer = Trainer(\n        model,\n        args=args,\n        train_dataset=train_dataset,\n        eval_dataset=valid_dataset,\n        compute_metrics=lambda eval_pred: compute_metrics(eval_pred, valid_gt_submission),\n    )\n\n    trainer.train()\n    \n    torch.save(model.state_dict(), f\"best_model_fold_{fold}.pt\")\n    \n    wandb.finish()\n\nfor fold in range(1):\n    train_fold(fold, train_df, train_series_descriptions_df, config, args)","metadata":{"execution":{"iopub.status.busy":"2024-09-12T07:39:56.403772Z","iopub.execute_input":"2024-09-12T07:39:56.404396Z"},"trusted":true},"execution_count":null,"outputs":[]}]}