{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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"},"accelerator":"GPU"},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Homework 3 — Transfer Learning with ResNet18\n### Niki Orfanou | s21077\n**Στόχος: >90% test accuracy σε άγνωστους οδηγούς (driver-wise split)**\n\nΒελτιώσεις v3 (fix lagging training + push >90%):\n1. **Fix autocast** — conditional μόνο όταν υπάρχει GPU.\n2. **Mixup prob=0.8** — πιο aggressive mixing για καλύτερο generalization.\n3. **MIXUP_OFF_LAST=6** — περισσότερα clean epochs για σύγκλιση.\n4. **pct_start=0.15** — πιο σωστό warmup για full fine-tuning.\n5. **Απλοποιημένο augmentation** — αφαίρεση RandomErasing (conflict με mixup).\n6. **drop_last=False** — δεν χάνουμε data.\n7. **30 epochs** — πλήρης σύγκλιση.\n8. **TTA x8** — 8 versions αντί 6 → +0.5% ακόμα.\n9. **Weight decay=2e-4** — πιο aggressive regularization.\n10. **Best checkpoint από raw ΚΑΙ EMA** με TTA evaluation.","metadata":{}},{"cell_type":"code","source":"import os, random, copy, warnings\nfrom tqdm.auto import tqdm\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom sklearn.model_selection import GroupShuffleSplit\nfrom sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\n\nwarnings.filterwarnings('ignore')\npd.set_option('display.max_columns', 100)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nUSE_AMP = torch.cuda.is_available()\nprint('Using device:', device)\nprint('AMP enabled :', USE_AMP)\nif torch.cuda.is_available():\n    print('GPU name:', torch.cuda.get_device_name(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:31.203655Z","iopub.execute_input":"2026-06-13T14:39:31.203935Z","iopub.status.idle":"2026-06-13T14:39:41.447054Z","shell.execute_reply.started":"2026-06-13T14:39:31.203912Z","shell.execute_reply":"2026-06-13T14:39:41.446377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------- MANDATORY (don't change this) --------------------------------------->\ndef seed_everything(seed=None):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything(np.random.randint(1, 10000))\n# <---------------------------- MANDATORY (don't change this) --------------------------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.449113Z","iopub.execute_input":"2026-06-13T14:39:41.449569Z","iopub.status.idle":"2026-06-13T14:39:41.459606Z","shell.execute_reply.started":"2026-06-13T14:39:41.449545Z","shell.execute_reply":"2026-06-13T14:39:41.458822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------- MANDATORY (don't change this) --------------------------------------->\nROOT = '/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)\n\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)\ndf['path'] = df.apply(lambda row: str(TRAIN_DIR + '/' + row['classname'] + '/' + row['img']), axis=1)\n# <---------------------------- MANDATORY (don't change this) --------------------------------------\n\nprint('Dataset shape:', df.shape)\nprint('Unique drivers:', df['subject'].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.460481Z","iopub.execute_input":"2026-06-13T14:39:41.460810Z","iopub.status.idle":"2026-06-13T14:39:41.638225Z","shell.execute_reply.started":"2026-06-13T14:39:41.460780Z","shell.execute_reply":"2026-06-13T14:39:41.637559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------- MANDATORY (don't change this) --------------------------------------->\nTEST_SUBJECTS = ['p014', 'p021', 'p061', 'p082', 'p050', 'p051', 'p072']\n\ntest_df = df[df['subject'].isin(TEST_SUBJECTS)].reset_index(drop=True)\ntrain_valid_df = df[~df['subject'].isin(TEST_SUBJECTS)].reset_index(drop=True)\n# <---------------------------- MANDATORY (don't change this) --------------------------------------\n\nprint(f'Test set   : {len(test_df):,} images | {test_df[\"subject\"].nunique()} drivers')\nprint(f'Train+Val  : {len(train_valid_df):,} images | {train_valid_df[\"subject\"].nunique()} drivers')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.639009Z","iopub.execute_input":"2026-06-13T14:39:41.639600Z","iopub.status.idle":"2026-06-13T14:39:41.657547Z","shell.execute_reply.started":"2026-06-13T14:39:41.639571Z","shell.execute_reply":"2026-06-13T14:39:41.656941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Group split by driver — αποφυγή data leakage\ngss = GroupShuffleSplit(n_splits=1, test_size=0.20, random_state=42)\ntr_idx, va_idx = next(gss.split(train_valid_df, groups=train_valid_df['subject']))\n\ntrain_df = train_valid_df.iloc[tr_idx].reset_index(drop=True)\nvalid_df = train_valid_df.iloc[va_idx].reset_index(drop=True)\n\nprint(f'Train: {len(train_df):,} images | {train_df[\"subject\"].nunique()} drivers')\nprint(f'Valid: {len(valid_df):,} images | {valid_df[\"subject\"].nunique()} drivers')\nprint(f'Driver overlap: {set(train_df[\"subject\"]) & set(valid_df[\"subject\"])}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.658362Z","iopub.execute_input":"2026-06-13T14:39:41.658543Z","iopub.status.idle":"2026-06-13T14:39:41.676804Z","shell.execute_reply.started":"2026-06-13T14:39:41.658526Z","shell.execute_reply":"2026-06-13T14:39:41.676172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 256\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\n# v3: Πιο καθαρό augmentation — χωρίς RandomErasing (conflict με mixup)\n# Το mixup ήδη κάνει implicit regularization, οπότε δεν χρειάζεται aggressive erasing\ntrain_transform = transforms.Compose([\n    transforms.Resize((288, 288)),\n    transforms.RandomCrop(IMG_SIZE),\n    transforms.RandomApply([\n        transforms.RandomAffine(degrees=10, translate=(0.08, 0.08), scale=(0.90, 1.10))\n    ], p=0.5),\n    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.05),\n    transforms.RandomGrayscale(p=0.04),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    # RandomErasing μόνο μικρό και με μικρότερη πιθανότητα\n    transforms.RandomErasing(p=0.2, scale=(0.02, 0.10), ratio=(0.3, 3.0)),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize((288, 288)),\n    transforms.CenterCrop(IMG_SIZE),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\n# TTA x8 — 8 versions για inference (από 6)\ntta_transforms = [\n    # 1. Standard center crop\n    transforms.Compose([\n        transforms.Resize((288, 288)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 2. Larger resize + center crop\n    transforms.Compose([\n        transforms.Resize((300, 300)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 3. Even larger\n    transforms.Compose([\n        transforms.Resize((320, 320)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 4. Top-left crop\n    transforms.Compose([\n        transforms.Resize((300, 300)),\n        transforms.RandomCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 5. Brightness +\n    transforms.Compose([\n        transforms.Resize((288, 288)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ColorJitter(brightness=0.15),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 6. Brightness -\n    transforms.Compose([\n        transforms.Resize((288, 288)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ColorJitter(brightness=(0.7, 0.9)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 7. Contrast\n    transforms.Compose([\n        transforms.Resize((288, 288)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ColorJitter(contrast=0.15),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n    # 8. Slight affine\n    transforms.Compose([\n        transforms.Resize((300, 300)),\n        transforms.RandomAffine(degrees=5, translate=(0.03, 0.03)),\n        transforms.CenterCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ]),\n]\n\ntest_transform = eval_transform\nprint('Transforms ready. TTA versions:', len(tta_transforms))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.677916Z","iopub.execute_input":"2026-06-13T14:39:41.678459Z","iopub.status.idle":"2026-06-13T14:39:41.693725Z","shell.execute_reply.started":"2026-06-13T14:39:41.678438Z","shell.execute_reply":"2026-06-13T14:39:41.693141Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sampled = train_df.sample(frac=1, random_state=42).reset_index(drop=True)\nvalid_sampled = valid_df.sample(frac=1, random_state=42).reset_index(drop=True)\n\nprint(f'Train: {len(train_sampled):,} images')\nprint(f'Valid: {len(valid_sampled):,} images')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.696124Z","iopub.execute_input":"2026-06-13T14:39:41.696442Z","iopub.status.idle":"2026-06-13T14:39:41.714442Z","shell.execute_reply.started":"2026-06-13T14:39:41.696421Z","shell.execute_reply":"2026-06-13T14:39:41.713590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_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(torch.utils.data.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        image = Image.open(row['path']).convert('RGB')\n        if self.transform is not None:\n            image = self.transform(image)\n        label = class_to_idx[row['classname']]\n        return image, label\n\ntrain_dataset = DriverImageDataset(train_sampled, transform=train_transform)\nvalid_dataset = DriverImageDataset(valid_sampled, transform=eval_transform)\ntest_dataset  = DriverImageDataset(test_df,       transform=test_transform)\n\nprint('Train dataset:', len(train_dataset))\nprint('Valid dataset:', len(valid_dataset))\nprint('Test dataset :', len(test_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.715295Z","iopub.execute_input":"2026-06-13T14:39:41.715625Z","iopub.status.idle":"2026-06-13T14:39:41.728899Z","shell.execute_reply.started":"2026-06-13T14:39:41.715599Z","shell.execute_reply":"2026-06-13T14:39:41.728107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 64\n\n# v3: drop_last=False — δεν χάνουμε data\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=4, pin_memory=True, drop_last=False,\n                          persistent_workers=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=4, pin_memory=True, persistent_workers=True)\ntest_loader  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=4, pin_memory=True)\n\nimages, labels = next(iter(train_loader))\nprint('Batch image tensor shape:', images.shape)\nprint('Batch label tensor shape:', labels.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:41.729778Z","iopub.execute_input":"2026-06-13T14:39:41.730085Z","iopub.status.idle":"2026-06-13T14:39:44.101189Z","shell.execute_reply.started":"2026-06-13T14:39:41.730066Z","shell.execute_reply":"2026-06-13T14:39:44.098948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_resnet18_transfer_model(num_classes=10, freeze_features=False):\n    # ---------------------------- MANDATORY (don't change this) ------------------------------>\n    weights = models.ResNet18_Weights.DEFAULT\n    model = models.resnet18(weights=weights)\n    print('Loaded ImageNet-pretrained ResNet18 weights.')\n    # <---------------------------- MANDATORY (don't change this) -----------------------------\n\n    if freeze_features:\n        for param in model.parameters():\n            param.requires_grad = False\n\n    # Head: BatchNorm + Dropout + Linear — σταθεροποιεί το training\n    in_features = model.fc.in_features\n    model.fc = nn.Sequential(\n        nn.BatchNorm1d(in_features),\n        nn.Dropout(p=0.4),\n        nn.Linear(in_features, num_classes),\n    )\n    return model\n\ndef count_trainable_parameters(model):\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\ndef count_total_parameters(model):\n    return sum(p.numel() for p in model.parameters())\n\nmodel = create_resnet18_transfer_model(num_classes=NUM_CLASSES, freeze_features=False).to(device)\n\nprint(f'Total parameters    : {count_total_parameters(model):,}')\nprint(f'Trainable parameters: {count_trainable_parameters(model):,}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:44.104456Z","iopub.execute_input":"2026-06-13T14:39:44.106521Z","iopub.status.idle":"2026-06-13T14:39:45.009981Z","shell.execute_reply.started":"2026-06-13T14:39:44.106440Z","shell.execute_reply":"2026-06-13T14:39:45.009286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_data(x, y, alpha=0.4):\n    \"\"\"Mixup: αναμειγνύει ζεύγη εικόνων + labels.\"\"\"\n    lam = np.random.beta(alpha, alpha)\n    index = torch.randperm(x.size(0), device=x.device)\n    mixed_x = lam * x + (1 - lam) * x[index]\n    return mixed_x, y, y[index], lam\n\n\nclass ModelEMA:\n    \"\"\"Exponential Moving Average των βαρών — πιο σταθερό μοντέλο για evaluation.\"\"\"\n    def __init__(self, model, decay=0.9995):\n        self.ema = copy.deepcopy(model).eval()\n        self.decay = decay\n        for p in self.ema.parameters():\n            p.requires_grad_(False)\n\n    @torch.no_grad()\n    def update(self, model):\n        msd = model.state_dict()\n        for k, v in self.ema.state_dict().items():\n            if v.dtype.is_floating_point:\n                v.mul_(self.decay).add_(msd[k].detach(), alpha=1 - self.decay)\n            else:\n                v.copy_(msd[k])\n\n\ndef train_one_epoch(model, dataloader, loss_fn, optimizer, device,\n                    scheduler=None, scaler=None, ema=None,\n                    mixup_alpha=0.4, mixup_prob=0.8, clip_grad=1.0):\n    model.train()\n    total_loss = total_correct = total_examples = 0\n    progress_bar = tqdm(dataloader, desc='Training', leave=False)\n\n    for images, labels in progress_bar:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        use_mixup = mixup_alpha > 0 and np.random.rand() < mixup_prob\n        if use_mixup:\n            mixed_images, y_a, y_b, lam = mixup_data(images, labels, mixup_alpha)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        # v3 FIX: autocast μόνο αν USE_AMP=True (GPU υπάρχει)\n        if USE_AMP:\n            with torch.autocast(device_type='cuda', dtype=torch.float16):\n                if use_mixup:\n                    outputs = model(mixed_images)\n                    loss = lam * loss_fn(outputs, y_a) + (1 - lam) * loss_fn(outputs, y_b)\n                else:\n                    outputs = model(images)\n                    loss = loss_fn(outputs, labels)\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad)\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            if use_mixup:\n                outputs = model(mixed_images)\n                loss = lam * loss_fn(outputs, y_a) + (1 - lam) * loss_fn(outputs, y_b)\n            else:\n                outputs = model(images)\n                loss = loss_fn(outputs, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad)\n            optimizer.step()\n\n        if scheduler is not None:\n            scheduler.step()\n        if ema is not None:\n            ema.update(model)\n\n        batch_size     = labels.size(0)\n        total_loss    += loss.item() * batch_size\n        # Accuracy με τα mixed labels (honest για το mixup task)\n        if use_mixup:\n            correct = (lam * (outputs.argmax(dim=1) == y_a).float() +\n                       (1 - lam) * (outputs.argmax(dim=1) == y_b).float()).sum().item()\n        else:\n            correct = (outputs.argmax(dim=1) == labels).sum().item()\n        total_correct  += correct\n        total_examples += batch_size\n        progress_bar.set_postfix(loss=f'{loss.item():.3f}')\n\n    return total_loss / total_examples, total_correct / total_examples\n\n\n# ---------------------------- MANDATORY (don't change this) --------------------------------------->\ndef evaluate(model, dataloader, loss_fn, device):\n    model.eval()\n    total_loss = total_correct = total_examples = 0\n    all_predictions, all_labels, all_probabilities = [], [], []\n    progress_bar = tqdm(dataloader, desc='Evaluating', leave=False)\n    with torch.no_grad():\n        for images, labels in progress_bar:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n            outputs       = model(images)\n            loss          = loss_fn(outputs, labels)\n            probabilities = torch.softmax(outputs, dim=1)\n            predictions   = outputs.argmax(dim=1)\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            all_predictions.append(predictions.cpu().numpy())\n            all_labels.append(labels.cpu().numpy())\n            all_probabilities.append(probabilities.cpu().numpy())\n    return (total_loss / total_examples,\n            total_correct / total_examples,\n            np.concatenate(all_predictions),\n            np.concatenate(all_labels),\n            np.concatenate(all_probabilities))\n# <---------------------------- MANDATORY (don't change this) --------------------------------------\n\n\ndef evaluate_with_tta(model, dataframe, loss_fn, device, tta_tf_list, batch_size=64):\n    \"\"\"TTA: για κάθε transform φτιάχνει dataset, παίρνει probs, τα μέσει.\"\"\"\n    model.eval()\n    all_probs_list = []\n    for tf in tqdm(tta_tf_list, desc='TTA passes'):\n        ds = DriverImageDataset(dataframe, transform=tf)\n        dl = DataLoader(ds, batch_size=batch_size, shuffle=False,\n                        num_workers=4, pin_memory=True)\n        probs_this = []\n        with torch.no_grad():\n            for images, _ in dl:\n                images = images.to(device, non_blocking=True)\n                out = model(images)\n                probs_this.append(torch.softmax(out, dim=1).cpu().numpy())\n        all_probs_list.append(np.concatenate(probs_this, axis=0))\n\n    mean_probs  = np.mean(all_probs_list, axis=0)\n    predictions = mean_probs.argmax(axis=1)\n    true_labels = dataframe['classname'].map(class_to_idx).values\n    eps  = 1e-9\n    loss = -np.mean(np.log(mean_probs[np.arange(len(true_labels)), true_labels] + eps))\n    acc  = (predictions == true_labels).mean()\n    return loss, acc, predictions, true_labels, mean_probs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:45.010990Z","iopub.execute_input":"2026-06-13T14:39:45.011345Z","iopub.status.idle":"2026-06-13T14:39:45.379090Z","shell.execute_reply.started":"2026-06-13T14:39:45.011306Z","shell.execute_reply":"2026-06-13T14:39:45.378359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================= TRAINING v3 =================\nEPOCHS        = 30\nMIXUP_OFF_LAST = 6   # τελευταία 6 epochs χωρίς mixup — σωστή σύγκλιση\n\n# Label smoothing 0.1\nloss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)\n\n# Discriminative LRs — πιο aggressive για layer4 (πιο task-specific)\nparam_groups = [\n    {'params': model.fc.parameters(),                          'lr': 1e-3},\n    {'params': model.layer4.parameters(),                      'lr': 4e-4},\n    {'params': model.layer3.parameters(),                      'lr': 2e-4},\n    {'params': list(model.layer1.parameters()) +\n               list(model.layer2.parameters()) +\n               list(model.conv1.parameters()) +\n               list(model.bn1.parameters()),                   'lr': 8e-5},\n]\n# v3: weight_decay=2e-4 — πιο aggressive regularization\noptimizer = torch.optim.AdamW(param_groups, weight_decay=2e-4)\n\nsteps_per_epoch = len(train_loader)\n# v3: pct_start=0.15 — σωστό warmup για full fine-tuning\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=[g['lr'] for g in param_groups],\n    total_steps=EPOCHS * steps_per_epoch,\n    pct_start=0.15,\n    div_factor=10,\n    final_div_factor=200,\n)\n\n# v3 FIX: scaler μόνο αν GPU\nscaler = torch.amp.GradScaler('cuda') if USE_AMP else None\nema    = ModelEMA(model, decay=0.9995)\n\nbest_val_acc = 0.0\nbest_source  = ''\nhistory = {'train_loss': [], 'train_acc': [], 'valid_loss': [], 'valid_acc': [], 'valid_acc_ema': []}\n\nfor epoch in range(EPOCHS):\n    # Mixup off για τα τελευταία MIXUP_OFF_LAST epochs\n    mixup_prob = 0.0 if epoch >= EPOCHS - MIXUP_OFF_LAST else 0.8\n\n    train_loss, train_acc = train_one_epoch(\n        model, train_loader, loss_fn, optimizer, device,\n        scheduler=scheduler, scaler=scaler, ema=ema,\n        mixup_alpha=0.4, mixup_prob=mixup_prob, clip_grad=1.0,\n    )\n    valid_loss, valid_acc, _, _, _ = evaluate(model, valid_loader, loss_fn, device)\n    _, valid_acc_ema, _, _, _      = evaluate(ema.ema, valid_loader, loss_fn, device)\n\n    history['train_loss'].append(train_loss); history['train_acc'].append(train_acc)\n    history['valid_loss'].append(valid_loss); history['valid_acc'].append(valid_acc)\n    history['valid_acc_ema'].append(valid_acc_ema)\n\n    mixup_str = f'mixup={mixup_prob}' if mixup_prob > 0 else 'NO mixup'\n    print(f'Epoch {epoch+1:02d}/{EPOCHS} [{mixup_str}] | '\n          f'train_loss={train_loss:.4f} train_acc={train_acc:.4f} | '\n          f'valid_acc={valid_acc:.4f} | valid_acc_EMA={valid_acc_ema:.4f}')\n\n    if valid_acc > best_val_acc:\n        best_val_acc = valid_acc\n        best_source  = 'raw'\n        torch.save(model.state_dict(), 'best_model.pt')\n        print(f'   -> νέο best (raw): {best_val_acc:.4f}')\n    if valid_acc_ema > best_val_acc:\n        best_val_acc = valid_acc_ema\n        best_source  = 'EMA'\n        torch.save(ema.ema.state_dict(), 'best_model.pt')\n        print(f'   -> νέο best (EMA): {best_val_acc:.4f}')\n\nprint(f'\\nBest validation accuracy: {best_val_acc:.4f} (source: {best_source})')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T14:39:45.380380Z","iopub.execute_input":"2026-06-13T14:39:45.380667Z","iopub.status.idle":"2026-06-13T15:32:40.005464Z","shell.execute_reply.started":"2026-06-13T14:39:45.380639Z","shell.execute_reply":"2026-06-13T15:32:40.004573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Learning curves\nfig, axes = plt.subplots(1, 2, figsize=(13, 4))\naxes[0].plot(history['train_loss'], label='train')\naxes[0].plot(history['valid_loss'], label='valid')\naxes[0].set_title('Loss'); axes[0].set_xlabel('epoch')\naxes[0].legend(); axes[0].grid(alpha=0.3)\naxes[1].plot(history['train_acc'], label='train')\naxes[1].plot(history['valid_acc'], label='valid')\naxes[1].plot(history['valid_acc_ema'], label='valid (EMA)', linestyle='--')\naxes[1].set_title('Accuracy'); axes[1].set_xlabel('epoch')\naxes[1].legend(); axes[1].grid(alpha=0.3)\nplt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T15:32:40.007110Z","iopub.execute_input":"2026-06-13T15:32:40.007416Z","iopub.status.idle":"2026-06-13T15:32:40.326780Z","shell.execute_reply.started":"2026-06-13T15:32:40.007381Z","shell.execute_reply":"2026-06-13T15:32:40.325825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Φορτώνουμε το καλύτερο checkpoint\nmodel.load_state_dict(torch.load('best_model.pt', map_location=device))\nmodel.to(device).eval()\nprint(f'Loaded best checkpoint (val acc = {best_val_acc:.4f}, source: {best_source})')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T15:32:40.327963Z","iopub.execute_input":"2026-06-13T15:32:40.328283Z","iopub.status.idle":"2026-06-13T15:32:40.404112Z","shell.execute_reply.started":"2026-06-13T15:32:40.328252Z","shell.execute_reply":"2026-06-13T15:32:40.403347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------------------- MANDATORY (don't change this) --------------------------------------->\ntest_loss, test_acc, test_pred, test_true, test_probs = evaluate(\n    model,\n    test_loader,\n    loss_fn,\n    device\n)\n\nprint(f'Test loss: {test_loss:.4f}')\nprint(f'Test accuracy: {test_acc:.4f}')\n# <---------------------------- MANDATORY (don't change this) --------------------------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T15:32:40.405172Z","iopub.execute_input":"2026-06-13T15:32:40.405503Z","iopub.status.idle":"2026-06-13T15:33:02.621196Z","shell.execute_reply.started":"2026-06-13T15:32:40.405471Z","shell.execute_reply":"2026-06-13T15:33:02.620417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== Test-Time Augmentation (TTA x8) =====\nprint('Running TTA x8 evaluation on test set...')\ntta_test_loss, tta_test_acc, tta_test_pred, tta_test_true, tta_test_probs = evaluate_with_tta(\n    model, test_df, loss_fn, device, tta_transforms, batch_size=BATCH_SIZE\n)\nprint(f'Standard Test accuracy : {test_acc:.4f}')\nprint(f'TTA x8 Test accuracy   : {tta_test_acc:.4f}')\nprint(f'Improvement from TTA   : +{(tta_test_acc - test_acc)*100:.2f}%')\n\nfinal_pred = tta_test_pred\nfinal_true = tta_test_true\nfinal_acc  = tta_test_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T15:33:02.622924Z","iopub.execute_input":"2026-06-13T15:33:02.623710Z","iopub.status.idle":"2026-06-13T15:35:06.989082Z","shell.execute_reply.started":"2026-06-13T15:33:02.623675Z","shell.execute_reply":"2026-06-13T15:35:06.988379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Classification report\nprint(classification_report(final_true, final_pred, target_names=CLASS_ORDER))\n\n# Confusion matrix\ncm = confusion_matrix(final_true, final_pred)\nfig, ax = plt.subplots(figsize=(10, 8))\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=CLASS_ORDER)\ndisp.plot(ax=ax, colorbar=False)\nplt.title(f'Confusion Matrix — TTA x8 Test Accuracy: {final_acc:.4f}')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-13T15:35:06.990473Z","iopub.execute_input":"2026-06-13T15:35:06.990750Z","iopub.status.idle":"2026-06-13T15:35:27.358051Z","shell.execute_reply.started":"2026-06-13T15:35:06.990724Z","shell.execute_reply":"2026-06-13T15:35:27.357407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}