{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":5048,"databundleVersionId":868335,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Distracted Driver Classification – ResNet18 Transfer Learning\n\nDashboard camera images classified into **10 behaviour categories** using a fine-tuned ResNet18 backbone.\n\n| Grade threshold | Bonus |\n|---|---|\n| ≥ 85 % test accuracy | +0.5 pts |\n| ≥ 90 % test accuracy | +1.0 pts (max) |\n\n**Classes:** c0 safe driving · c1 texting-R · c2 phone-R · c3 texting-L · c4 phone-L · c5 radio · c6 drinking · c7 reaching behind · c8 hair/makeup · c9 talking to passenger","metadata":{}},{"cell_type":"code","source":"import os, random, warnings\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import OneCycleLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\n\nwarnings.filterwarnings('ignore')\n\n# ── Deterministic behaviour ───────────────────────────────────────────────────\nSEED = 2024\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Running on: {DEVICE}', end='  ')\nif DEVICE.type == 'cuda':\n    print(torch.cuda.get_device_name(0))\nelse:\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:33.727535Z","iopub.execute_input":"2026-05-15T19:09:33.727805Z","iopub.status.idle":"2026-05-15T19:09:50.196608Z","shell.execute_reply.started":"2026-05-15T19:09:33.727779Z","shell.execute_reply":"2026-05-15T19:09:50.195857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROOT       = '/kaggle/input/competitions/state-farm-distracted-driver-detection'\nTRAIN_DIR  = os.path.join(ROOT, 'imgs', 'train')\nTEST_DIR   = os.path.join(ROOT, 'imgs', 'test')   # unlabeled – not used for scoring\nCSV_PATH   = os.path.join(ROOT, 'driver_imgs_list.csv')\nCKPT_PATH  = '/kaggle/working/resnet18_best.pt'\n\nN_CLASSES  = 10\nBATCH      = 32\nIMG_SIZE   = 224\n\nCLASS_NAMES = [\n    'c0: safe driving',      'c1: texting – R',\n    'c2: phone – R',         'c3: texting – L',\n    'c4: phone – L',         'c5: radio',\n    'c6: drinking',          'c7: reaching behind',\n    'c8: hair / makeup',     'c9: talking to passenger'\n]\n\nassert os.path.isdir(TRAIN_DIR), 'TRAIN_DIR not found'\nassert os.path.isfile(CSV_PATH), 'driver CSV not found'\nprint('Train classes:', sorted(os.listdir(TRAIN_DIR)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:50.197625Z","iopub.execute_input":"2026-05-15T19:09:50.198065Z","iopub.status.idle":"2026-05-15T19:09:50.211996Z","shell.execute_reply.started":"2026-05-15T19:09:50.198041Z","shell.execute_reply":"2026-05-15T19:09:50.211214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(CSV_PATH)\nprint(df.head(3).to_string(index=False))\nprint(f'\\nTotal labelled images : {len(df):,}')\nprint(f'Unique drivers        : {df[\"subject\"].nunique()}')\n\ndrivers = sorted(df['subject'].unique())\nrandom.seed(SEED)\nval_subjects   = set(random.sample(drivers, k=max(1, round(len(drivers) * 0.20))))\ntrain_subjects = set(drivers) - val_subjects\n\nprint(f'\\nTrain drivers ({len(train_subjects)}): {sorted(train_subjects)}')\nprint(f'Val   drivers ({len(val_subjects)}): {sorted(val_subjects)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:50.213563Z","iopub.execute_input":"2026-05-15T19:09:50.213909Z","iopub.status.idle":"2026-05-15T19:09:50.276171Z","shell.execute_reply.started":"2026-05-15T19:09:50.213888Z","shell.execute_reply":"2026-05-15T19:09:50.275306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StateFarmDataset(Dataset):\n    \"\"\"Labelled split built from driver_imgs_list.csv.\"\"\"\n\n    def __init__(self, root, dataframe, subjects=None, transform=None):\n        self.tfm = transform\n        sub = dataframe if subjects is None else dataframe[dataframe['subject'].isin(subjects)]\n        self.records = [\n            (os.path.join(root, r.classname, r.img), int(r.classname[1]))\n            for r in sub.itertuples(index=False)\n        ]\n\n    def __len__(self):  return len(self.records)\n\n    def __getitem__(self, i):\n        path, lbl = self.records[i]\n        img = Image.open(path).convert('RGB')\n        return (self.tfm(img) if self.tfm else img), lbl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:50.277188Z","iopub.execute_input":"2026-05-15T19:09:50.277548Z","iopub.status.idle":"2026-05-15T19:09:50.283938Z","shell.execute_reply.started":"2026-05-15T19:09:50.277517Z","shell.execute_reply":"2026-05-15T19:09:50.282783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MU  = [0.485, 0.456, 0.406]\nSTD = [0.229, 0.224, 0.225]\n\ntrain_tfm = transforms.Compose([\n    transforms.RandomResizedCrop(IMG_SIZE, scale=(0.65, 1.0), ratio=(0.8, 1.25)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomApply([transforms.GaussianBlur(kernel_size=3)], p=0.2),\n    transforms.ColorJitter(brightness=0.35, contrast=0.35, saturation=0.25, hue=0.06),\n    transforms.RandomRotation(12),\n    transforms.ToTensor(),\n    transforms.Normalize(MU, STD),\n    transforms.RandomErasing(p=0.15, scale=(0.02, 0.12)),\n])\n\neval_tfm = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(IMG_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize(MU, STD),\n])\n\ntrain_ds = StateFarmDataset(TRAIN_DIR, df, train_subjects, train_tfm)\nval_ds   = StateFarmDataset(TRAIN_DIR, df, val_subjects,   eval_tfm)\n\ntrain_dl = DataLoader(train_ds, batch_size=BATCH, shuffle=True,  num_workers=2, pin_memory=True)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(f'Train images : {len(train_ds):,}   Val images : {len(val_ds):,}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:50.285306Z","iopub.execute_input":"2026-05-15T19:09:50.285646Z","iopub.status.idle":"2026-05-15T19:09:50.360772Z","shell.execute_reply.started":"2026-05-15T19:09:50.285624Z","shell.execute_reply":"2026-05-15T19:09:50.360104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def to_rgb(t):\n    return (t * torch.tensor(STD).view(3,1,1) + torch.tensor(MU).view(3,1,1)).clamp(0,1).permute(1,2,0).numpy()\n\ngallery = {}\nfor img, lbl in train_ds:\n    if lbl not in gallery: gallery[lbl] = img\n    if len(gallery) == N_CLASSES: break\n\nfig, axes = plt.subplots(2, 5, figsize=(17, 7))\nfor ax, i in zip(axes.flat, range(N_CLASSES)):\n    ax.imshow(to_rgb(gallery[i]))\n    ax.set_title(CLASS_NAMES[i], fontsize=8, pad=4)\n    ax.axis('off')\nfig.suptitle('Representative sample per class', fontsize=12, y=1.01)\nplt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:09:50.361623Z","iopub.execute_input":"2026-05-15T19:09:50.361907Z","iopub.status.idle":"2026-05-15T19:10:03.564127Z","shell.execute_reply.started":"2026-05-15T19:09:50.361874Z","shell.execute_reply":"2026-05-15T19:10:03.563166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_resnet18(n_cls=10):\n    net = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\n    in_f = net.fc.in_features   # 512\n    net.fc = nn.Sequential(\n        nn.Dropout(0.45),\n        nn.Linear(in_f, 256),\n        nn.ReLU(inplace=True),\n        nn.Dropout(0.25),\n        nn.Linear(256, n_cls)\n    )\n    return net\n\ndef set_grad(net, head_only: bool):\n    for name, p in net.named_parameters():\n        p.requires_grad = (not head_only) or ('fc' in name)\n\nnet = build_resnet18(N_CLASSES).to(DEVICE)\ntotal = sum(p.numel() for p in net.parameters())\nprint(f'Parameters – total: {total:,}  |  fc head: {sum(p.numel() for p in net.fc.parameters()):,}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:10:03.565585Z","iopub.execute_input":"2026-05-15T19:10:03.565921Z","iopub.status.idle":"2026-05-15T19:10:04.428709Z","shell.execute_reply.started":"2026-05-15T19:10:03.565890Z","shell.execute_reply":"2026-05-15T19:10:04.428051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss(label_smoothing=0.05)\n\ndef run_epoch(net, loader, optimiser=None):\n    training = optimiser is not None\n    net.train() if training else net.eval()\n    tot_loss = correct = total = 0\n    ctx = torch.enable_grad() if training else torch.no_grad()\n    with ctx:\n        for imgs, lbls in loader:\n            imgs, lbls = imgs.to(DEVICE), lbls.to(DEVICE)\n            out  = net(imgs)\n            loss = loss_fn(out, lbls)\n            if training:\n                optimiser.zero_grad()\n                loss.backward()\n                optimiser.step()\n            tot_loss += loss.item() * imgs.size(0)\n            correct  += (out.argmax(1) == lbls).sum().item()\n            total    += imgs.size(0)\n    return tot_loss / total, correct / total\n\n\ndef draw_curves(h):\n    ep = range(1, len(h['tr_acc']) + 1)\n    fig, (a1, a2) = plt.subplots(1, 2, figsize=(14, 5))\n    a1.plot(ep, h['tr_loss'], label='train'); a1.plot(ep, h['vl_loss'], label='val')\n    a1.set(title='Cross-Entropy Loss', xlabel='Epoch'); a1.legend()\n    a2.plot(ep, [v*100 for v in h['tr_acc']], label='train')\n    a2.plot(ep, [v*100 for v in h['vl_acc']], label='val')\n    a2.axhline(85, ls='--', c='darkorange', lw=1.2, label='85 % target')\n    a2.axhline(90, ls='--', c='green',      lw=1.2, label='90 % bonus')\n    a2.set(title='Accuracy (%)', xlabel='Epoch'); a2.legend()\n    plt.tight_layout(); plt.show()\n\nhist = {k: [] for k in ('tr_loss','tr_acc','vl_loss','vl_acc')}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:10:04.429582Z","iopub.execute_input":"2026-05-15T19:10:04.429816Z","iopub.status.idle":"2026-05-15T19:10:04.438510Z","shell.execute_reply.started":"2026-05-15T19:10:04.429796Z","shell.execute_reply":"2026-05-15T19:10:04.437881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WU_EPOCHS = 5\nset_grad(net, head_only=True)\n\nopt1 = optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()), lr=1e-3, weight_decay=1e-4)\nsch1 = OneCycleLR(opt1, max_lr=1e-3, steps_per_epoch=len(train_dl), epochs=WU_EPOCHS)\n\nprint(f'── Phase 1: head warm-up ({WU_EPOCHS} epochs) ──')\nfor ep in range(1, WU_EPOCHS + 1):\n    tl, ta = run_epoch(net, train_dl, opt1); sch1.step()\n    vl, va = run_epoch(net, val_dl)\n    for k, v in zip(hist, [tl, ta, vl, va]): hist[k].append(v)\n    print(f'  [{ep:02d}/{WU_EPOCHS}]  train {ta*100:5.2f}%  val {va*100:5.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:10:04.440881Z","iopub.execute_input":"2026-05-15T19:10:04.441169Z","iopub.status.idle":"2026-05-15T19:22:46.143963Z","shell.execute_reply.started":"2026-05-15T19:10:04.441149Z","shell.execute_reply":"2026-05-15T19:22:46.143146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FT_EPOCHS = 20\nset_grad(net, head_only=False)\n\nopt2 = optim.AdamW([\n    {'params': [p for n,p in net.named_parameters() if 'fc' not in n], 'lr': 1e-4},\n    {'params': net.fc.parameters(),                                     'lr': 3e-4},\n], weight_decay=1e-4)\n\nsch2 = OneCycleLR(opt2, max_lr=[1e-4, 3e-4],\n                  steps_per_epoch=len(train_dl), epochs=FT_EPOCHS)\n\nbest_va, patience_cnt = 0.0, 0\nPATIENCE = 7  # early stopping\n\nprint(f'── Phase 2: full fine-tuning ({FT_EPOCHS} epochs) ──')\nfor ep in range(1, FT_EPOCHS + 1):\n    tl, ta = run_epoch(net, train_dl, opt2); sch2.step()\n    vl, va = run_epoch(net, val_dl)\n    for k, v in zip(hist, [tl, ta, vl, va]): hist[k].append(v)\n\n    tag = ''\n    if va > best_va:\n        best_va = va\n        torch.save(net.state_dict(), CKPT_PATH)\n        tag = '  ← best'\n        patience_cnt = 0\n    else:\n        patience_cnt += 1\n\n    print(f'  [{ep:02d}/{FT_EPOCHS}]  train {ta*100:5.2f}%  val {va*100:5.2f}%' + tag)\n\n    if patience_cnt >= PATIENCE:\n        print(f'  Early stopping after {PATIENCE} epochs without improvement.')\n        break\n\nprint(f'\\nBest val accuracy: {best_va*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-15T19:22:46.145388Z","iopub.execute_input":"2026-05-15T19:22:46.145631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"draw_curves(hist)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"net.load_state_dict(torch.load(CKPT_PATH, map_location=DEVICE))\nnet.eval()\nprint('Best checkpoint loaded.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_pred, all_true = [], []\nwith torch.no_grad():\n    for imgs, lbls in val_dl:\n        all_pred.extend(net(imgs.to(DEVICE)).argmax(1).cpu().tolist())\n        all_true.extend(lbls.tolist())\n\nshort = [f'c{i}' for i in range(N_CLASSES)]\ncm    = confusion_matrix(all_true, all_pred)\n\nplt.figure(figsize=(10, 8))\nsns.heatmap(cm, annot=True, fmt='d', cmap='YlOrRd',\n            xticklabels=short, yticklabels=short, linewidths=0.4)\nplt.xlabel('Predicted label'); plt.ylabel('True label')\nplt.title('Confusion Matrix – Validation Set')\nplt.tight_layout(); plt.show()\n\nprint(classification_report(all_true, all_pred, target_names=short, digits=3))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n_, test_acc = run_epoch(net, val_dl)\nprint(f\"Test accuracy: {test_acc:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cc = [0] * N_CLASSES\nct = [0] * N_CLASSES\nfor p, t in zip(all_pred, all_true):\n    ct[t] += 1\n    if p == t: cc[t] += 1\n\nprint(f'  {\"Class\":<30s}  {\"Correct\":>7}  {\"Total\":>7}  {\"Acc\":>7}')\nprint('  ' + '-' * 55)\nfor i in range(N_CLASSES):\n    a = cc[i] / ct[i] if ct[i] else 0\n    print(f'  {CLASS_NAMES[i]:<30s}  {cc[i]:>7d}  {ct[i]:>7d}  {a*100:>6.2f}%')\nprint('  ' + '-' * 55)\nprint(f'  {\"OVERALL\":<30s}  {sum(cc):>7d}  {sum(ct):>7d}  {test_acc*100:>6.2f}%')\n\nprint()\nif test_acc >= 0.90:\n    print(' >=90% — BONUS achieved')\nelif test_acc >= 0.85:\n    print(' >=85% — Target achieved')\nelse:\n    print('   Below 85%')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}