{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30761,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport re\nimport glob\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\nimport torch\nimport timm\nimport albumentations as A\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import AdamW\nfrom torch.cuda.amp import autocast\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold\n\n# Configuration class\nclass Config:\n    RD = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    OUTPUT_DIR = '/kaggle/working/models'\n    DEVICE = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')\n    N_WORKERS = os.cpu_count()\n    USE_AMP = True\n    SEED = 8620\n    IMG_SIZE = [512, 512]\n    IN_CHANS = 42\n    N_LABELS = 25\n    N_CLASSES = 3 * N_LABELS\n    N_FOLDS = 5\n    MODEL_NAME = \"edgenext_base.in21k_ft_in1k\"\n    BATCH_SIZE = 1\n    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    LEVELS = [\n        'l1_l2',\n        'l2_l3',\n        'l3_l4',\n        'l4_l5',\n        'l5_s1',\n    ]\n\n# Utility functions\ndef atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    return [atoi(c) for c in re.split(r'(\\d+)', text)]\n\n# Dataset class\nclass RSNA24TestDataset(Dataset):\n    def __init__(self, df, study_ids, transform=None):\n        self.df = df\n        self.study_ids = study_ids\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.df[self.df['study_id'] == study_id]\n        pdf_ = pdf[pdf['series_description'] == series_desc]\n        allimgs = []\n        for _, row in pdf_.iterrows():\n            pimgs = glob.glob(f'{Config.RD}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm')\n            pimgs = sorted(pimgs, key=natural_keys)\n            allimgs.extend(pimgs)\n        return allimgs\n    \n    def read_dcm_ret_arr(self, src_path):\n        dicom_data = pydicom.dcmread(src_path)\n        image = dicom_data.pixel_array\n        image = (image - image.min()) / (image.max() - image.min() + 1e-6) * 255\n        img = cv2.resize(image, (Config.IMG_SIZE[0], Config.IMG_SIZE[1]), interpolation=cv2.INTER_CUBIC)\n        assert img.shape == (Config.IMG_SIZE[0], Config.IMG_SIZE[1])\n        return img\n\n    def __getitem__(self, idx):\n        x = np.zeros((Config.IMG_SIZE[0], Config.IMG_SIZE[1], Config.IN_CHANS), dtype=np.uint8)\n        st_id = self.study_ids[idx]\n        \n        def process_series(series_desc, start_idx):\n            allimgs = self.get_img_paths(st_id, series_desc)\n            if len(allimgs) == 0:\n                print(f'{st_id}: {series_desc}, has no images')\n                return start_idx\n            step = len(allimgs) / 14.0\n            st = len(allimgs) / 2.0 - 6.0 * step\n            end = len(allimgs) + 0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i - 0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs[ind2])\n                    x[..., start_idx + j] = img.astype(np.uint8)\n                except:\n                    print(f'Failed to load on {st_id}, {series_desc}')\n            return start_idx + 14\n\n        start_idx = process_series('Sagittal T1', 0)\n        start_idx = process_series('Sagittal T2/STIR', start_idx)\n        process_series('Axial T2', start_idx)\n\n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n        \n        x = x.transpose(2, 0, 1)\n        return x, str(st_id)\n\n# Model class\nclass RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=42, n_classes=75, pretrained=True, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            features_only=features_only,\n            in_chans=in_c,\n            num_classes=n_classes,\n            global_pool='avg'\n        )\n    \n    def forward(self, x):\n        return self.model(x)\n\n# Training function\ndef train_and_save_model(fold, train_loader):\n    model = RSNA24Model(Config.MODEL_NAME, Config.IN_CHANS, Config.N_CLASSES, pretrained=True)\n    model.to(Config.DEVICE)\n    \n    optimizer = AdamW(model.parameters(), lr=1e-4)\n    criterion = nn.CrossEntropyLoss()\n    \n    model.train()\n    for epoch in range(5):  # Assume 5 epochs for simplicity\n        for x, _ in train_loader:\n            optimizer.zero_grad()\n            x = x.to(Config.DEVICE)\n            outputs = model(x)\n            loss = criterion(outputs, torch.zeros_like(outputs))  # Dummy target\n            loss.backward()\n            optimizer.step()\n    \n    torch.save(model.state_dict(), f'{Config.OUTPUT_DIR}/model_fold-{fold}.pt')\n    print(f'Model for fold {fold} saved!')\n\n# Loading models\ndef load_models(model_paths):\n    models = []\n    for path in model_paths:\n        print(f'Loading {path}...')\n        model = RSNA24Model(Config.MODEL_NAME, Config.IN_CHANS, Config.N_CLASSES, pretrained=False)\n        model.load_state_dict(torch.load(path))\n        model.eval().half().to(Config.DEVICE)\n        models.append(model)\n    return models\n\n# Prediction\ndef predict(models, test_dl):\n    autocast_enabled = Config.USE_AMP and torch.cuda.is_available()\n    y_preds = []\n    row_names = []\n\n    with tqdm(test_dl, leave=True) as pbar:\n        with torch.no_grad():\n            for x, si in pbar:\n                x = x.to(Config.DEVICE)\n                pred_per_study = np.zeros((Config.N_LABELS, 3))\n\n                for cond in Config.CONDITIONS:\n                    for level in Config.LEVELS:\n                        row_names.append(f'{si[0]}_{cond}_{level}')\n                \n                with autocast(enabled=autocast_enabled, dtype=torch.half):\n                    for model in models:\n                        y = model(x)[0]\n                        for col in range(Config.N_LABELS):\n                            pred = y[col*3:col*3+3]\n                            y_pred = pred.float().softmax(0).cpu().numpy()\n                            pred_per_study[col] += y_pred / len(models)\n                    y_preds.append(pred_per_study)\n\n    y_preds = np.concatenate(y_preds, axis=0)\n    sub = pd.DataFrame()\n    sub['row_id'] = row_names\n    sub[Config.LABELS] = y_preds\n    return sub\n\n# Main execution\ndef main():\n    df = pd.read_csv(f'{Config.RD}/test_series_descriptions.csv')\n    study_ids = list(df['study_id'].unique())\n    sample_sub = pd.read_csv(f'{Config.RD}/sample_submission.csv')\n    Config.LABELS = list(sample_sub.columns[1:])\n\n    transforms_test = A.Compose([\n        A.Resize(Config.IMG_SIZE[0], Config.IMG_SIZE[1]),\n        A.Normalize(mean=0.5, std=0.5)\n    ])\n    \n    test_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\n    test_dl = DataLoader(\n        test_ds,\n        batch_size=Config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=Config.N_WORKERS,\n        pin_memory=True,\n        drop_last=False\n    )\n\n    # Check if models exist; if not, train and save them\n    if not os.path.exists(Config.OUTPUT_DIR):\n        os.makedirs(Config.OUTPUT_DIR)\n\n    model_paths = [\n        f\"{Config.OUTPUT_DIR}/model_fold-{i}.pt\" for i in range(3)\n    ]\n    \n    # Dummy train data loader\n    train_dl = DataLoader(\n        test_ds,\n        batch_size=Config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=Config.N_WORKERS,\n        pin_memory=True,\n        drop_last=False\n    )\n\n    for fold, model_path in enumerate(model_paths):\n        if not os.path.exists(model_path):\n            print(f\"Training and saving model for fold {fold}...\")\n            train_and_save_model(fold, train_dl)\n        else:\n            print(f\"Model for fold {fold} already exists. Skipping training...\")\n\n    models = load_models(model_paths)\n    submission = predict(models, test_dl)\n    submission.to_csv('/kaggle/working/submission.csv', index=False)\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"execution":{"iopub.status.busy":"2024-09-04T17:01:34.568268Z","iopub.execute_input":"2024-09-04T17:01:34.568729Z","iopub.status.idle":"2024-09-04T17:01:38.726688Z","shell.execute_reply.started":"2024-09-04T17:01:34.568683Z","shell.execute_reply":"2024-09-04T17:01:38.724858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}