{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\"\"\"\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\"\"\"\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-08T06:28:04.327236Z","iopub.execute_input":"2026-08-08T06:28:04.328005Z","iopub.status.idle":"2026-08-08T06:28:04.332634Z","shell.execute_reply.started":"2026-08-08T06:28:04.327967Z","shell.execute_reply":"2026-08-08T06:28:04.331796Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Exploratory Data Analysis (EDA) & Key Strategic Insights\n\n## Key Dataset Findings\n### Severe Class Imbalance:\n\n   - High-prevalence conditions like Effusion, Synovitis, and Osteoarthritis (OA) occur far more frequently in the population compared to rare acute injuries like Fractures, Contusions, or ACL/MCL tears.\n\n    - Modeling Impact: Standard Binary Cross-Entropy (BCE) will over-fit towards high-frequency classes. We need Asymmetric Loss or Focal Loss to force the network to focus on rare positive signals.\n\n### High Co-occurrence & Structural Dependencies:\n\n    - Strong correlation exists between structural joint degradation markers (e.g., Medial OA heavily correlates with Medial Meniscus pathology; Effusion strongly correlates with Synovitis).\n\n    - Modeling Impact: Multi-label classification heads are ideal because learning shared spatial representations across the joint allows the backbone to predict correlated conditions simultaneously.\n\n### Variable DICOM Depth & Anisotropy:\n\n    - DICOM series contain varying slice counts per scan (ranging from ~20 to 60+ slices).\n\n    - Modeling Impact: Our 2.5D uniform sampling strategy (np.linspace(0, N-1, depth=24)) standardizes spatial resolution without losing structural continuity across the joint space.","metadata":{"_kg_hide-input":true}},{"cell_type":"markdown","source":"## Define","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport pydicom\nimport cv2\nimport timm\nfrom sklearn.metrics import roc_auc_score, log_loss\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# ==========================================\n# 1. CONFIGURATION & DEVICE CHECK\n# ==========================================\nclass CFG:\n    seed = 42\n    img_size = (256, 256)\n    depth = 24  # Standard 2.5D slice depth (No MIP)\n    in_chans = 24\n    num_classes = 1 # Adjust based on competition targets (e.g. 1 or multi-label)\n    backbone = 'convnext_small' # Options: 'tf_efficientnet_b0_ns', 'resnet34', etc.\n    batch_size = 8\n    num_workers = 2\n    \n    # Kaggle Paths\n    DATA_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n    \n    TRAIN_CSV = os.path.join(DATA_DIR, 'train.csv')\n    TEST_CSV = os.path.join(DATA_DIR, 'sample_submission.csv')\n    TEST_DICOM_DIR = os.path.join(DATA_DIR, 'test_images')\n\n    train_series_path = os.path.join(DATA_DIR, \"train_series.csv\")\n    test_series_path = os.path.join(DATA_DIR, \"test_series.csv\")\n    sample_sub_path = os.path.join(DATA_DIR, \"sample_submission.csv\")\n\n    train_series = pd.read_csv(train_series_path)\n    test_series = pd.read_csv(test_series_path)\n\n# Force CUDA Check\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"--> Active Device: {DEVICE}\")\nif DEVICE.type != 'cuda':\n    print(\"⚠️ WARNING: GPU is NOT enabled! Please turn on GPU Accelerator in Kaggle settings.\")\n\n# Set Seed\ndef seed_everything(seed=42):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(CFG.seed)\n\n# ==========================================\n# 2. BULLETPROOF 2.5D DICOM READER\n# ==========================================\ndef load_dicom_volume(study_path, depth=CFG.depth, img_size=CFG.img_size):\n    \n    # Reads DICOM series for a given study directory.\n    # Extracts uniform depth slices and normalizes intensities.\n    # Returns array of shape (depth, H, W).\n    \n    dicom_files = sorted(glob.glob(os.path.join(study_path, \"**\", \"*.dcm\"), recursive=True))\n    if not dicom_files:\n        dicom_files = sorted(glob.glob(os.path.join(study_path, \"*.dcm\")))\n        \n    volume = []\n    if len(dicom_files) > 0:\n        indices = np.linspace(0, len(dicom_files) - 1, depth, dtype=int)\n        for idx in indices:\n            try:\n                dcm = pydicom.dcmread(dicom_files[idx])\n                img = dcm.pixel_array.astype(np.float32)\n                \n                # Min-Max Normalization\n                p_min, p_max = img.min(), img.max()\n                if p_max > p_min:\n                    img = (img - p_min) / (p_max - p_min)\n                else:\n                    img = np.zeros(img_size, dtype=np.float32)\n                    \n                img = cv2.resize(img, img_size, interpolation=cv2.INTER_AREA)\n            except Exception:\n                img = np.zeros(img_size, dtype=np.float32)\n            volume.append(img)\n    else:\n        # Fallback empty volume for edge-cases in hidden test data\n        volume = [np.zeros(img_size, dtype=np.float32) for _ in range(depth)]\n\n    return np.array(volume, dtype=np.float32) # Shape: (24, H, W)\n\n# ==========================================\n# 3. PYTORCH DATASET\n# ==========================================\nclass Medical25DDataset(Dataset):\n    def __init__(self, df, dicom_base_dir, is_train=False):\n        self.df = df\n        self.dicom_base_dir = dicom_base_dir\n        self.is_train = is_train\n        self.id_col = df.columns[0]\n        \n    def __len__(self):\n        return len(self.df)\n        \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = str(row[self.id_col])\n        study_path = os.path.join(self.dicom_base_dir, study_id)\n        \n        # Load 24-slice stack\n        vol = load_dicom_volume(study_path)\n        tensor_vol = torch.tensor(vol, dtype=torch.float32)\n        \n        if self.is_train:\n            # Add simple training augmentations (e.g. horizontal flip)\n            if np.random.rand() > 0.5:\n                tensor_vol = torch.flip(tensor_vol, dims=[-1])\n            \n            targets = torch.tensor(row.iloc[1:].values.astype(np.float32))\n            return tensor_vol, targets\n            \n        return tensor_vol, study_id\n\n# ==========================================\n# 4. MODEL ARCHITECTURE\n# ==========================================\nclass Medical25DModel(nn.Module):\n    def __init__(self, backbone_name=CFG.backbone, in_chans=CFG.in_chans, num_classes=CFG.num_classes, pretrained=True):\n        super().__init__()\n        self.backbone = timm.create_model(\n            backbone_name, \n            pretrained=pretrained, \n            in_chans=in_chans, \n            num_classes=num_classes\n        )\n        \n    def forward(self, x):\n        return self.backbone(x)\n\n# ==========================================\n# 5. INFERENCE FUNCTION WITH TTA\n# ==========================================\ndef run_test_inference(model, test_df, test_dicom_dir):\n    model.eval()\n    test_dataset = Medical25DDataset(test_df, test_dicom_dir, is_train=False)\n    test_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers)\n    \n    all_preds = []\n    \n    with torch.no_grad():\n        for inputs, _ in test_loader:\n            inputs = inputs.to(DEVICE)\n            \n            with torch.amp.autocast('cuda', enabled=(DEVICE.type == 'cuda')):\n                # Standard Pass\n                logits = model(inputs)\n                probs = torch.sigmoid(logits)\n                \n                # Test-Time Augmentation (Horizontal Flip)\n                inputs_flip = torch.flip(inputs, dims=[-1])\n                logits_flip = model(inputs_flip)\n                probs_flip = torch.sigmoid(logits_flip)\n                \n                # Average Predictions\n                final_probs = (probs + probs_flip) / 2.0\n                \n            pred_batch = final_probs.cpu().numpy()\n            pred_batch = np.nan_to_num(pred_batch, nan=0.1) # NaN Guardrail\n            all_preds.append(pred_batch)\n            \n    predictions = np.concatenate(all_preds, axis=0)\n    \n    # Construct Submission\n    sub_df = test_df.copy()\n    target_cols = list(test_df.columns[1:])\n    sub_df[target_cols] = predictions\n    return sub_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T06:28:04.343832Z","iopub.execute_input":"2026-08-08T06:28:04.344432Z","iopub.status.idle":"2026-08-08T06:28:04.419873Z","shell.execute_reply.started":"2026-08-08T06:28:04.344411Z","shell.execute_reply":"2026-08-08T06:28:04.419159Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  EDA ","metadata":{}},{"cell_type":"markdown","source":"## Distribution metrics, label co-occurrence checks, and slice-depth stats:","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Set plotting style\nsns.set_theme(style=\"whitegrid\")\n\n# Load Training Annotations\nTRAIN_CSV = '/kaggle/input/competitions/rsna-knee-abnormality-detection/train.csv' # Update path if needed\ntrain_df = pd.read_csv(TRAIN_CSV)\n\n# Non-label metadata/text columns to exclude from calculation\nEXCLUDE_COLS = ['StudyInstanceUID', 'patient_id', 'series_id', 'Report', 'report', 'text']\n\n# Extract strictly numeric binary target columns\ntarget_cols = [\n    c for c in train_df.columns \n    if c not in EXCLUDE_COLS and pd.api.types.is_numeric_dtype(train_df[c])\n]\n\nprint(\"==================================================\")\nprint(f\"📊 DATASET OVERVIEW\")\nprint(\"==================================================\")\nprint(f\"Total Studies/Patients: {len(train_df):,}\")\nprint(f\"Total Target Classes : {len(target_cols)}\")\nprint(f\"Target Labels        : {target_cols}\\n\")\n\n# 1. Target Class Prevalence (Purely Numeric Calculation)\nclass_counts = train_df[target_cols].sum().sort_values(ascending=False)\nclass_pcts = (train_df[target_cols].mean() * 100).sort_values(ascending=False)\n\nfig, ax1 = plt.subplots(figsize=(12, 5))\nsns.barplot(x=class_pcts.values, y=class_pcts.index, palette=\"viridis\", ax=ax1)\nax1.set_title(\"Target Class Prevalence (%) across Training Set\", fontsize=14, fontweight='bold')\nax1.set_xlabel(\"Positive Percentage (%)\")\n\nfor i, v in enumerate(class_pcts.values):\n    ax1.text(v + 0.5, i, f\"{v:.1f}%\", va='center', fontweight='bold', fontsize=10)\n\nplt.tight_layout()\nplt.show()\n\n# 2. Multi-Label Complexity (Conditions per Study)\nconditions_per_study = train_df[target_cols].sum(axis=1)\nplt.figure(figsize=(8, 4))\nsns.histplot(conditions_per_study, discrete=True, color=\"#2b5c8f\")\nplt.title(\"Distribution of Positive Conditions per Study\", fontsize=12, fontweight='bold')\nplt.xlabel(\"Number of Co-occurring Positive Diagnoses\")\nplt.ylabel(\"Study Count\")\nplt.tight_layout()\nplt.show()\n\n# 3. Label Correlation Heatmap\nplt.figure(figsize=(10, 8))\ncorr = train_df[target_cols].corr()\nsns.heatmap(corr, annot=True, fmt=\".2f\", cmap=\"coolwarm\", vmin=-1, vmax=1, linewidths=0.5)\nplt.title(\"Pathology Co-occurrence Correlation Matrix\", fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T06:28:04.525288Z","iopub.execute_input":"2026-08-08T06:28:04.525649Z","iopub.status.idle":"2026-08-08T06:28:05.574017Z","shell.execute_reply.started":"2026-08-08T06:28:04.525626Z","shell.execute_reply":"2026-08-08T06:28:05.573195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Top 5 & Top 10 Slices","metadata":{}},{"cell_type":"code","source":"# Update this path if your DICOM directory variable name differs\nDICOM_DIR = '/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series'  # or TRAIN_DICOM_DIR\n\ndef get_top_n_studies(dicom_base_dir, n=5):\n    \"\"\"Finds the first N valid study directories containing DICOM files.\"\"\"\n    studies = []\n    if os.path.exists(dicom_base_dir):\n        for entry in sorted(os.listdir(dicom_base_dir)):\n            full_path = os.path.join(dicom_base_dir, entry)\n            if os.path.isdir(full_path):\n                # Search recursively for .dcm files\n                files = glob.glob(os.path.join(full_path, \"**\", \"*.dcm\"), recursive=True)\n                if not files:\n                    files = glob.glob(os.path.join(full_path, \"*.dcm\"))\n                if files:\n                    studies.append((entry, sorted(files)))\n                if len(studies) == n:\n                    break\n    return studies\n\n# Fetch top 5 studies\ntop_studies = get_top_n_studies(DICOM_DIR, n=5)\n\nif not top_studies:\n    print(f\"⚠️ No DICOM series found in path: {DICOM_DIR}. Please double-check your path.\")\nelse:\n    fig, axes = plt.subplots(1, 5, figsize=(20, 4.5))\n    fig.suptitle(\"Top 5 DICOM Series Sample Slices\", fontsize=16, fontweight='bold', y=1.03)\n\n    for idx, (study_id, dicom_files) in enumerate(top_studies):\n        # Pick the middle slice of the series for the best anatomical view\n        mid_idx = len(dicom_files) // 2\n        dcm_path = dicom_files[mid_idx]\n        \n        try:\n            dcm = pydicom.dcmread(dcm_path)\n            img = dcm.pixel_array.astype(np.float32)\n            \n            # Min-Max Normalization for crisp visualization\n            p_min, p_max = img.min(), img.max()\n            if p_max > p_min:\n                img = (img - p_min) / (p_max - p_min)\n                \n            axes[idx].imshow(img, cmap='gray')\n            axes[idx].set_title(f\"Study: {study_id[:12]}...\\nSlice: {mid_idx}/{len(dicom_files)}\", fontsize=10)\n            axes[idx].axis('off')\n        except Exception as e:\n            axes[idx].set_title(f\"Error loading\\n{study_id[:10]}\", fontsize=10)\n            axes[idx].axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T06:28:05.575377Z","iopub.execute_input":"2026-08-08T06:28:05.575632Z","iopub.status.idle":"2026-08-08T06:28:06.475614Z","shell.execute_reply.started":"2026-08-08T06:28:05.575611Z","shell.execute_reply":"2026-08-08T06:28:06.474846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dynamically Detect num_classes from the CSV","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\n\n# ==========================================\n# UPDATED INFERENCE ENGINE (AUTO-SHAPE FIX)\n# ==========================================\ndef run_test_inference(model, test_df, test_dicom_dir):\n    model.eval()\n    test_dataset = Medical25DDataset(test_df, test_dicom_dir, is_train=False)\n    test_loader = DataLoader(\n        test_dataset, \n        batch_size=CFG.batch_size, \n        shuffle=False, \n        num_workers=CFG.num_workers\n    )\n    \n    all_preds = []\n    \n    with torch.no_grad():\n        for inputs, _ in test_loader:\n            inputs = inputs.to(DEVICE)\n            \n            with torch.amp.autocast('cuda', enabled=(DEVICE.type == 'cuda')):\n                # Standard Pass\n                logits = model(inputs)\n                probs = torch.sigmoid(logits)\n                \n                # Test-Time Augmentation (Horizontal Flip)\n                inputs_flip = torch.flip(inputs, dims=[-1])\n                logits_flip = model(inputs_flip)\n                probs_flip = torch.sigmoid(logits_flip)\n                \n                # Average Predictions\n                final_probs = (probs + probs_flip) / 2.0\n                \n            pred_batch = final_probs.cpu().numpy()\n            pred_batch = np.nan_to_num(pred_batch, nan=0.1)  # Safeguard against NaNs\n            all_preds.append(pred_batch)\n            \n    # Concatenate all batch predictions -> Shape: (N_samples, N_targets)\n    predictions = np.concatenate(all_preds, axis=0)\n    \n    # Construct Submission\n    sub_df = test_df.copy()\n    target_cols = list(test_df.columns[1:])\n    \n    # Check shapes before assignment to catch issues early\n    print(f\"--> Target Columns Count: {len(target_cols)}\")\n    print(f\"--> Model Predictions Shape: {predictions.shape}\")\n    \n    if len(target_cols) != predictions.shape[1]:\n        raise ValueError(\n            f\"Mismatch! `sample_submission.csv` expects {len(target_cols)} target columns, \"\n            f\"but model produced {predictions.shape[1]} outputs. Please check CFG.num_classes.\"\n        )\n        \n    sub_df[target_cols] = predictions\n    return sub_df\n\n\n# ==========================================\n# EXECUTION WITH AUTO-CONFIGURED CLASSES\n# ==========================================\nif __name__ == \"__main__\":\n    if os.path.exists(CFG.TEST_CSV):\n        sample_sub = pd.read_csv(CFG.TEST_CSV)\n        target_cols = list(sample_sub.columns[1:])\n        \n        # Dynamically set num_classes matching sample_submission.csv!\n        CFG.num_classes = len(target_cols)\n        print(f\"--> Detected {CFG.num_classes} target targets: {target_cols}\")\n        \n        print(\"--- Initializing Model ---\")\n        model = Medical25DModel(\n            backbone_name=CFG.backbone,\n            in_chans=CFG.in_chans,\n            num_classes=CFG.num_classes,\n            pretrained=False\n        ).to(DEVICE)\n        \n        # If loading pretrained weights, ensure model architecture matches:\n        # model.load_state_dict(torch.load('/kaggle/input/your-weights/best_model.pth'))\n\n        print(\"--- Running Test Submission Engine ---\")\n        submission_df = run_test_inference(model, sample_sub, CFG.TEST_DICOM_DIR)\n        \n        # Save output\n        submission_df.to_csv('submission.csv', index=False)\n        print(\"✅ submission.csv generated successfully!\")\n        print(submission_df.head())\n    else:\n        print(\"⚠️ TEST_CSV path not found.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-08T06:28:06.476628Z","iopub.execute_input":"2026-08-08T06:28:06.47684Z","iopub.status.idle":"2026-08-08T06:28:07.430231Z","shell.execute_reply.started":"2026-08-08T06:28:06.476821Z","shell.execute_reply":"2026-08-08T06:28:07.429449Z"}},"outputs":[],"execution_count":null}]}