{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.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"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# Path where Kaggle mounts the competition dataset\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# List what’s inside\nprint(os.listdir(data_path))\n\n# Load train.csv\ntrain_df = pd.read_csv(os.path.join(data_path, \"train.csv\"))\n\ntrain_df.shape\ntrain_df.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:20.724175Z","iopub.execute_input":"2025-10-03T13:12:20.724715Z","iopub.status.idle":"2025-10-03T13:12:22.135235Z","shell.execute_reply.started":"2025-10-03T13:12:20.724688Z","shell.execute_reply":"2025-10-03T13:12:22.134512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label_coordinates = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.135967Z","iopub.execute_input":"2025-10-03T13:12:22.136283Z","iopub.status.idle":"2025-10-03T13:12:22.302383Z","shell.execute_reply.started":"2025-10-03T13:12:22.136242Z","shell.execute_reply":"2025-10-03T13:12:22.301625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_label_coordinates","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.303854Z","iopub.execute_input":"2025-10-03T13:12:22.304066Z","iopub.status.idle":"2025-10-03T13:12:22.319676Z","shell.execute_reply.started":"2025-10-03T13:12:22.304048Z","shell.execute_reply":"2025-10-03T13:12:22.318737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndata_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\"\n\n# list a few study folders\nprint(os.listdir(os.path.join(data_path, \"train_images\"))[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.320412Z","iopub.execute_input":"2025-10-03T13:12:22.320837Z","iopub.status.idle":"2025-10-03T13:12:22.359996Z","shell.execute_reply.started":"2025-10-03T13:12:22.320815Z","shell.execute_reply":"2025-10-03T13:12:22.359391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_id = \"4003253\"\nstudy_path = os.path.join(data_path, \"train_images\", study_id)\n\nprint(\"Series inside study:\", os.listdir(study_path)[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.360477Z","iopub.execute_input":"2025-10-03T13:12:22.360673Z","iopub.status.idle":"2025-10-03T13:12:22.375548Z","shell.execute_reply.started":"2025-10-03T13:12:22.360655Z","shell.execute_reply":"2025-10-03T13:12:22.374840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_id = os.listdir(study_path)[0]  # take first series\nseries_path = os.path.join(study_path, series_id)\n\nprint(\"DICOM slices inside series:\", os.listdir(series_path)[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.376458Z","iopub.execute_input":"2025-10-03T13:12:22.376711Z","iopub.status.idle":"2025-10-03T13:12:22.391976Z","shell.execute_reply.started":"2025-10-03T13:12:22.376689Z","shell.execute_reply":"2025-10-03T13:12:22.391343Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\ndcm_file = os.path.join(series_path, os.listdir(series_path)[0])  # first slice\ndcm = pydicom.dcmread(dcm_file)\n\nprint(dcm)  # metadata\nplt.imshow(dcm.pixel_array, cmap='gray')\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:22.392734Z","iopub.execute_input":"2025-10-03T13:12:22.392955Z","iopub.status.idle":"2025-10-03T13:12:23.217150Z","shell.execute_reply.started":"2025-10-03T13:12:22.392935Z","shell.execute_reply":"2025-10-03T13:12:23.216444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\nimport os\n\n# Pick one study and series (from your earlier code)\nstudy_id = \"4003253\"\nseries_id = os.listdir(os.path.join(data_path, \"train_images\", study_id))[0]\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices in this series and sort by filename (important!)\ndcm_files = sorted(os.listdir(series_path))\n\n# Pick 5 slices spread across the series\nsample_slices = [dcm_files[i] for i in [0, len(dcm_files)//4, len(dcm_files)//2, 3*len(dcm_files)//4, -1]]\n\n# Plot them\nfig, axes = plt.subplots(1, 5, figsize=(20, 5))\n\nfor ax, fname in zip(axes, sample_slices):\n    dcm_path = os.path.join(series_path, fname)\n    dcm = pydicom.dcmread(dcm_path)\n    ax.imshow(dcm.pixel_array, cmap='gray')\n    ax.set_title(fname)\n    ax.axis('off')\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:23.217959Z","iopub.execute_input":"2025-10-03T13:12:23.218221Z","iopub.status.idle":"2025-10-03T13:12:23.739583Z","shell.execute_reply.started":"2025-10-03T13:12:23.218197Z","shell.execute_reply":"2025-10-03T13:12:23.738509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\ncoords_path = os.path.join(data_path, \"train_label_coordinates.csv\")\ncoords_df = pd.read_csv(coords_path)\n\nprint(coords_df.shape)\nprint(coords_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:23.742487Z","iopub.execute_input":"2025-10-03T13:12:23.742736Z","iopub.status.idle":"2025-10-03T13:12:23.815396Z","shell.execute_reply.started":"2025-10-03T13:12:23.742716Z","shell.execute_reply":"2025-10-03T13:12:23.814639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\nimport os\n\n# Pick one row from coords_df\nrow = coords_df.iloc[0]\n\nstudy_id = str(row.study_id)\nseries_id = str(row.series_id)\ninstance_number = int(row.instance_number)\n\n# Build path to the series folder\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices (files) in this series\ndcm_files = sorted(os.listdir(series_path))\n\n# instance_number in metadata starts from 1, so adjust index\ndcm_file = os.path.join(series_path, dcm_files[instance_number - 1])\n\n# Load the DICOM slice\ndcm = pydicom.dcmread(dcm_file)\nimg = dcm.pixel_array\n\n# Plot image with coordinate marked\nplt.figure(figsize=(6,6))\nplt.imshow(img, cmap='gray')\nplt.scatter(row.x, row.y, c='red', s=40, label=f\"{row.condition} {row.level}\")\nplt.legend()\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:23.816217Z","iopub.execute_input":"2025-10-03T13:12:23.816520Z","iopub.status.idle":"2025-10-03T13:12:24.079389Z","shell.execute_reply.started":"2025-10-03T13:12:23.816497Z","shell.execute_reply":"2025-10-03T13:12:24.078433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\n# Pick a study to visualize\nstudy_id = \"4003253\"\n\n# Filter coords for this study\nstudy_coords = coords_df[coords_df.study_id == int(study_id)]\n\n# Pick one series from this study\nseries_id = str(study_coords.series_id.iloc[0])\nseries_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n\n# Get all slices in this series\ndcm_files = sorted(os.listdir(series_path))\n\n# Choose a slice number that has multiple labels (for demonstration)\nslice_num = study_coords.instance_number.iloc[0]\ndcm_file = os.path.join(series_path, dcm_files[slice_num - 1])\n\n# Load the slice\ndcm = pydicom.dcmread(dcm_file)\nimg = dcm.pixel_array\n\n# Plot slice with all coordinates from this slice\nplt.figure(figsize=(7,7))\nplt.imshow(img, cmap='gray')\n\nfor _, row in study_coords[study_coords.instance_number == slice_num].iterrows():\n    plt.scatter(row.x, row.y, c='red', s=40)\n    plt.text(row.x+5, row.y+5, f\"{row.level}\", color='yellow', fontsize=9)\n\nplt.title(f\"Study {study_id} - Slice {slice_num}\")\nplt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:24.080329Z","iopub.execute_input":"2025-10-03T13:12:24.080588Z","iopub.status.idle":"2025-10-03T13:12:24.410709Z","shell.execute_reply.started":"2025-10-03T13:12:24.080566Z","shell.execute_reply":"2025-10-03T13:12:24.409845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_study_ids = set(train_df.study_id.astype(str))\nimage_study_ids = set(os.listdir(os.path.join(data_path, \"train_images\")))\n\nmissing_studies = train_study_ids - image_study_ids\nextra_studies = image_study_ids - train_study_ids\n\nprint(\"Missing studies in images:\", len(missing_studies))\nprint(\"Extra studies not in train.csv:\", len(extra_studies))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:12:24.411646Z","iopub.execute_input":"2025-10-03T13:12:24.411933Z","iopub.status.idle":"2025-10-03T13:12:24.421182Z","shell.execute_reply.started":"2025-10-03T13:12:24.411911Z","shell.execute_reply":"2025-10-03T13:12:24.420453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport pydicom\nfrom glob import glob\nfrom tqdm import tqdm\n\nBASE_PATH = data_path\nLABEL_CSV = os.path.join(BASE_PATH, \"train_label_coordinates.csv\")\nTRAIN_DIR = os.path.join(BASE_PATH, \"train_images\")\n\n# Load labels\ndf = pd.read_csv(LABEL_CSV)\n\n# Use study_id + series_id from folder structure\nseries_index = {}\n\nprint(\"Indexing DICOM files...\")\ndicom_files = glob(os.path.join(TRAIN_DIR, \"**\", \"*.dcm\"), recursive=True)\n\nfor fp in tqdm(dicom_files):\n    try:\n        ds = pydicom.dcmread(fp, stop_before_pixels=True)\n        inst_no = int(ds.InstanceNumber)\n\n        # Extract study_id and series_id from folder path\n        parts = fp.split(os.sep)\n        study_id = parts[-3]\n        series_id = parts[-2]\n        key = (study_id, series_id)\n\n        if key not in series_index:\n            series_index[key] = set()\n        series_index[key].add(inst_no)\n    except Exception:\n        continue\n\nprint(f\"Indexed {len(series_index)} (study_id, series_id) pairs from {len(dicom_files)} DICOMs.\")\n\n# Now check labels against index\nmissing_series = []\nmissing_instances = []\n\nprint(\"Checking label references...\")\nfor idx, row in tqdm(df.iterrows(), total=len(df)):\n    study_id = str(row[\"study_id\"])\n    series_id = str(row[\"series_id\"])\n    inst_no = int(row[\"instance_number\"])\n    key = (study_id, series_id)\n\n    if key not in series_index:\n        missing_series.append((idx, study_id, series_id))\n    elif inst_no not in series_index[key]:\n        missing_instances.append((idx, study_id, series_id, inst_no))\n\nprint(f\"Total rows checked: {len(df)}\")\nprint(f\"Missing series: {len(missing_series)}\")\nprint(f\"Missing instances: {len(missing_instances)}\")\n\n# Save reports\npd.DataFrame(missing_series, columns=[\"row_index\",\"study_id\",\"series_id\"]).to_csv(\"/kaggle/working/missing_series.csv\", index=False)\npd.DataFrame(missing_instances, columns=[\"row_index\",\"study_id\",\"series_id\",\"instance_number\"]).to_csv(\"/kaggle/working/missing_instances.csv\", index=False)\n\nprint(\"Reports saved in /kaggle/working/: missing_series.csv, missing_instances.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:19:15.358404Z","iopub.execute_input":"2025-10-03T13:19:15.358861Z","iopub.status.idle":"2025-10-03T13:47:36.789325Z","shell.execute_reply.started":"2025-10-03T13:19:15.358838Z","shell.execute_reply":"2025-10-03T13:47:36.788370Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport cv2\nfrom tqdm import tqdm\nimport pandas as pd\n\n# Parameters\nPATCH_SIZE = 128\nOUTPUT_DIR = \"/kaggle/working/patches\"\n\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# Load label coordinates\ncoords_df = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))\n\ndef safe_crop(img, center_x, center_y, size=PATCH_SIZE):\n    h, w = img.shape\n    half = size // 2\n    x1, x2 = center_x - half, center_x + half\n    y1, y2 = center_y - half, center_y + half\n    patch = np.zeros((size, size), dtype=img.dtype)\n    x1_img, x2_img = max(0, x1), min(w, x2)\n    y1_img, y2_img = max(0, y1), min(h, y2)\n    x1_patch, y1_patch = x1_img - x1, y1_img - y1\n    x2_patch, y2_patch = x1_patch + (x2_img - x1_img), y1_patch + (y2_img - y1_img)\n    patch[y1_patch:y2_patch, x1_patch:x2_patch] = img[y1_img:y2_img, x1_img:x2_img]\n    return patch\n\nsaved, skipped = 0, 0\n\nfor idx, row in tqdm(coords_df.iterrows(), total=len(coords_df)):\n    study_id = str(row[\"study_id\"])\n    series_id = str(row[\"series_id\"])\n    inst_no = int(row[\"instance_number\"])\n    x, y = int(row[\"x\"]), int(row[\"y\"])\n\n    # 🔹 Sanitize condition and level for filenames\n    condition = str(row['condition']).replace(\" \", \"_\")\n    level = str(row['level']).replace(\"/\", \"-\")\n\n    series_path = os.path.join(data_path, \"train_images\", study_id, series_id)\n    \n    try:\n        dcm_file = os.path.join(series_path, f\"{inst_no}.dcm\")\n        dcm = pydicom.dcmread(dcm_file)\n        img = dcm.pixel_array\n\n        patch = safe_crop(img, x, y, PATCH_SIZE)\n\n        out_path = os.path.join(OUTPUT_DIR, f\"{study_id}_{series_id}_{inst_no}_{condition}_{level}.png\")\n        ok = cv2.imwrite(out_path, patch)\n\n        if ok:\n            saved += 1\n        else:\n            skipped += 1\n    except Exception as e:\n        skipped += 1\n        continue\n\nprint(f\"✅ Done. Saved: {saved} patches | Skipped: {skipped}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:47:46.584354Z","iopub.execute_input":"2025-10-03T13:47:46.584625Z","iopub.status.idle":"2025-10-03T13:54:17.883411Z","shell.execute_reply.started":"2025-10-03T13:47:46.584605Z","shell.execute_reply":"2025-10-03T13:54:17.882758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil, os\n\n# Zip the patches\nshutil.make_archive(\"/kaggle/working/patches\", 'zip', \"/kaggle/working/patches\")\n\n# Check zip file size\n!ls -lh /kaggle/working/patches.zip\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:54:32.971144Z","iopub.execute_input":"2025-10-03T13:54:32.971737Z","iopub.status.idle":"2025-10-03T13:54:56.625791Z","shell.execute_reply.started":"2025-10-03T13:54:32.971713Z","shell.execute_reply":"2025-10-03T13:54:56.625025Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ****EDA****","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\n\n# Load label coordinates\ncoords_df = pd.read_csv(os.path.join(data_path, \"train_label_coordinates.csv\"))\n\n# Distribution by condition\ncond_counts = coords_df['condition'].value_counts()\n\nplt.figure(figsize=(8,5))\ncond_counts.plot(kind='bar')\nplt.title(\"Distribution of Conditions\")\nplt.ylabel(\"Count\")\nplt.show()\n\n# Distribution by level (L1/L2 … L5/S1)\nlevel_counts = coords_df['level'].value_counts()\n\nplt.figure(figsize=(8,5))\nlevel_counts.plot(kind='bar')\nplt.title(\"Distribution of Levels\")\nplt.ylabel(\"Count\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-03T15:06:44.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Create cross-tab\ncrosstab = pd.crosstab(coords_df['level'], coords_df['condition'])\n\nplt.figure(figsize=(12,6))\nsns.heatmap(crosstab, annot=True, fmt=\"d\", cmap=\"Blues\")\n\nplt.title(\"Condition Counts per Spinal Level\")\nplt.ylabel(\"Spinal Level\")\nplt.xlabel(\"Condition\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-03T15:06:44.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_counts = coords_df.groupby(\"study_id\")['condition'].count()\n\nplt.figure(figsize=(8,5))\nstudy_counts.hist(bins=50)\nplt.title(\"Number of Annotations per Study\")\nplt.xlabel(\"Annotations per study\")\nplt.ylabel(\"Frequency\")\nplt.show()\n\nprint(\"Min annotations:\", study_counts.min())\nprint(\"Max annotations:\", study_counts.max())\nprint(\"Average annotations:\", study_counts.mean())\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-03T15:06:44.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport random\nimport matplotlib.pyplot as plt\nimport os\n\nPATCH_DIR = \"/kaggle/working/patches\"\n\nall_patches = os.listdir(PATCH_DIR)\nprint(f\"Found {len(all_patches)} patches\")\n\n# Pick up to 12 random patches (or fewer if not enough available)\nn_samples = min(12, len(all_patches))\nsample_files = random.sample(all_patches, n_samples)\n\nplt.figure(figsize=(12,8))\nfor i, fname in enumerate(sample_files):\n    img = cv2.imread(os.path.join(PATCH_DIR, fname), cv2.IMREAD_GRAYSCALE)\n    plt.subplot(3,4,i+1)\n    plt.imshow(img, cmap='gray')\n    plt.title(fname.split(\"_\")[3])  # condition\n    plt.axis('off')\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-03T15:06:44.647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nprint(\"Does patch dir exist?\", os.path.exists(\"/kaggle/working/patches\"))\nprint(\"How many files inside?\", len(os.listdir(\"/kaggle/working/patches\")) if os.path.exists(\"/kaggle/working/patches\") else 0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:55:11.316172Z","iopub.execute_input":"2025-10-03T13:55:11.316819Z","iopub.status.idle":"2025-10-03T13:55:11.357206Z","shell.execute_reply.started":"2025-10-03T13:55:11.316787Z","shell.execute_reply":"2025-10-03T13:55:11.356564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\nPATCH_DIR = \"/kaggle/working/patches\"\n\n# Collect all patches\nall_patches = os.listdir(PATCH_DIR)\n\n# Extract labels (condition from filename)\ndata = []\nfor fname in all_patches:\n    parts = fname.split(\"_\")\n    condition = parts[-2]     # second last part\n    level = parts[-1].replace(\".png\", \"\")\n    data.append([fname, condition, level])\n\ndf = pd.DataFrame(data, columns=[\"filename\", \"condition\", \"level\"])\n\n# Split into train/val\ntrain_df, val_df = train_test_split(df, test_size=0.2, stratify=df[\"condition\"], random_state=42)\n\n# Image generators with augmentation\ntrain_datagen = ImageDataGenerator(\n    rescale=1./255,\n    rotation_range=10,\n    width_shift_range=0.1,\n    height_shift_range=0.1,\n    zoom_range=0.1,\n    horizontal_flip=True\n)\n\nval_datagen = ImageDataGenerator(rescale=1./255)\n\n# Flow from dataframe\ntrain_gen = train_datagen.flow_from_dataframe(\n    train_df,\n    directory=PATCH_DIR,\n    x_col=\"filename\",\n    y_col=\"condition\",\n    target_size=(128, 128),\n    class_mode=\"categorical\",\n    batch_size=32\n)\n\nval_gen = val_datagen.flow_from_dataframe(\n    val_df,\n    directory=PATCH_DIR,\n    x_col=\"filename\",\n    y_col=\"condition\",\n    target_size=(128, 128),\n    class_mode=\"categorical\",\n    batch_size=32\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:55:33.568078Z","iopub.execute_input":"2025-10-03T13:55:33.568712Z","iopub.status.idle":"2025-10-03T13:55:46.411054Z","shell.execute_reply.started":"2025-10-03T13:55:33.568689Z","shell.execute_reply":"2025-10-03T13:55:46.410399Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras import layers, models\nimport matplotlib.pyplot as plt\n\n# Determine number of classes from your train DataFrame\nnum_classes = train_df['condition'].nunique()\n\n# ======================\n# 1. Define Simple CNN\n# ======================\nmodel = models.Sequential([\n    layers.Conv2D(32, (3,3), activation='relu', input_shape=(128,128,3)),\n    layers.MaxPooling2D((2,2)),\n    \n    layers.Conv2D(64, (3,3), activation='relu'),\n    layers.MaxPooling2D((2,2)),\n    \n    layers.Conv2D(128, (3,3), activation='relu'),\n    layers.MaxPooling2D((2,2)),\n    \n    layers.Flatten(),\n    layers.Dense(128, activation='relu'),\n    layers.Dropout(0.5),\n    layers.Dense(num_classes, activation='softmax')  # fixed here\n])\n\n# ======================\n# 2. Compile\n# ======================\nmodel.compile(\n    optimizer='adam',\n    loss='categorical_crossentropy',\n    metrics=['accuracy']\n)\n\nmodel.summary()\n\n# ======================\n# 3. Train\n# ======================\nhistory = model.fit(\n    train_gen,\n    validation_data=val_gen,\n    epochs=10\n)\n\n# ======================\n# 4. Plot Accuracy & Loss\n# ======================\nplt.figure(figsize=(12,5))\n\n# Accuracy\nplt.subplot(1,2,1)\nplt.plot(history.history['accuracy'], label='Train Acc')\nplt.plot(history.history['val_accuracy'], label='Val Acc')\nplt.title(\"Model Accuracy\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.legend()\n\n# Loss\nplt.subplot(1,2,2)\nplt.plot(history.history['loss'], label='Train Loss')\nplt.plot(history.history['val_loss'], label='Val Loss')\nplt.title(\"Model Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T13:59:42.643783Z","iopub.execute_input":"2025-10-03T13:59:42.644091Z","iopub.status.idle":"2025-10-03T14:26:46.692254Z","shell.execute_reply.started":"2025-10-03T13:59:42.644070Z","shell.execute_reply":"2025-10-03T14:26:46.691544Z"}},"outputs":[],"execution_count":null}]}