{"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":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport time\nimport zipfile\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom IPython.display import display\n\nfrom sklearn.model_selection import GroupShuffleSplit, train_test_split\nfrom sklearn.metrics import (\n    accuracy_score,\n    classification_report,\n    confusion_matrix,\n    ConfusionMatrixDisplay,\n)\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nwarnings.filterwarnings(\"ignore\")\n\npd.set_option(\"display.max_columns\", 100)\npd.set_option(\"display.max_rows\", 100)\n\ndef seed_everything(seed=42):\n    \"\"\"\n    Make results more reproducible.\n\n    In deep learning, perfect reproducibility is not always guaranteed,\n    especially on GPU, but this helps a lot.\n    \"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything(42)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\nif torch.cuda.is_available():\n    print(\"GPU name:\", torch.cuda.get_device_name(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:07.676065Z","iopub.execute_input":"2026-07-25T11:36:07.676652Z","iopub.status.idle":"2026-07-25T11:36:17.217191Z","shell.execute_reply.started":"2026-07-25T11:36:07.676617Z","shell.execute_reply":"2026-07-25T11:36:17.216340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nprint(torch.__version__, torch.version.cuda)\nprint(torch.cuda.get_device_capability())  # likely (6, 0)\nprint(torch.cuda.get_arch_list())          # check if 'sm_60' is missing","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:23.525364Z","iopub.execute_input":"2026-07-25T11:36:23.525923Z","iopub.status.idle":"2026-07-25T11:36:23.531319Z","shell.execute_reply.started":"2026-07-25T11:36:23.525893Z","shell.execute_reply":"2026-07-25T11:36:23.530560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROOT = \"/kaggle/input/competitions/state-farm-distracted-driver-detection\"\nTRAIN_DIR = \"/kaggle/input/competitions/state-farm-distracted-driver-detection/imgs/train\"\ncsv_path = ROOT +'/'+ \"driver_imgs_list.csv\"\ndf = pd.read_csv(csv_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:27.795552Z","iopub.execute_input":"2026-07-25T11:36:27.795804Z","iopub.status.idle":"2026-07-25T11:36:27.825720Z","shell.execute_reply.started":"2026-07-25T11:36:27.795784Z","shell.execute_reply":"2026-07-25T11:36:27.825147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:32.675114Z","iopub.execute_input":"2026-07-25T11:36:32.675912Z","iopub.status.idle":"2026-07-25T11:36:32.707652Z","shell.execute_reply.started":"2026-07-25T11:36:32.675874Z","shell.execute_reply":"2026-07-25T11:36:32.706819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nCLASS_NAMES = {\n    \"c0\": \"safe driving\",\n    \"c1\": \"texting - right\",\n    \"c2\": \"talking on the phone - right\",\n    \"c3\": \"texting - left\",\n    \"c4\": \"talking on the phone - left\",\n    \"c5\": \"operating the radio\",\n    \"c6\": \"drinking\",\n    \"c7\": \"reaching behind\",\n    \"c8\": \"hair and makeup\",\n    \"c9\": \"talking to passenger\",\n}\n\nCLASS_ORDER = [f\"c{i}\" for i in range(10)]\nNUM_CLASSES = len(CLASS_ORDER)\n\ndf[\"class_name\"] = df[\"classname\"].map(CLASS_NAMES)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:36.075500Z","iopub.execute_input":"2026-07-25T11:36:36.076527Z","iopub.status.idle":"2026-07-25T11:36:36.084335Z","shell.execute_reply.started":"2026-07-25T11:36:36.076481Z","shell.execute_reply":"2026-07-25T11:36:36.083599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:40.995449Z","iopub.execute_input":"2026-07-25T11:36:40.996048Z","iopub.status.idle":"2026-07-25T11:36:41.007019Z","shell.execute_reply.started":"2026-07-25T11:36:40.996013Z","shell.execute_reply":"2026-07-25T11:36:41.006426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df[\"path\"] = df.apply(lambda row: str(TRAIN_DIR +'/'+ row[\"classname\"] +'/'+ row[\"img\"]), axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:44.696296Z","iopub.execute_input":"2026-07-25T11:36:44.697017Z","iopub.status.idle":"2026-07-25T11:36:44.831809Z","shell.execute_reply.started":"2026-07-25T11:36:44.696989Z","shell.execute_reply":"2026-07-25T11:36:44.831255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:49.275300Z","iopub.execute_input":"2026-07-25T11:36:49.275903Z","iopub.status.idle":"2026-07-25T11:36:49.287080Z","shell.execute_reply.started":"2026-07-25T11:36:49.275833Z","shell.execute_reply":"2026-07-25T11:36:49.286094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Number of images:\", len(df))\nprint(\"Number of unique drivers:\", df[\"subject\"].nunique())\nprint(\"Number of classes:\", df[\"classname\"].nunique())\n\nprint(\"\\nDuplicate rows:\")\nprint(df.duplicated(subset=[\"subject\", \"classname\", \"img\"]).sum())\n\nclass_counts = (\n    df[\"classname\"]\n    .value_counts()\n    .reindex(CLASS_ORDER)\n    .rename_axis(\"class_id\")\n    .reset_index(name=\"count\")\n)\n\nclass_counts[\"class_name\"] = class_counts[\"class_id\"].map(CLASS_NAMES)\nclass_counts[\"percent\"] = (class_counts[\"count\"] / class_counts[\"count\"].sum() * 100).round(2)\n\ndisplay(class_counts)\n\nplt.figure(figsize=(10, 4))\nplt.bar(class_counts[\"class_id\"], class_counts[\"count\"])\nplt.xticks(rotation=0)\nplt.xlabel(\"Class\")\nplt.ylabel(\"Number of images\")\nplt.title(\"Class distribution\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:36:54.920181Z","iopub.execute_input":"2026-07-25T11:36:54.920670Z","iopub.status.idle":"2026-07-25T11:36:55.128722Z","shell.execute_reply.started":"2026-07-25T11:36:54.920643Z","shell.execute_reply":"2026-07-25T11:36:55.128109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Driver / subject distribution\n\nsubject_counts = (\n    df[\"subject\"]\n    .value_counts()\n    .rename_axis(\"subject\")\n    .reset_index(name=\"image_count\")\n)\n\nprint(\"Images per driver:\")\ndisplay(subject_counts.describe().T)\n\nplt.figure(figsize=(10, 4))\nplt.bar(subject_counts[\"subject\"], subject_counts[\"image_count\"])\nplt.xticks(rotation=90)\nplt.ylabel(\"Number of images\")\nplt.title(\"Images per driver / subject\")\nplt.show()\n\nprint(\"\\nFirst rows:\")\ndisplay(subject_counts.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:37:02.895506Z","iopub.execute_input":"2026-07-25T11:37:02.896251Z","iopub.status.idle":"2026-07-25T11:37:03.113603Z","shell.execute_reply.started":"2026-07-25T11:37:02.896221Z","shell.execute_reply":"2026-07-25T11:37:03.112703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check image dimensions on a sample.\n\nsample_for_size = df.sample(min(500, len(df)), random_state=42)\n\nsizes = []\nfor p in sample_for_size[\"path\"]:\n    with Image.open(p) as img:\n        width, height = img.size\n        sizes.append((width, height))\n\nsize_df = pd.DataFrame(sizes, columns=[\"width\", \"height\"])\n\nprint(\"Image size summary from a sample:\")\ndisplay(size_df.describe().T)\n\nprint(\"\\nMost common sizes:\")\ndisplay(size_df.value_counts().head(10).to_frame(\"count\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:37:12.576017Z","iopub.execute_input":"2026-07-25T11:37:12.576295Z","iopub.status.idle":"2026-07-25T11:37:15.045469Z","shell.execute_reply.started":"2026-07-25T11:37:12.576274Z","shell.execute_reply":"2026-07-25T11:37:15.044650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_samples_by_class(dataframe, n_per_class=2, seed=42):\n    \"\"\"\n    Show a few images from each class.\n    \"\"\"\n    rng = np.random.default_rng(seed)\n\n    rows = []\n    for class_id in CLASS_ORDER:\n        class_df = dataframe[dataframe[\"classname\"] == class_id]\n        if len(class_df) == 0:\n            continue\n\n        take_n = min(n_per_class, len(class_df))\n        sampled_idx = rng.choice(class_df.index, size=take_n, replace=False)\n        rows.extend(sampled_idx)\n\n    n_images = len(rows)\n    n_cols = n_per_class\n    n_rows = int(np.ceil(n_images / n_cols))\n\n    plt.figure(figsize=(4 * n_cols, 3 * n_rows))\n\n    for i, idx in enumerate(rows):\n        row = dataframe.loc[idx]\n        img = Image.open(row[\"path\"]).convert(\"RGB\")\n\n        ax = plt.subplot(n_rows, n_cols, i + 1)\n        ax.imshow(img)\n        ax.set_title(f\"{row['classname']} - {row['class_name']}\\nsubject={row['subject']}\")\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\nshow_samples_by_class(df, n_per_class=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:37:18.855252Z","iopub.execute_input":"2026-07-25T11:37:18.855507Z","iopub.status.idle":"2026-07-25T11:37:22.172654Z","shell.execute_reply.started":"2026-07-25T11:37:18.855486Z","shell.execute_reply":"2026-07-25T11:37:22.171642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_samples_by_subject(dataframe, n_per_subject=3, seed=42, max_subjects=None):\n    \"\"\"\n    Show a few images from each subject / driver.\n    \"\"\"\n    rng = np.random.default_rng(seed)\n\n    # Make sure class_name exists\n    dataframe = dataframe.copy()\n    if \"class_name\" not in dataframe.columns:\n        dataframe[\"class_name\"] = dataframe[\"classname\"].map(CLASS_NAMES)\n\n    subjects = sorted(dataframe[\"subject\"].unique())\n\n    if max_subjects is not None:\n        subjects = subjects[:max_subjects]\n\n    rows = []\n\n    for subject in subjects:\n        subject_df = dataframe[dataframe[\"subject\"] == subject]\n\n        if len(subject_df) == 0:\n            continue\n\n        take_n = min(n_per_subject, len(subject_df))\n        sampled_idx = rng.choice(subject_df.index, size=take_n, replace=False)\n        rows.extend(sampled_idx)\n\n    n_images = len(rows)\n    n_cols = n_per_subject\n    n_rows = int(np.ceil(n_images / n_cols))\n\n    plt.figure(figsize=(4 * n_cols, 3 * n_rows))\n\n    for i, idx in enumerate(rows):\n        row = dataframe.loc[idx]\n        img = Image.open(row[\"path\"]).convert(\"RGB\")\n\n        ax = plt.subplot(n_rows, n_cols, i + 1)\n        ax.imshow(img)\n        ax.set_title(\n            f\"subject={row['subject']}\\n\"\n            f\"{row['classname']} - {row['class_name']}\"\n        )\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\nshow_samples_by_subject(df, n_per_subject=5, max_subjects=10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:37:33.410375Z","iopub.execute_input":"2026-07-25T11:37:33.411192Z","iopub.status.idle":"2026-07-25T11:37:39.009541Z","shell.execute_reply.started":"2026-07-25T11:37:33.411156Z","shell.execute_reply":"2026-07-25T11:37:39.008370Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Class distribution per subject\n\nCLASS_ORDER = [f\"c{i}\" for i in range(10)]\n\n# Make sure class_name exists\nif \"class_name\" not in df.columns:\n    df[\"class_name\"] = df[\"classname\"].map(CLASS_NAMES)\n\n# Count images for each subject-class pair\nsubject_class_counts = pd.crosstab(\n    df[\"subject\"],\n    df[\"classname\"]\n).reindex(columns=CLASS_ORDER, fill_value=0)\n\n# Convert counts to percentages within each subject\nsubject_class_percent = (\n    subject_class_counts\n    .div(subject_class_counts.sum(axis=1), axis=0)\n    .mul(100)\n    .round(2)\n)\n\nsubject_class_percent.plot(\n    kind=\"bar\",\n    stacked=True,\n    figsize=(14, 6),\n    width=0.85\n)\n\nplt.axhline(100, linewidth=1)\nplt.xlabel(\"Subject\")\nplt.ylabel(\"Percentage of images\")\nplt.title(\"Class composition per subject\")\nplt.legend(title=\"Class\", bbox_to_anchor=(1.02, 1), loc=\"upper left\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:38:13.025559Z","iopub.execute_input":"2026-07-25T11:38:13.026070Z","iopub.status.idle":"2026-07-25T11:38:13.607564Z","shell.execute_reply.started":"2026-07-25T11:38:13.026037Z","shell.execute_reply":"2026-07-25T11:38:13.606694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# normal random split\n\nrandom_train_df, random_valid_df = train_test_split(\n    df,\n    test_size=0.20,\n    random_state=42,\n    stratify=df[\"classname\"]\n)\n\nrandom_train_subjects = set(random_train_df[\"subject\"])\nrandom_valid_subjects = set(random_valid_df[\"subject\"])\nrandom_overlap = random_train_subjects.intersection(random_valid_subjects)\n\nprint(\"RANDOM SPLIT\")\nprint(\"Train images:\", len(random_train_df))\nprint(\"Valid images:\", len(random_valid_df))\nprint(\"Unique train drivers:\", len(random_train_subjects))\nprint(\"Unique valid drivers:\", len(random_valid_subjects))\nprint(\"Drivers appearing in BOTH train and validation:\", len(random_overlap))\nprint(\"Example overlapping drivers:\", sorted(list(random_overlap))[:10])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# group split by driver / subject\n\ngss = GroupShuffleSplit(\n    n_splits=1,\n    test_size=0.20,\n    random_state=42\n)\n\ntrain_idx, valid_idx = next(\n    gss.split(\n        df,\n        y=df[\"classname\"],\n        groups=df[\"subject\"]\n    )\n)\n\ntrain_df = df.iloc[train_idx].reset_index(drop=True)\nvalid_df = df.iloc[valid_idx].reset_index(drop=True)\n\ntrain_subjects = set(train_df[\"subject\"])\nvalid_subjects = set(valid_df[\"subject\"])\noverlap = train_subjects.intersection(valid_subjects)\n\nprint(\"GROUP SPLIT BY DRIVER\")\nprint(\"Train images:\", len(train_df))\nprint(\"Valid images:\", len(valid_df))\nprint(\"Unique train drivers:\", len(train_subjects))\nprint(\"Unique valid drivers:\", len(valid_subjects))\nprint(\"Drivers appearing in BOTH train and validation:\", len(overlap))\n\nprint(\"\\nTrain class distribution:\")\ndisplay(train_df[\"classname\"].value_counts().reindex(CLASS_ORDER).to_frame(\"train_count\"))\n\nprint(\"\\nValidation class distribution:\")\ndisplay(valid_df[\"classname\"].value_counts().reindex(CLASS_ORDER).to_frame(\"valid_count\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:38:20.466474Z","iopub.execute_input":"2026-07-25T11:38:20.467240Z","iopub.status.idle":"2026-07-25T11:38:20.505786Z","shell.execute_reply.started":"2026-07-25T11:38:20.467208Z","shell.execute_reply":"2026-07-25T11:38:20.504916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_TRAIN_IMAGES_PER_CLASS = 1500\nMAX_VALID_IMAGES_PER_CLASS = 400\n\ndef sample_per_class(dataframe, max_per_class=None, seed=42):\n    \"\"\"\n    Keep at most max_per_class images for each class.\n    \"\"\"\n    if max_per_class is None:\n        return dataframe.sample(frac=1, random_state=seed).reset_index(drop=True)\n\n    parts = []\n    for class_id in CLASS_ORDER:\n        class_df = dataframe[dataframe[\"classname\"] == class_id]\n        n = min(max_per_class, len(class_df))\n        parts.append(class_df.sample(n=n, random_state=seed))\n\n    return pd.concat(parts).sample(frac=1, random_state=seed).reset_index(drop=True)\n# random_train_df\ntrain_df_small = sample_per_class(train_df, MAX_TRAIN_IMAGES_PER_CLASS)\nvalid_df_small = sample_per_class(valid_df, MAX_VALID_IMAGES_PER_CLASS)\n\nprint(\"Training subset shape:\", train_df_small.shape)\nprint(\"Validation subset shape:\", valid_df_small.shape)\n\nprint(\"\\nTraining subset class counts:\")\ndisplay(train_df_small[\"classname\"].value_counts().reindex(CLASS_ORDER).to_frame(\"count\"))\n\nprint(\"\\nValidation subset class counts:\")\ndisplay(valid_df_small[\"classname\"].value_counts().reindex(CLASS_ORDER).to_frame(\"count\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:39:59.067623Z","iopub.execute_input":"2026-07-25T11:39:59.068435Z","iopub.status.idle":"2026-07-25T11:39:59.132290Z","shell.execute_reply.started":"2026-07-25T11:39:59.068405Z","shell.execute_reply":"2026-07-25T11:39:59.131661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\n\nMEAN = [0.5, 0.5, 0.5]\nSTD = [0.5, 0.5, 0.5]\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n\n    # Small rotation can make model more robust.\n    transforms.RandomRotation(degrees=5),\n\n    # Lighting/color changes are realistic.\n    transforms.ColorJitter(\n        brightness=0.15,\n        contrast=0.15,\n        saturation=0.10,\n        hue=0.02\n    ),\n\n    transforms.ToTensor(),\n    transforms.Normalize(mean=MEAN, std=STD),\n])\n\nvalid_transform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=MEAN, std=STD),\n])\n\nclass_to_idx = {class_id: i for i, class_id in enumerate(CLASS_ORDER)}\nidx_to_class = {i: class_id for class_id, i in class_to_idx.items()}\n\nclass DriverImageDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n\n        # Open image\n        image = Image.open(row[\"path\"]).convert(\"RGB\")\n\n        # Transform image\n        if self.transform is not None:\n            image = self.transform(image)\n\n        # Convert class id to integer label\n        label = class_to_idx[row[\"classname\"]]\n\n        return image, label\n\ntrain_dataset = DriverImageDataset(train_df_small, transform=train_transform)\nvalid_dataset = DriverImageDataset(valid_df_small, transform=valid_transform)\n\n\nprint(\"Train dataset length:\", len(train_dataset))\nprint(\"Valid dataset length:\", len(valid_dataset))\n\none_image, one_label = train_dataset[0]\n\nprint(\"\\nOne image tensor shape:\", one_image.shape)\nprint(\"One label:\", one_label, \"->\", idx_to_class[one_label], \"->\", CLASS_NAMES[idx_to_class[one_label]])\nprint(\"Pixel min:\", one_image.min().item())\nprint(\"Pixel mean:\", one_image.mean().item())\nprint(\"Pixel max:\", one_image.max().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:40:06.529728Z","iopub.execute_input":"2026-07-25T11:40:06.530118Z","iopub.status.idle":"2026-07-25T11:40:06.658936Z","shell.execute_reply.started":"2026-07-25T11:40:06.530090Z","shell.execute_reply":"2026-07-25T11:40:06.658155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 64\n\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    pin_memory=torch.cuda.is_available()\n)\n\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    pin_memory=torch.cuda.is_available()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:40:12.598508Z","iopub.execute_input":"2026-07-25T11:40:12.598756Z","iopub.status.idle":"2026-07-25T11:40:12.603621Z","shell.execute_reply.started":"2026-07-25T11:40:12.598736Z","shell.execute_reply":"2026-07-25T11:40:12.602823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision.models import vit_b_16, ViT_B_16_Weights\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, num_classes=10, pretrained=True):\n        super().__init__()\n\n        if pretrained:\n            weights = ViT_B_16_Weights.DEFAULT\n            self.model = vit_b_16(weights=weights)\n        else:\n            self.model = vit_b_16(weights=None)\n\n        # -----------------------------\n        # Freeze ALL parameters\n        # -----------------------------\n        for param in self.model.parameters():\n            param.requires_grad = False\n\n        # -----------------------------\n        # Unfreeze last 6 encoder blocks\n        # (ViT-B/16 has 12 encoder blocks)\n        # -----------------------------\n        for block in self.model.encoder.layers[6:]:\n            for param in block.parameters():\n                param.requires_grad = True\n\n        # -----------------------------\n        # Replace classification head\n        # -----------------------------\n        in_features = self.model.heads.head.in_features\n        self.model.heads.head = nn.Linear(in_features, num_classes)\n\n        # Ensure classifier is trainable\n        for param in self.model.heads.parameters():\n            param.requires_grad = True\n\n    def forward(self, x):\n        return self.model(x)\n\n\n# ==============================\n# Create Model\n# ==============================\nmodel = VisionTransformer(\n    num_classes=NUM_CLASSES,\n    pretrained=True\n).to(device)\n\nprint(model)\n\n# ==============================\n# Count Parameters\n# ==============================\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"Total Parameters: {total_params:,}\")\nprint(f\"Trainable Parameters: {trainable_params:,}\")\nprint(f\"Frozen Parameters: {total_params - trainable_params:,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:40:18.212785Z","iopub.execute_input":"2026-07-25T11:40:18.213343Z","iopub.status.idle":"2026-07-25T11:40:21.718077Z","shell.execute_reply.started":"2026-07-25T11:40:18.213314Z","shell.execute_reply":"2026-07-25T11:40:21.717246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, loss_fn, optimizer, device):\n    model.train()\n\n    total_loss = 0.0\n    total_correct = 0\n    total_examples = 0\n\n    for images, labels in dataloader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        # Forward\n        outputs = model(images)\n        loss = loss_fn(outputs, labels)\n\n        # Backward\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        # Metrics\n        batch_size = labels.size(0)\n        total_loss += loss.item() * batch_size\n\n        predictions = outputs.argmax(dim=1)\n        total_correct += (predictions == labels).sum().item()\n        total_examples += batch_size\n\n    avg_loss = total_loss / total_examples\n    accuracy = total_correct / total_examples\n\n    return avg_loss, accuracy\n\n\ndef evaluate(model, dataloader, loss_fn, device):\n    model.eval()\n\n    total_loss = 0.0\n    total_correct = 0\n    total_examples = 0\n\n    all_predictions = []\n    all_labels = []\n    all_probabilities = []\n\n    with torch.no_grad():\n        for images, labels in dataloader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = loss_fn(outputs, labels)\n\n            probabilities = torch.softmax(outputs, dim=1)\n            predictions = outputs.argmax(dim=1)\n\n            batch_size = labels.size(0)\n            total_loss += loss.item() * batch_size\n            total_correct += (predictions == labels).sum().item()\n            total_examples += batch_size\n\n            all_predictions.append(predictions.cpu())\n            all_labels.append(labels.cpu())\n            all_probabilities.append(probabilities.cpu())\n\n    avg_loss = total_loss / total_examples\n    accuracy = total_correct / total_examples\n\n    all_predictions = torch.cat(all_predictions).numpy()\n    all_labels = torch.cat(all_labels).numpy()\n    all_probabilities = torch.cat(all_probabilities).numpy()\n\n    return avg_loss, accuracy, all_predictions, all_labels, all_probabilities","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:40:37.517763Z","iopub.execute_input":"2026-07-25T11:40:37.518203Z","iopub.status.idle":"2026-07-25T11:40:37.528162Z","shell.execute_reply.started":"2026-07-25T11:40:37.518173Z","shell.execute_reply":"2026-07-25T11:40:37.527261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss()\n\n# If classes were very imbalanced, we could add class weights.\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n\nEPOCHS = 3\n\nhistory = []\n\nfor epoch in range(EPOCHS):\n    start = time.time()\n\n    train_loss, train_acc = train_one_epoch(\n        model,\n        train_loader,\n        loss_fn,\n        optimizer,\n        device\n    )\n\n    valid_loss, valid_acc, valid_pred, valid_true, valid_probs = evaluate(\n        model,\n        valid_loader,\n        loss_fn,\n        device\n    )\n\n    epoch_time = time.time() - start\n\n    history.append({\n        \"epoch\": epoch + 1,\n        \"train_loss\": train_loss,\n        \"train_acc\": train_acc,\n        \"valid_loss\": valid_loss,\n        \"valid_acc\": valid_acc,\n        \"time_sec\": epoch_time,\n    })\n\n    print(\n        f\"Epoch {epoch + 1}/{EPOCHS} | \"\n        f\"train_loss={train_loss:.4f} | \"\n        f\"train_acc={train_acc:.4f} | \"\n        f\"valid_loss={valid_loss:.4f} | \"\n        f\"valid_acc={valid_acc:.4f} | \"\n        f\"time={epoch_time:.1f}s\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T11:40:41.928196Z","iopub.execute_input":"2026-07-25T11:40:41.928922Z","iopub.status.idle":"2026-07-25T12:16:04.571402Z","shell.execute_reply.started":"2026-07-25T11:40:41.928896Z","shell.execute_reply":"2026-07-25T12:16:04.570236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_df = pd.DataFrame(history)\ndisplay(history_df)\n\nplt.figure(figsize=(7, 4))\nplt.plot(history_df[\"epoch\"], history_df[\"train_loss\"], marker=\"o\", label=\"train loss\")\nplt.plot(history_df[\"epoch\"], history_df[\"valid_loss\"], marker=\"o\", label=\"valid loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs validation loss\")\nplt.legend()\nplt.show()\n\nplt.figure(figsize=(7, 4))\nplt.plot(history_df[\"epoch\"], history_df[\"train_acc\"], marker=\"o\", label=\"train acc\")\nplt.plot(history_df[\"epoch\"], history_df[\"valid_acc\"], marker=\"o\", label=\"valid acc\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Training vs validation accuracy\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:35:14.972645Z","iopub.execute_input":"2026-07-25T12:35:14.973283Z","iopub.status.idle":"2026-07-25T12:35:15.258284Z","shell.execute_reply.started":"2026-07-25T12:35:14.973253Z","shell.execute_reply":"2026-07-25T12:35:15.257664Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run final evaluation again, to make sure variables are fresh.\nvalid_loss, valid_acc, valid_pred, valid_true, valid_probs = evaluate(\n    model,\n    valid_loader,\n    loss_fn,\n    device\n)\n\nprint(f\"Validation loss: {valid_loss:.4f}\")\nprint(f\"Validation accuracy: {valid_acc:.4f}\")\n\ntarget_names = [f\"{class_id}: {CLASS_NAMES[class_id]}\" for class_id in CLASS_ORDER]\n\nprint(\"\\nClassification report:\")\nprint(classification_report(valid_true, valid_pred, target_names=target_names))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:35:24.618404Z","iopub.execute_input":"2026-07-25T12:35:24.618801Z","iopub.status.idle":"2026-07-25T12:36:41.633927Z","shell.execute_reply.started":"2026-07-25T12:35:24.618774Z","shell.execute_reply":"2026-07-25T12:36:41.632804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import (\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\naccuracy = accuracy_score(valid_true, valid_pred)\nprecision = precision_score(valid_true, valid_pred, average=\"weighted\")\nrecall = recall_score(valid_true, valid_pred, average=\"weighted\")\nf1 = f1_score(valid_true, valid_pred, average=\"weighted\")\n\nprint(\"=\"*40)\nprint(\"Overall Performance Metrics\")\nprint(\"=\"*40)\nprint(f\"Accuracy : {accuracy:.4f}\")\nprint(f\"Precision: {precision:.4f}\")\nprint(f\"Recall   : {recall:.4f}\")\nprint(f\"F1-Score : {f1:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:38:43.307627Z","iopub.execute_input":"2026-07-25T12:38:43.308210Z","iopub.status.idle":"2026-07-25T12:38:43.324449Z","shell.execute_reply.started":"2026-07-25T12:38:43.308181Z","shell.execute_reply":"2026-07-25T12:38:43.323771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\ntarget_names = [CLASS_NAMES[i] for i in CLASS_ORDER]\n\nprint(\n    classification_report(\n        valid_true,\n        valid_pred,\n        target_names=target_names,\n        digits=4\n    )\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:38:48.395640Z","iopub.execute_input":"2026-07-25T12:38:48.396399Z","iopub.status.idle":"2026-07-25T12:38:48.411829Z","shell.execute_reply.started":"2026-07-25T12:38:48.396371Z","shell.execute_reply":"2026-07-25T12:38:48.411150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\ncm = confusion_matrix(valid_true, valid_pred)\n\nplt.figure(figsize=(12,10))\nsns.heatmap(\n    cm,\n    annot=True,\n    fmt=\"d\",\n    cmap=\"Blues\",\n    xticklabels=target_names,\n    yticklabels=target_names\n)\n\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.title(\"Confusion Matrix\")\nplt.xticks(rotation=90)\nplt.yticks(rotation=0)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:38:53.868136Z","iopub.execute_input":"2026-07-25T12:38:53.868936Z","iopub.status.idle":"2026-07-25T12:38:54.579428Z","shell.execute_reply.started":"2026-07-25T12:38:53.868904Z","shell.execute_reply":"2026-07-25T12:38:54.578394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"=\"*50)\nprint(\"Precision\")\nprint(\"=\"*50)\n\nprint(\"Macro   :\", precision_score(valid_true, valid_pred, average=\"macro\"))\nprint(\"Weighted:\", precision_score(valid_true, valid_pred, average=\"weighted\"))\nprint(\"Micro   :\", precision_score(valid_true, valid_pred, average=\"micro\"))\n\nprint(\"\\n\")\n\nprint(\"=\"*50)\nprint(\"Recall\")\nprint(\"=\"*50)\n\nprint(\"Macro   :\", recall_score(valid_true, valid_pred, average=\"macro\"))\nprint(\"Weighted:\", recall_score(valid_true, valid_pred, average=\"weighted\"))\nprint(\"Micro   :\", recall_score(valid_true, valid_pred, average=\"micro\"))\n\nprint(\"\\n\")\n\nprint(\"=\"*50)\nprint(\"F1 Score\")\nprint(\"=\"*50)\n\nprint(\"Macro   :\", f1_score(valid_true, valid_pred, average=\"macro\"))\nprint(\"Weighted:\", f1_score(valid_true, valid_pred, average=\"weighted\"))\nprint(\"Micro   :\", f1_score(valid_true, valid_pred, average=\"micro\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:39:00.870500Z","iopub.execute_input":"2026-07-25T12:39:00.871243Z","iopub.status.idle":"2026-07-25T12:39:00.903330Z","shell.execute_reply.started":"2026-07-25T12:39:00.871214Z","shell.execute_reply":"2026-07-25T12:39:00.902712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\nfrom sklearn.preprocessing import label_binarize\n\nnum_classes = len(CLASS_ORDER)\n\ny_true_bin = label_binarize(valid_true, classes=range(num_classes))\n\nroc_auc = roc_auc_score(\n    y_true_bin,\n    valid_probs,\n    average=\"macro\",\n    multi_class=\"ovr\"\n)\n\nprint(f\"Macro ROC-AUC: {roc_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:39:12.286766Z","iopub.execute_input":"2026-07-25T12:39:12.287547Z","iopub.status.idle":"2026-07-25T12:39:12.312940Z","shell.execute_reply.started":"2026-07-25T12:39:12.287517Z","shell.execute_reply":"2026-07-25T12:39:12.312277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(valid_true, valid_pred, labels=list(range(NUM_CLASSES)))\n\ndisp = ConfusionMatrixDisplay(\n    confusion_matrix=cm,\n    display_labels=CLASS_ORDER\n)\n\nfig, ax = plt.subplots(figsize=(8, 8))\ndisp.plot(ax=ax, values_format=\"d\", xticks_rotation=45)\nplt.title(\"Validation confusion matrix\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T12:39:16.266695Z","iopub.execute_input":"2026-07-25T12:39:16.267407Z","iopub.status.idle":"2026-07-25T12:39:16.551304Z","shell.execute_reply.started":"2026-07-25T12:39:16.267377Z","shell.execute_reply":"2026-07-25T12:39:16.550661Z"}},"outputs":[],"execution_count":null}]}