{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":4922010,"datasetId":2854386,"databundleVersionId":4989769},{"sourceType":"modelInstanceVersion","sourceId":139093,"databundleVersionId":9881870,"modelInstanceId":117776,"modelId":141013},{"sourceType":"modelInstanceVersion","sourceId":845256,"databundleVersionId":16876209,"modelInstanceId":642869,"modelId":654862}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":33.126564,"end_time":"2026-04-24T16:45:04.100480+00:00","environment_variables":{},"exception":true,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-04-24T16:44:30.973916+00:00","version":"2.7.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Environment Setup","metadata":{"papermill":{"duration":0.007184,"end_time":"2026-04-24T16:44:33.746484+00:00","exception":false,"start_time":"2026-04-24T16:44:33.739300+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Ultralytics install through pip (Uncomment if run on Kaggle/Google Colab)\n!pip install ultralytics -q\n\n# Neccesity libraries\nimport json, random, shutil\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nimport yaml\nimport pathlib\nfrom pathlib import Path\nfrom types import SimpleNamespace\n\n# Visualization\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\n\n# Image processing\nimport cv2\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms.functional as TF\n\n# Ultralytics\nfrom ultralytics import YOLO\nfrom ultralytics.utils.loss import v8DetectionLoss\nfrom ultralytics.utils.nms import non_max_suppression\nfrom ultralytics.utils.metrics import DetMetrics, ConfusionMatrix, box_iou\n\n# PyTorch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\n\n# Training visualization\nfrom tqdm import tqdm\n\n# Ensure all notebook prints are flushed immediately (cleaner Kaggle logs).\nimport builtins\nfrom functools import partial\nprint = partial(builtins.print, flush=True)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:09.513086Z","iopub.execute_input":"2026-04-24T18:44:09.513474Z","iopub.status.idle":"2026-04-24T18:44:24.768571Z","shell.execute_reply.started":"2026-04-24T18:44:09.513425Z","shell.execute_reply":"2026-04-24T18:44:24.767695Z"},"papermill":{"duration":18.676822,"end_time":"2026-04-24T16:44:52.429313+00:00","exception":false,"start_time":"2026-04-24T16:44:33.752491+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Configuration","metadata":{"papermill":{"duration":0.005818,"end_time":"2026-04-24T16:44:52.442953+00:00","exception":false,"start_time":"2026-04-24T16:44:52.437135+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Check for CUDA\nDEVICE = 'cuda' if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:24.769904Z","iopub.execute_input":"2026-04-24T18:44:24.770227Z","iopub.status.idle":"2026-04-24T18:44:25.020888Z","shell.execute_reply.started":"2026-04-24T18:44:24.770199Z","shell.execute_reply":"2026-04-24T18:44:25.020047Z"},"papermill":{"duration":0.279444,"end_time":"2026-04-24T16:44:52.728298+00:00","exception":false,"start_time":"2026-04-24T16:44:52.448854+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nzeroshot: Experiment 1 only\nfinetune: Experiment 2 only\nscratch: Experiment 3 only\ncompare: Results comparison + charts (requires all 3 runs done)\nall: All 3 experiments + comparison (single notebook)\n'''\nEXPERIMENT_TO_RUN = \"scratch\"","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.022094Z","iopub.execute_input":"2026-04-24T18:44:25.022445Z","iopub.status.idle":"2026-04-24T18:44:25.050092Z","shell.execute_reply.started":"2026-04-24T18:44:25.022411Z","shell.execute_reply":"2026-04-24T18:44:25.049213Z"},"papermill":{"duration":0.012484,"end_time":"2026-04-24T16:44:52.747199+00:00","exception":false,"start_time":"2026-04-24T16:44:52.734715+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Path setup (local + Kaggle-safe)\nIS_KAGGLE = Path(\"/kaggle/working\").exists()\n\n# Dataset path (pick the first existing candidate)\n_dataset_candidates = [\n    Path(\"dataset\"),\n    Path(\"/kaggle/input/datasets/lakshaytyagi01/fruit-detection/Fruits-detection\"),\n]\nDATASET_ROOT = next((p for p in _dataset_candidates if p.exists()), _dataset_candidates[0])\n\n# Output/working dir (always create it to avoid missing-directory failures)\nOUTPUT = Path(\"/kaggle/working\") if IS_KAGGLE else Path(\"output\")\nRUNS_DIR = OUTPUT / \"runs\"\n\n# Model checkpoints: separate roots for loading and saving\n_model_ckp_load_candidates = [\n    OUTPUT.parent / \"ckp\",   # local sibling of output\n    Path(\"/kaggle/input/models/mnhlfuch/yolov8-fruit-detection/pytorch/default/1\")\n]\nMODEL_CKP_LOAD = next((p for p in _model_ckp_load_candidates if p.exists()), _model_ckp_load_candidates[0])\nMODEL_CKP_SAVE = OUTPUT.parent / \"ckp\" if not IS_KAGGLE else Path(\"/kaggle/working/ckp\")\n\nOUTPUT.mkdir(parents=True, exist_ok=True)\nRUNS_DIR.mkdir(parents=True, exist_ok=True)\nMODEL_CKP_SAVE.mkdir(parents=True, exist_ok=True)\n(MODEL_CKP_SAVE / \"finetune\").mkdir(parents=True, exist_ok=True)\n(MODEL_CKP_SAVE / \"scratch\").mkdir(parents=True, exist_ok=True)\n\n# Dataset YAML — stored in writable output dir\nDATASET_YAML = OUTPUT / \"fruit_fixed.yaml\"\n\n# YOLOv8 base checkpoint (pick first existing path)\n_ckpt_candidates = [\n    Path(\"yolov8_pre-trained/yolov8m.pt\"),\n    Path(\"/kaggle/input/models/ultralytics/yolov8/pytorch/default/1/yolov8m.pt\"),\n]\nYOLOV8_CKP = next((p for p in _ckpt_candidates if p.exists()), _ckpt_candidates[0])","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.051975Z","iopub.execute_input":"2026-04-24T18:44:25.052276Z","iopub.status.idle":"2026-04-24T18:44:25.068704Z","shell.execute_reply.started":"2026-04-24T18:44:25.052250Z","shell.execute_reply":"2026-04-24T18:44:25.067907Z"},"papermill":{"duration":0.021006,"end_time":"2026-04-24T16:44:52.774325+00:00","exception":false,"start_time":"2026-04-24T16:44:52.753319+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Class names (6 fruit classes) \nCLASS_NAMES = [\"Apple\", \"Banana\", \"Grapes\", \"Orange\", \"Pineapple\", \"Watermelon\"]","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.069701Z","iopub.execute_input":"2026-04-24T18:44:25.070099Z","iopub.status.idle":"2026-04-24T18:44:25.076283Z","shell.execute_reply.started":"2026-04-24T18:44:25.070072Z","shell.execute_reply":"2026-04-24T18:44:25.075430Z"},"papermill":{"duration":0.012773,"end_time":"2026-04-24T16:44:52.793745+00:00","exception":false,"start_time":"2026-04-24T16:44:52.780972+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  Image / Training hyper-parameters \nIMG_SIZE        = 640\nBATCH_SIZE      = 32\nEPOCHS_FINETUNE = 100\nEPOCHS_SCRATCH  = 100\nCONF_THRESHOLD  = 0.25\nIOU_THRESHOLD   = 0.5\n\n#  Fine-tuning hyper-parameters \nFT_LR0           = 0.005\nFT_LRF           = 0.1   # stronger LR decay for better convergence\nFT_MOMENTUM      = 0.937\nFT_WEIGHT_DECAY  = 0.0005\nFT_WARMUP_EPOCHS = 5\nFT_OPTIMIZER     = \"AdamW\"\nFT_PATIENCE      = 10\n\n#  Shared augmentation hyper-parameters (used by both fine-tuning and scratch)\nAUG_MOSAIC       = 1.0\nAUG_MIXUP        = 0.0   # dataset is already pre-augmented\nAUG_COPY_PASTE   = 0.0   # dataset is already pre-augmented\nAUG_SCALE        = 0.5\nAUG_TRANSLATE    = 0.1\nAUG_DEGREES      = 15.0\nAUG_FLIPLR       = 0.5\nAUG_FLIPUD       = 0.0   # fruits rarely appear upside-down\nAUG_HSV_H        = 0.02\nAUG_HSV_S        = 0.7\nAUG_HSV_V        = 0.4\n\n#  From-scratch hyper-parameters \nSC_LR0           = 0.01\nSC_LRF           = 0.1\nSC_MOMENTUM      = 0.937\nSC_WEIGHT_DECAY  = 0.0005\nSC_WARMUP_EPOCHS = 5\nSC_OPTIMIZER     = \"SGD\"\nSC_PATIENCE      = 10","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.077196Z","iopub.execute_input":"2026-04-24T18:44:25.077437Z","iopub.status.idle":"2026-04-24T18:44:25.090315Z","shell.execute_reply.started":"2026-04-24T18:44:25.077412Z","shell.execute_reply":"2026-04-24T18:44:25.089470Z"},"papermill":{"duration":0.015481,"end_time":"2026-04-24T16:44:52.815903+00:00","exception":false,"start_time":"2026-04-24T16:44:52.800422+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nThis cell is for fixing the YAML file in the dataset\nCreate a new YAML file to alter the dataset's YAML file\n'''\ndef stringify_paths(obj):\n    if isinstance(obj, pathlib.PurePath):\n        return str(obj)\n    elif isinstance(obj, dict):\n        return {k: stringify_paths(v) for k, v in obj.items()}\n    elif isinstance(obj, list):\n        return [stringify_paths(v) for v in obj]\n    return obj\n\n_orig = yaml.safe_load(open(f\"{DATASET_ROOT}/data.yaml\"))\n\nfixed = {\n    \"path\" : DATASET_ROOT,\n    \"train\": f\"{DATASET_ROOT}/train/images\",\n    \"val\"  : f\"{DATASET_ROOT}/valid/images\",\n    \"test\" : f\"{DATASET_ROOT}/test/images\",\n    \"nc\"   : len(CLASS_NAMES),\n    \"names\": CLASS_NAMES,\n}\n\nfixed = stringify_paths(fixed)\n\nDATASET_YAML.parent.mkdir(parents=True, exist_ok=True)\nwith open(DATASET_YAML, \"w\") as f:\n    yaml.safe_dump(fixed, f)\nprint(\"Fixed yaml written:\", DATASET_YAML)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.091207Z","iopub.execute_input":"2026-04-24T18:44:25.091706Z","iopub.status.idle":"2026-04-24T18:44:25.115549Z","shell.execute_reply.started":"2026-04-24T18:44:25.091650Z","shell.execute_reply":"2026-04-24T18:44:25.115008Z"},"papermill":{"duration":0.021254,"end_time":"2026-04-24T16:44:52.843486+00:00","exception":false,"start_time":"2026-04-24T16:44:52.822232+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Device            : {DEVICE}\")\nprint(f\"Experiment to run : {EXPERIMENT_TO_RUN}\")\nprint(f\"Dataset root      : {DATASET_ROOT}\")\nprint(f\"Output dir        : {OUTPUT}\")\nprint(f\"Runs dir          : {RUNS_DIR}\")\nprint(f\"Model CKP load    : {MODEL_CKP_LOAD}\")\nprint(f\"Model CKP save    : {MODEL_CKP_SAVE}\")\nprint(f\"Dataset YAML      : {DATASET_YAML}\")\nprint(f\"YOLOv8 checkpoint : {YOLOV8_CKP}\")\nprint(f\"Classes           : {CLASS_NAMES}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.116413Z","iopub.execute_input":"2026-04-24T18:44:25.116711Z","iopub.status.idle":"2026-04-24T18:44:25.129576Z","shell.execute_reply.started":"2026-04-24T18:44:25.116675Z","shell.execute_reply":"2026-04-24T18:44:25.128783Z"},"papermill":{"duration":0.022655,"end_time":"2026-04-24T16:44:52.872350+00:00","exception":false,"start_time":"2026-04-24T16:44:52.849695+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Dataset Preparation","metadata":{"papermill":{"duration":0.006312,"end_time":"2026-04-24T16:44:52.885463+00:00","exception":false,"start_time":"2026-04-24T16:44:52.879151+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Verify dataset path\ndataset_root = Path(DATASET_ROOT)\nassert dataset_root.exists(), f\"Dataset not found at {DATASET_ROOT}.\"","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.130442Z","iopub.execute_input":"2026-04-24T18:44:25.130829Z","iopub.status.idle":"2026-04-24T18:44:25.138007Z","shell.execute_reply.started":"2026-04-24T18:44:25.130801Z","shell.execute_reply":"2026-04-24T18:44:25.137336Z"},"papermill":{"duration":0.013609,"end_time":"2026-04-24T16:44:52.906031+00:00","exception":false,"start_time":"2026-04-24T16:44:52.892422+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Open YAML file\nwith open(DATASET_YAML) as f:\n    yaml_cfg = yaml.safe_load(f)\n\nprint(\"Data.yaml contents:\")\nfor k, v in yaml_cfg.items():\n    print(f\"  {k}: {v}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.140755Z","iopub.execute_input":"2026-04-24T18:44:25.141038Z","iopub.status.idle":"2026-04-24T18:44:25.157227Z","shell.execute_reply.started":"2026-04-24T18:44:25.141012Z","shell.execute_reply":"2026-04-24T18:44:25.156443Z"},"papermill":{"duration":0.02105,"end_time":"2026-04-24T16:44:52.933720+00:00","exception":false,"start_time":"2026-04-24T16:44:52.912670+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check for number of images in each sub dataset\nfor split in [\"train\", \"valid\", \"test\"]:\n    img_dir = dataset_root / split / \"images\"\n    n = len(glob(str(img_dir / \"*.jpg\")) + glob(str(img_dir / \"*.png\")))\n    print(f\"  {split}: {n} images\")","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.158279Z","iopub.execute_input":"2026-04-24T18:44:25.158634Z","iopub.status.idle":"2026-04-24T18:44:25.934932Z","shell.execute_reply.started":"2026-04-24T18:44:25.158582Z","shell.execute_reply":"2026-04-24T18:44:25.934199Z"},"papermill":{"duration":0.55854,"end_time":"2026-04-24T16:44:53.499673+00:00","exception":false,"start_time":"2026-04-24T16:44:52.941133+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualise a few training samples\ndef show_samples(img_dir, lbl_dir, class_names, n=4, figsize=(16, 8)):\n    imgs = glob(str(Path(img_dir) / \"*.jpg\")) + glob(str(Path(img_dir) / \"*.png\"))\n    imgs = random.sample(imgs, min(n, len(imgs)))\n\n    fig, axes = plt.subplots(1, len(imgs), figsize=figsize)\n    if len(imgs) == 1:\n        axes = [axes]\n\n    for ax, img_path in zip(axes, imgs):\n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        h, w = img.shape[:2]\n        lbl_path = Path(lbl_dir) / (Path(img_path).stem + \".txt\")\n        ax.imshow(img)\n        if lbl_path.exists():\n            with open(lbl_path) as f:\n                for line in f:\n                    parts = line.strip().split()\n                    if len(parts) < 5:\n                        continue\n                    cls_id, cx, cy, bw, bh = int(parts[0]), *map(float, parts[1:])\n                    x1 = (cx - bw / 2) * w\n                    y1 = (cy - bh / 2) * h\n                    rect = patches.Rectangle(\n                        (x1, y1), bw * w, bh * h,\n                        linewidth=2, edgecolor=\"lime\", facecolor=\"none\")\n                    ax.add_patch(rect)\n                    label = class_names[cls_id] if cls_id < len(class_names) else str(cls_id)\n                    ax.text(x1, y1 - 4, label, color=\"lime\", fontsize=8,\n                            bbox=dict(facecolor=\"black\", alpha=0.4, pad=1))\n        ax.axis(\"off\")\n        ax.set_title(Path(img_path).name, fontsize=8)\n\n    plt.suptitle(\"Fruit Detection — Sample Training Images\", fontsize=13, y=1.02)\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.935930Z","iopub.execute_input":"2026-04-24T18:44:25.936284Z","iopub.status.idle":"2026-04-24T18:44:25.945000Z","shell.execute_reply.started":"2026-04-24T18:44:25.936257Z","shell.execute_reply":"2026-04-24T18:44:25.944309Z"},"papermill":{"duration":0.01991,"end_time":"2026-04-24T16:44:53.526873+00:00","exception":false,"start_time":"2026-04-24T16:44:53.506963+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_samples(\n    img_dir     = dataset_root / \"train\" / \"images\",\n    lbl_dir     = dataset_root / \"train\" / \"labels\",\n    class_names = CLASS_NAMES,\n)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:25.946568Z","iopub.execute_input":"2026-04-24T18:44:25.946946Z","iopub.status.idle":"2026-04-24T18:44:26.776880Z","shell.execute_reply.started":"2026-04-24T18:44:25.946896Z","shell.execute_reply":"2026-04-24T18:44:26.776015Z"},"papermill":{"duration":0.898916,"end_time":"2026-04-24T16:44:54.432966+00:00","exception":false,"start_time":"2026-04-24T16:44:53.534050+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  PyTorch Dataset & DataLoader \nclass FruitDetectionDataset(Dataset):\n    \"\"\"\n    Minimal PyTorch Dataset for YOLO-format fruit detection data.\n    Returns (image_tensor, labels) where labels is a list of\n    [class_id, cx, cy, w, h] rows (normalised, as stored on disk).\n    \"\"\"\n\n    def __init__(self, img_dir: str, lbl_dir: str, img_size: int = IMG_SIZE, augment: bool = False, aug_cfg: dict = None):\n        self.img_paths = sorted(\n            glob(str(Path(img_dir) / \"*.jpg\")) +\n            glob(str(Path(img_dir) / \"*.png\"))\n        )\n        self.lbl_dir  = Path(lbl_dir)\n        self.img_size = img_size\n        self.augment  = augment\n        self.aug_cfg  = aug_cfg or {}\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def _load_image_and_labels(self, idx):\n        img_path = self.img_paths[idx]\n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, (self.img_size, self.img_size))\n\n        lbl_path = self.lbl_dir / (Path(img_path).stem + \".txt\")\n        labels = []\n        if lbl_path.exists():\n            with open(lbl_path) as f:\n                for line in f:\n                    parts = line.strip().split()\n                    if len(parts) == 5:\n                        labels.append([float(x) for x in parts])\n        return img, labels\n\n    def _xywhn_to_xyxy(self, labels):\n        boxes = []\n        s = float(self.img_size)\n        for cls_id, cx, cy, bw, bh in labels:\n            x1 = (cx - bw / 2.0) * s\n            y1 = (cy - bh / 2.0) * s\n            x2 = (cx + bw / 2.0) * s\n            y2 = (cy + bh / 2.0) * s\n            boxes.append([float(cls_id), x1, y1, x2, y2])\n        return boxes\n\n    def _xyxy_to_xywhn(self, boxes):\n        labels = []\n        s = float(self.img_size)\n        for cls_id, x1, y1, x2, y2 in boxes:\n            x1 = np.clip(x1, 0, s)\n            y1 = np.clip(y1, 0, s)\n            x2 = np.clip(x2, 0, s)\n            y2 = np.clip(y2, 0, s)\n            if x2 <= x1 or y2 <= y1:\n                continue\n            cx = ((x1 + x2) / 2.0) / s\n            cy = ((y1 + y2) / 2.0) / s\n            bw = (x2 - x1) / s\n            bh = (y2 - y1) / s\n            labels.append([float(cls_id), cx, cy, bw, bh])\n        return labels\n\n    def _mosaic_augment(self, idx):\n        s = self.img_size\n        mosaic_img = np.full((2 * s, 2 * s, 3), 114, dtype=np.uint8)\n        labels_out = []\n\n        indices = [idx] + random.choices(range(len(self.img_paths)), k=3)\n        positions = [(0, 0), (0, s), (s, 0), (s, s)]\n\n        for sample_idx, (y0, x0) in zip(indices, positions):\n            img_i, labels_i = self._load_image_and_labels(sample_idx)\n            mosaic_img[y0:y0 + s, x0:x0 + s] = img_i\n\n            for cls_id, cx, cy, bw, bh in labels_i:\n                x1 = (cx - bw / 2.0) * s + x0\n                y1 = (cy - bh / 2.0) * s + y0\n                x2 = (cx + bw / 2.0) * s + x0\n                y2 = (cy + bh / 2.0) * s + y0\n\n                x1 = np.clip(x1, 0, 2 * s)\n                y1 = np.clip(y1, 0, 2 * s)\n                x2 = np.clip(x2, 0, 2 * s)\n                y2 = np.clip(y2, 0, 2 * s)\n                if x2 <= x1 or y2 <= y1:\n                    continue\n\n                cx_n = ((x1 + x2) / 2.0) / (2.0 * s)\n                cy_n = ((y1 + y2) / 2.0) / (2.0 * s)\n                bw_n = (x2 - x1) / (2.0 * s)\n                bh_n = (y2 - y1) / (2.0 * s)\n                labels_out.append([float(cls_id), cx_n, cy_n, bw_n, bh_n])\n\n        mosaic_img = cv2.resize(mosaic_img, (s, s))\n        return mosaic_img, labels_out\n\n    def _random_affine(self, img, labels):\n        degrees = self.aug_cfg.get(\"degrees\", 0.0)\n        scale = self.aug_cfg.get(\"scale\", 0.0)\n        translate = self.aug_cfg.get(\"translate\", 0.0)\n        if degrees <= 0 and scale <= 0 and translate <= 0:\n            return img, labels\n\n        s = self.img_size\n        angle = random.uniform(-degrees, degrees)\n        zoom = random.uniform(max(0.1, 1.0 - scale), 1.0 + scale)\n        tx = random.uniform(-translate, translate) * s\n        ty = random.uniform(-translate, translate) * s\n\n        m = cv2.getRotationMatrix2D((s / 2.0, s / 2.0), angle, zoom)\n        m[:, 2] += [tx, ty]\n        img_warp = cv2.warpAffine(img, m, (s, s), flags=cv2.INTER_LINEAR, borderValue=(114, 114, 114))\n\n        boxes = self._xywhn_to_xyxy(labels)\n        if not boxes:\n            return img_warp, labels\n\n        boxes_out = []\n        for cls_id, x1, y1, x2, y2 in boxes:\n            corners = np.array([[x1, y1], [x2, y1], [x2, y2], [x1, y2]], dtype=np.float32)\n            ones = np.ones((4, 1), dtype=np.float32)\n            warped = np.hstack((corners, ones)) @ m.T\n            nx1, ny1 = warped[:, 0].min(), warped[:, 1].min()\n            nx2, ny2 = warped[:, 0].max(), warped[:, 1].max()\n            boxes_out.append([cls_id, nx1, ny1, nx2, ny2])\n\n        return img_warp, self._xyxy_to_xywhn(boxes_out)\n\n    def _copy_paste(self, img, labels):\n        p = self.aug_cfg.get(\"copy_paste\", 0.0)\n        if p <= 0 or random.random() >= p:\n            return img, labels\n\n        donor_idx = random.randrange(len(self.img_paths))\n        donor_img, donor_labels = self._load_image_and_labels(donor_idx)\n        donor_boxes = self._xywhn_to_xyxy(donor_labels)\n        if not donor_boxes:\n            return img, labels\n\n        out = img.copy()\n        out_boxes = self._xywhn_to_xyxy(labels)\n\n        cls_id, x1, y1, x2, y2 = random.choice(donor_boxes)\n        x1, y1, x2, y2 = map(int, [x1, y1, x2, y2])\n        x1 = np.clip(x1, 0, self.img_size - 1)\n        y1 = np.clip(y1, 0, self.img_size - 1)\n        x2 = np.clip(x2, x1 + 1, self.img_size)\n        y2 = np.clip(y2, y1 + 1, self.img_size)\n        patch = donor_img[y1:y2, x1:x2]\n        if patch.size == 0:\n            return img, labels\n\n        ph, pw = patch.shape[:2]\n        max_x = max(0, self.img_size - pw)\n        max_y = max(0, self.img_size - ph)\n        px = random.randint(0, max_x) if max_x > 0 else 0\n        py = random.randint(0, max_y) if max_y > 0 else 0\n        out[py:py + ph, px:px + pw] = patch\n        out_boxes.append([cls_id, px, py, px + pw, py + ph])\n\n        return out, self._xyxy_to_xywhn(out_boxes)\n\n    def _mixup(self, img, labels):\n        p = self.aug_cfg.get(\"mixup\", 0.0)\n        if p <= 0 or random.random() >= p:\n            return img, labels\n\n        donor_idx = random.randrange(len(self.img_paths))\n        img2, labels2 = self._load_image_and_labels(donor_idx)\n        lam = 0.5\n        mixed = (img.astype(np.float32) * lam + img2.astype(np.float32) * (1.0 - lam)).astype(np.uint8)\n        return mixed, labels + labels2\n\n    def _augment_image_and_labels(self, img, labels):\n        fliplr_p = self.aug_cfg.get(\"fliplr\", 0.5)\n        flipud_p = self.aug_cfg.get(\"flipud\", 0.0)\n        hsv_h = self.aug_cfg.get(\"hsv_h\", 0.015)\n        hsv_s = self.aug_cfg.get(\"hsv_s\", 0.7)\n        hsv_v = self.aug_cfg.get(\"hsv_v\", 0.4)\n\n        img, labels = self._random_affine(img, labels)\n\n        if random.random() < fliplr_p:\n            img = cv2.flip(img, 1)\n            for row in labels:\n                row[1] = 1.0 - row[1]\n\n        if random.random() < flipud_p:\n            img = cv2.flip(img, 0)\n            for row in labels:\n                row[2] = 1.0 - row[2]\n\n        if random.random() < 0.5:\n            hsv = cv2.cvtColor(img, cv2.COLOR_RGB2HSV).astype(np.float32)\n            hsv[..., 0] = (hsv[..., 0] + random.uniform(-hsv_h * 180.0, hsv_h * 180.0)) % 180.0\n            hsv[..., 1] = np.clip(hsv[..., 1] * random.uniform(1.0 - hsv_s * 0.5, 1.0 + hsv_s * 0.5), 0, 255)\n            hsv[..., 2] = np.clip(hsv[..., 2] * random.uniform(1.0 - hsv_v * 0.5, 1.0 + hsv_v * 0.5), 0, 255)\n            img = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB)\n\n        img, labels = self._copy_paste(img, labels)\n        img, labels = self._mixup(img, labels)\n        return img, labels\n\n    def __getitem__(self, idx):\n        mosaic_p = self.aug_cfg.get(\"mosaic\", 0.0)\n\n        if self.augment and random.random() < mosaic_p:\n            img, labels = self._mosaic_augment(idx)\n        else:\n            img, labels = self._load_image_and_labels(idx)\n\n        if self.augment:\n            img, labels = self._augment_image_and_labels(img, labels)\n\n        # HWC uint8  →  CHW float [0, 1]\n        img_tensor = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n\n        return img_tensor, labels","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:26.777964Z","iopub.execute_input":"2026-04-24T18:44:26.778295Z","iopub.status.idle":"2026-04-24T18:44:26.807609Z","shell.execute_reply.started":"2026-04-24T18:44:26.778268Z","shell.execute_reply":"2026-04-24T18:44:26.806742Z"},"papermill":{"duration":0.054902,"end_time":"2026-04-24T16:44:54.506182+00:00","exception":false,"start_time":"2026-04-24T16:44:54.451280+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def yolo_collate_fn(batch):\n    imgs, labels = zip(*batch)\n    imgs = torch.stack(imgs)          # (B, 3, H, W)\n\n    cls_list, box_list, idx_list = [], [], []\n    for i, sample_labels in enumerate(labels):\n        for row in sample_labels:     # each row: [cls, cx, cy, w, h]\n            cls_list.append(row[0])\n            box_list.append(row[1:])\n            idx_list.append(i)\n\n    return {\n        \"img\":      imgs,\n        \"cls\":      torch.tensor(cls_list,  dtype=torch.float32),\n        \"bboxes\":   torch.tensor(box_list,  dtype=torch.float32),\n        \"batch_idx\": torch.tensor(idx_list, dtype=torch.float32),\n    }","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:26.808672Z","iopub.execute_input":"2026-04-24T18:44:26.809125Z","iopub.status.idle":"2026-04-24T18:44:26.826034Z","shell.execute_reply.started":"2026-04-24T18:44:26.809084Z","shell.execute_reply":"2026-04-24T18:44:26.825099Z"},"papermill":{"duration":0.024571,"end_time":"2026-04-24T16:44:54.547526+00:00","exception":false,"start_time":"2026-04-24T16:44:54.522955+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build loaders for each split\nscratch_aug_cfg = {\n    \"mosaic\": AUG_MOSAIC,\n    \"mixup\": AUG_MIXUP,\n    \"copy_paste\": AUG_COPY_PASTE,\n    \"scale\": AUG_SCALE,\n    \"translate\": AUG_TRANSLATE,\n    \"degrees\": AUG_DEGREES,\n    \"fliplr\": AUG_FLIPLR,\n    \"flipud\": AUG_FLIPUD,\n    \"hsv_h\": AUG_HSV_H,\n    \"hsv_s\": AUG_HSV_S,\n    \"hsv_v\": AUG_HSV_V,\n}\n\ntrain_dataset = FruitDetectionDataset(\n    img_dir = dataset_root / \"train\" / \"images\",\n    lbl_dir = dataset_root / \"train\" / \"labels\",\n    augment = True,\n    aug_cfg = scratch_aug_cfg,\n)\nval_dataset = FruitDetectionDataset(\n    img_dir = dataset_root / \"valid\" / \"images\",\n    lbl_dir = dataset_root / \"valid\" / \"labels\",\n    augment = False,\n)\ntest_dataset = FruitDetectionDataset(\n    img_dir = dataset_root / \"test\" / \"images\",\n    lbl_dir = dataset_root / \"test\" / \"labels\",\n    augment = False,\n)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:26.826849Z","iopub.execute_input":"2026-04-24T18:44:26.827430Z","iopub.status.idle":"2026-04-24T18:44:26.867920Z","shell.execute_reply.started":"2026-04-24T18:44:26.827404Z","shell.execute_reply":"2026-04-24T18:44:26.867372Z"},"papermill":{"duration":0.054592,"end_time":"2026-04-24T16:44:54.618741+00:00","exception":false,"start_time":"2026-04-24T16:44:54.564149+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=2, pin_memory=True, collate_fn=yolo_collate_fn)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=2, pin_memory=True, collate_fn=yolo_collate_fn)\ntest_loader  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=2, pin_memory=True, collate_fn=yolo_collate_fn)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:26.868647Z","iopub.execute_input":"2026-04-24T18:44:26.868994Z","iopub.status.idle":"2026-04-24T18:44:26.873392Z","shell.execute_reply.started":"2026-04-24T18:44:26.868968Z","shell.execute_reply":"2026-04-24T18:44:26.872789Z"},"papermill":{"duration":0.024015,"end_time":"2026-04-24T16:44:54.659552+00:00","exception":false,"start_time":"2026-04-24T16:44:54.635537+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Train : {len(train_dataset):,} images  ({len(train_loader)} batches)\")\nprint(f\"Val   : {len(val_dataset):,} images  ({len(val_loader)} batches)\")\nprint(f\"Test  : {len(test_dataset):,} images  ({len(test_loader)} batches)\")\n\n\ndef _count_label_files(split_name: str):\n    img_dir = dataset_root / split_name / \"images\"\n    lbl_dir = dataset_root / split_name / \"labels\"\n    n_img = len(list(img_dir.glob(\"*.jpg\"))) + len(list(img_dir.glob(\"*.png\")))\n    n_lbl = len(list(lbl_dir.glob(\"*.txt\"))) if lbl_dir.exists() else 0\n    return n_img, n_lbl, img_dir, lbl_dir\n\n\nprint(\"\\nLabel sanity check:\")\nfor split in [\"train\", \"valid\", \"test\"]:\n    n_img, n_lbl, img_dir, lbl_dir = _count_label_files(split)\n    status = \"OK\" if n_lbl > 0 else \"MISSING_LABELS\"\n    print(f\"  {split:<5} images={n_img:<5} labels={n_lbl:<5} [{status}]\")\n    print(f\"       img_dir={img_dir}\")\n    print(f\"       lbl_dir={lbl_dir}\")\n\nif _count_label_files(\"test\")[1] == 0:\n    print(\"\\nWARNING: test labels are missing. Zero-shot/fine-tune test metrics will be 0.0 until labels exist.\")","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:26.874372Z","iopub.execute_input":"2026-04-24T18:44:26.874676Z","iopub.status.idle":"2026-04-24T18:44:27.599336Z","shell.execute_reply.started":"2026-04-24T18:44:26.874651Z","shell.execute_reply":"2026-04-24T18:44:27.598767Z"},"papermill":{"duration":0.535088,"end_time":"2026-04-24T16:44:55.211788+00:00","exception":false,"start_time":"2026-04-24T16:44:54.676700+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Quick visual sanity-check for scratch augmentations\n# Shows samples from train_dataset *after* augmentations (mosaic/mixup/copy-paste/etc.)\ndef show_augmented_scratch_samples(dataset, class_names, n=6, cols=3, figsize=(14, 9)):\n    n = max(1, min(n, len(dataset)))\n    cols = max(1, cols)\n    rows = int(np.ceil(n / cols))\n\n    fig, axes = plt.subplots(rows, cols, figsize=figsize)\n    axes = np.array(axes).reshape(-1)\n\n    for i in range(n):\n        ax = axes[i]\n        sample_idx = random.randrange(len(dataset))\n        img_tensor, labels = dataset[sample_idx]\n\n        img = (img_tensor.permute(1, 2, 0).numpy() * 255.0).astype(np.uint8)\n        h, w = img.shape[:2]\n\n        ax.imshow(img)\n        for row in labels:\n            cls_id, cx, cy, bw, bh = row\n            x1 = (cx - bw / 2.0) * w\n            y1 = (cy - bh / 2.0) * h\n            rect = patches.Rectangle(\n                (x1, y1), bw * w, bh * h,\n                linewidth=1.8, edgecolor=\"yellow\", facecolor=\"none\"\n            )\n            ax.add_patch(rect)\n            label = class_names[int(cls_id)] if int(cls_id) < len(class_names) else str(int(cls_id))\n            ax.text(\n                x1, max(0, y1 - 4), label,\n                color=\"yellow\", fontsize=8,\n                bbox=dict(facecolor=\"black\", alpha=0.45, pad=1)\n            )\n\n        ax.set_title(f\"Aug sample #{i+1} (boxes={len(labels)})\", fontsize=9)\n        ax.axis(\"off\")\n\n    for j in range(n, len(axes)):\n        axes[j].axis(\"off\")\n\n    plt.suptitle(\"Scratch Augmentation Preview\", fontsize=13)\n    plt.tight_layout()\n    plt.show()\n\n\nshow_augmented_scratch_samples(train_dataset, CLASS_NAMES, n=6, cols=3)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:27.600326Z","iopub.execute_input":"2026-04-24T18:44:27.600640Z","iopub.status.idle":"2026-04-24T18:44:29.658889Z","shell.execute_reply.started":"2026-04-24T18:44:27.600613Z","shell.execute_reply":"2026-04-24T18:44:29.657869Z"},"papermill":{"duration":1.94484,"end_time":"2026-04-24T16:44:57.174431+00:00","exception":false,"start_time":"2026-04-24T16:44:55.229591+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Metrics Helper Methods","metadata":{"papermill":{"duration":0.035803,"end_time":"2026-04-24T16:44:57.246669+00:00","exception":false,"start_time":"2026-04-24T16:44:57.210866+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def extract_metrics(results, label: str) -> dict:\n    mp      = float(results.box.mp)\n    mr      = float(results.box.mr)\n    map50   = float(results.box.map50)\n    map5095 = float(results.box.map)\n    f1 = (2 * mp * mr / (mp + mr)) if (mp + mr) > 0 else 0.0\n\n    metrics = {\n        \"Experiment\"   : label,\n        \"Precision\"    : round(mp,      4),\n        \"Recall\"       : round(mr,      4),\n        \"F1\"           : round(f1,      4),\n        \"mAP@0.5\"      : round(map50,   4),\n        \"mAP@0.5:0.95\" : round(map5095, 4),\n    }\n    print(f\"\\n{'='*50}\\n  {label}\\n{'='*50}\")\n    for k, v in metrics.items():\n        if k != \"Experiment\":\n            print(f\"  {k:<18}: {v}\")\n    return metrics","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.660103Z","iopub.execute_input":"2026-04-24T18:44:29.660329Z","iopub.status.idle":"2026-04-24T18:44:29.666673Z","shell.execute_reply.started":"2026-04-24T18:44:29.660306Z","shell.execute_reply":"2026-04-24T18:44:29.665914Z"},"papermill":{"duration":0.042396,"end_time":"2026-04-24T16:44:57.321158+00:00","exception":false,"start_time":"2026-04-24T16:44:57.278762+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_training_curves(csv_path, title, color):\n    if not Path(csv_path).exists():\n        print(f\"results.csv not found at {csv_path}\")\n        return None\n    df = pd.read_csv(csv_path)\n    df.columns = df.columns.str.strip()\n    plots = [\n        (\"train/box_loss\",       \"Train Box Loss\"),\n        (\"train/cls_loss\",       \"Train Cls Loss\"),\n        (\"val/box_loss\",         \"Val Box Loss\"),\n        (\"metrics/precision(B)\", \"Precision\"),\n        (\"metrics/recall(B)\",    \"Recall\"),\n        (\"metrics/mAP50(B)\",     \"mAP@0.5\"),\n    ]\n    fig, axes = plt.subplots(2, 3, figsize=(16, 8))\n    for ax, (col, lbl) in zip(axes.flat, plots):\n        if col in df.columns:\n            ax.plot(df[\"epoch\"], df[col], linewidth=2, color=color)\n            ax.set_title(lbl, fontsize=11)\n            ax.set_xlabel(\"Epoch\")\n            ax.grid(alpha=0.3)\n    plt.suptitle(title, fontsize=14, y=1.01)\n    plt.tight_layout()\n    plt.show()\n    return df","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.667678Z","iopub.execute_input":"2026-04-24T18:44:29.668127Z","iopub.status.idle":"2026-04-24T18:44:29.688378Z","shell.execute_reply.started":"2026-04-24T18:44:29.668083Z","shell.execute_reply":"2026-04-24T18:44:29.687700Z"},"papermill":{"duration":0.040608,"end_time":"2026-04-24T16:44:57.393077+00:00","exception":false,"start_time":"2026-04-24T16:44:57.352469+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_metrics = []","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.689251Z","iopub.execute_input":"2026-04-24T18:44:29.689548Z","iopub.status.idle":"2026-04-24T18:44:29.702286Z","shell.execute_reply.started":"2026-04-24T18:44:29.689525Z","shell.execute_reply":"2026-04-24T18:44:29.701499Z"},"papermill":{"duration":0.040895,"end_time":"2026-04-24T16:44:57.466744+00:00","exception":false,"start_time":"2026-04-24T16:44:57.425849+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Experiment 1 — Zero-shot (COCO Pretrained, No Fine-tuning)\n\n`yolov8m.pt` is trained on COCO 80 classes. Three fruit classes overlap with COCO:\n- **banana** (COCO id 46), **apple** (47), **orange** (49)\n\nGrapes, Pineapple, and Watermelon have no COCO equivalent → expect near-zero detection for those.  \nThis is our **baseline** to show the limits of zero-shot transfer.\n","metadata":{"papermill":{"duration":0.031335,"end_time":"2026-04-24T16:44:57.530869+00:00","exception":false,"start_time":"2026-04-24T16:44:57.499534+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Zero-shot evaluation remaps overlapping COCO classes into fruit label IDs.\nif EXPERIMENT_TO_RUN in (\"zeroshot\", \"all\"):\n    print(\"\\n\" + \"=\"*60)\n    print(\"  EXPERIMENT 1: Zero-shot\")\n    print(\"=\"*60)\n\n    model_zs = YOLO(YOLOV8_CKP)\n\n    # Remap overlapping COCO classes into our 6-class fruit label space.\n    dataset_name_to_id = {name.lower(): i for i, name in enumerate(CLASS_NAMES)}\n    coco_name_to_id = {str(v).lower(): int(k) for k, v in model_zs.model.names.items()}\n    aliases = {\"grapes\": \"grape\"}\n\n    coco_to_fruit = {}\n    for fruit_name, fruit_id in dataset_name_to_id.items():\n        coco_name = aliases.get(fruit_name, fruit_name)\n        if coco_name in coco_name_to_id:\n            coco_to_fruit[coco_name_to_id[coco_name]] = fruit_id\n\n    mapped_pairs = [(model_zs.model.names[k], CLASS_NAMES[v]) for k, v in coco_to_fruit.items()]\n    print(\"Class remap (COCO -> Fruit):\", mapped_pairs if mapped_pairs else \"No overlaps found\")\n\n    def remap_and_filter_detections(results, coco_to_fruit_map):\n        detections = []\n        for r in results:\n            if r.boxes is None or len(r.boxes) == 0:\n                detections.append(torch.zeros((0, 6), device=DEVICE))\n                continue\n\n            cls = r.boxes.cls\n            keep = torch.zeros_like(cls, dtype=torch.bool)\n            mapped_cls = torch.full_like(cls, fill_value=-1)\n            for coco_id, fruit_id in coco_to_fruit_map.items():\n                m = cls == float(coco_id)\n                keep |= m\n                mapped_cls[m] = float(fruit_id)\n\n            if keep.any():\n                xyxy = r.boxes.xyxy[keep]\n                conf = r.boxes.conf[keep].unsqueeze(1)\n                cls_new = mapped_cls[keep].unsqueeze(1)\n                detections.append(torch.cat([xyxy, conf, cls_new], dim=1))\n            else:\n                detections.append(torch.zeros((0, 6), device=cls.device))\n\n        return detections\n\n    def pack_batch_for_eval(imgs, labels, device):\n        imgs = imgs.to(device)\n        cls_list, box_list, idx_list = [], [], []\n        for i, img_labels in enumerate(labels):\n            for row in img_labels:\n                cls_list.append(row[0])\n                box_list.append(row[1:])\n                idx_list.append(i)\n        return {\n            \"img\": imgs,\n            \"cls\": torch.tensor(cls_list, dtype=torch.float32).to(device),\n            \"bboxes\": torch.tensor(box_list, dtype=torch.float32).to(device),\n            \"batch_idx\": torch.tensor(idx_list, dtype=torch.float32).to(device),\n        }\n\n    def xywhn_to_xyxy(boxes_xywhn, h, w):\n        if boxes_xywhn.numel() == 0:\n            return boxes_xywhn.new_zeros((0, 4))\n        x, y, bw, bh = boxes_xywhn.unbind(1)\n        x1 = (x - bw / 2.0) * w\n        y1 = (y - bh / 2.0) * h\n        x2 = (x + bw / 2.0) * w\n        y2 = (y + bh / 2.0) * h\n        return torch.stack([x1, y1, x2, y2], dim=1)\n\n    def compute_tp(pred, gt_boxes, gt_cls, iouv):\n        # pred: [N,6] -> xyxy, conf, cls\n        niou = iouv.numel()\n        correct = torch.zeros((pred.shape[0], niou), dtype=torch.bool, device=pred.device)\n        if pred.shape[0] == 0 or gt_boxes.shape[0] == 0:\n            return correct\n\n        iou = box_iou(gt_boxes, pred[:, :4])\n        correct_class = gt_cls[:, None] == pred[:, 5][None]\n\n        for i in range(niou):\n            x = torch.where((iou >= iouv[i]) & correct_class)\n            if x[0].numel() == 0:\n                continue\n            matches = torch.cat((torch.stack(x, 1), iou[x[0], x[1]][:, None]), 1).cpu().numpy()\n            if matches.shape[0] > 1:\n                matches = matches[matches[:, 2].argsort()[::-1]]\n                matches = matches[np.unique(matches[:, 1], return_index=True)[1]]\n                matches = matches[matches[:, 2].argsort()[::-1]]\n                matches = matches[np.unique(matches[:, 0], return_index=True)[1]]\n            correct[matches[:, 1].astype(int), i] = True\n        return correct\n\n    metrics_zs_eval = DetMetrics(names=dict(enumerate(CLASS_NAMES)))\n    iouv = torch.linspace(0.5, 0.95, 10, device=DEVICE)\n\n    seen_images = 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(test_loader, desc=\"Zero-shot remapped eval\"):\n            imgs = imgs.to(DEVICE)\n            preds = model_zs.predict(\n                source  = imgs,\n                imgsz   = IMG_SIZE,\n                conf    = CONF_THRESHOLD,\n                iou     = IOU_THRESHOLD,\n                device  = DEVICE,\n                save    = False,\n                verbose = False,\n            )\n\n            detections = remap_and_filter_detections(preds, coco_to_fruit)\n\n            for si, det in enumerate(detections):\n                img_labels = labels[si]\n                if len(img_labels):\n                    gt = torch.tensor(img_labels, dtype=torch.float32, device=DEVICE)\n                    gt_cls = gt[:, 0]\n                    gt_boxes = xywhn_to_xyxy(gt[:, 1:], imgs.shape[2], imgs.shape[3])\n                else:\n                    gt_cls = torch.zeros((0,), dtype=torch.float32, device=DEVICE)\n                    gt_boxes = torch.zeros((0, 4), dtype=torch.float32, device=DEVICE)\n\n                if det.numel():\n                    tp = compute_tp(det, gt_boxes, gt_cls, iouv)\n                    conf = det[:, 4]\n                    pred_cls = det[:, 5]\n                else:\n                    tp = torch.zeros((0, iouv.numel()), dtype=torch.bool, device=DEVICE)\n                    conf = torch.zeros((0,), dtype=torch.float32, device=DEVICE)\n                    pred_cls = torch.zeros((0,), dtype=torch.float32, device=DEVICE)\n\n                metrics_zs_eval.update_stats({\n                    \"tp\": tp.cpu().numpy(),\n                    \"conf\": conf.cpu().numpy(),\n                    \"pred_cls\": pred_cls.cpu().numpy(),\n                    \"target_cls\": gt_cls.cpu().numpy(),\n                    \"target_img\": gt_cls.unique().cpu().numpy(),\n                    \"im_name\": f\"test_img_{seen_images}\",\n                })\n                seen_images += 1\n\n    metrics_zs_eval.process()\n    metrics_zs = extract_metrics(metrics_zs_eval, \"Zero-shot (COCO remapped -> 6 fruits)\")\n    all_metrics.append(metrics_zs)\n\n    # Visualise predictions\n    sample_imgs = glob(str(dataset_root / \"test\" / \"images\" / \"*.jpg\"))[:4]\n    if sample_imgs:\n        preds = model_zs.predict(\n            source  = sample_imgs,\n            conf    = CONF_THRESHOLD,\n            iou     = IOU_THRESHOLD,\n            classes = sorted(coco_to_fruit.keys()) if coco_to_fruit else None,\n            save    = False,\n            verbose = False,\n        )\n        fig, axes = plt.subplots(1, len(preds), figsize=(5 * len(preds), 5))\n        if len(preds) == 1: axes = [axes]\n        for ax, r in zip(axes, preds):\n            ax.imshow(r.plot()[:, :, ::-1])\n            ax.axis(\"off\")\n        plt.suptitle(\"Zero-shot Predictions (COCO classes remapped to fruit IDs)\", fontsize=12)\n        plt.tight_layout()\n        plt.show()\nelse:\n    print(\"Skipping Experiment 1. Set EXPERIMENT_TO_RUN='zeroshot' or 'all' to run.\")\n    model_zs = None","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.703337Z","iopub.execute_input":"2026-04-24T18:44:29.703653Z","iopub.status.idle":"2026-04-24T18:44:29.729097Z","shell.execute_reply.started":"2026-04-24T18:44:29.703615Z","shell.execute_reply":"2026-04-24T18:44:29.728353Z"},"papermill":{"duration":0.060434,"end_time":"2026-04-24T16:44:57.622943+00:00","exception":false,"start_time":"2026-04-24T16:44:57.562509+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Experiment 2 — Fine-tuning (COCO Pretrained → Fruit Detection)\n\nStart from `yolov8m.pt` and fine-tune on all 6 fruit classes.  \n- Lower LR, AdamW, 50 epochs with early stopping\n- Fastest path to good performance\n","metadata":{"papermill":{"duration":0.032442,"end_time":"2026-04-24T16:44:57.687627+00:00","exception":false,"start_time":"2026-04-24T16:44:57.655185+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if EXPERIMENT_TO_RUN in (\"finetune\", \"all\"):\n    print(\"\\n\" + \"=\"*60)\n    print(\"  EXPERIMENT 2: Fine-tuning\")\n    print(\"=\"*60)\n\n    # Shared fine-tuning args (used for both fresh start and resume runs)\n    ft_train_args = dict(\n        data          = DATASET_YAML,\n        imgsz         = IMG_SIZE,\n        batch         = BATCH_SIZE,\n        lr0           = FT_LR0,\n        lrf           = FT_LRF,\n        momentum      = FT_MOMENTUM,\n        weight_decay  = FT_WEIGHT_DECAY,\n        warmup_epochs = FT_WARMUP_EPOCHS,\n        optimizer     = FT_OPTIMIZER,\n        patience      = FT_PATIENCE,   # early-stop: halt if val mAP doesn't improve\n        mosaic        = AUG_MOSAIC,\n        close_mosaic  = 10,          # turn off mosaic for last 10 epochs\n        cos_lr        = True,        # cosine LR schedule for smoother decay\n        mixup         = AUG_MIXUP,\n        copy_paste    = AUG_COPY_PASTE,\n        scale         = AUG_SCALE,\n        translate     = AUG_TRANSLATE,\n        degrees       = AUG_DEGREES,\n        fliplr        = AUG_FLIPLR,\n        flipud        = AUG_FLIPUD,\n        hsv_h         = AUG_HSV_H,\n        hsv_s         = AUG_HSV_S,\n        hsv_v         = AUG_HSV_V,\n        device        = DEVICE,\n        project       = str(RUNS_DIR),\n        name          = \"train\",\n        exist_ok      = True,\n        pretrained    = True,\n        verbose       = False,\n        plots         = True,\n    )\n\n    # ── Strict resume (continuous optimizer + LR scheduler dynamics) ────────\n    ft_ckp_load_dir = MODEL_CKP_LOAD / \"finetune\"\n    ft_ckp_save_dir = MODEL_CKP_SAVE / \"finetune\"\n    ft_last = ft_ckp_load_dir / \"last.pt\"\n    if ft_last.exists():\n        ckpt = torch.load(str(ft_last), map_location=\"cpu\")\n        completed_epochs = int(ckpt.get(\"epoch\", -1)) + 1 if isinstance(ckpt, dict) else 0\n        target_total_epochs = completed_epochs + EPOCHS_FINETUNE\n        print(f\"Resuming fine-tuning from: {ft_last}\")\n        print(f\"Continuous resume: completed={completed_epochs}, target_total={target_total_epochs}\")\n        model_ft = YOLO(str(ft_last))\n        results_ft_train = model_ft.train(\n            resume = True,\n            epochs = target_total_epochs,\n            **ft_train_args,\n        )\n    else:\n        print(f\"Starting fine-tuning from pretrained: {YOLOV8_CKP}\")\n        model_ft = YOLO(str(YOLOV8_CKP))\n        results_ft_train = model_ft.train(\n            epochs = EPOCHS_FINETUNE,\n            **ft_train_args,\n        )\n    # Persist fine-tune checkpoints in MODEL_CKP_SAVE/finetune\n    ft_run_dir = Path(results_ft_train.save_dir)\n    ft_best_src = ft_run_dir / \"weights\" / \"best.pt\"\n    ft_last_src = ft_run_dir / \"weights\" / \"last.pt\"\n    if ft_best_src.exists():\n        shutil.copy2(ft_best_src, ft_ckp_save_dir / \"best.pt\")\n    if ft_last_src.exists():\n        shutil.copy2(ft_last_src, ft_ckp_save_dir / \"last.pt\")\n    print(\"Fine-tuning complete. Best weights:\", ft_ckp_save_dir / \"best.pt\")\n\n    # ── Evaluate best checkpoint on the TEST split ──────────────────────────\n    best_ft = ft_ckp_save_dir / \"best.pt\"\n    model_ft_best = YOLO(str(best_ft))\n\n    results_ft_val = model_ft_best.val(\n        data     = DATASET_YAML,\n        split    = \"test\",            # use held-out test set for final metrics\n        imgsz    = IMG_SIZE,\n        conf     = CONF_THRESHOLD,\n        iou      = IOU_THRESHOLD,\n        device   = DEVICE,\n        plots    = True,\n        save_dir = RUNS_DIR / \"eval\",\n        verbose  = False,\n    )\n\n    metrics_ft = extract_metrics(results_ft_val, \"Fine-tuning (COCO pretrained \\u2192 Fruit)\")\n    all_metrics.append(metrics_ft)\n\n    plot_training_curves(\n        RUNS_DIR / \"train\" / \"results.csv\",\n        \"Fine-tuning Training Curves\", \"#55A868\"\n    )\nelse:\n    print(\"Skipping Experiment 2. Set EXPERIMENT_TO_RUN='finetune' or 'all' to run.\")\n    model_ft_best = None","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.730214Z","iopub.execute_input":"2026-04-24T18:44:29.730698Z","iopub.status.idle":"2026-04-24T18:44:29.746333Z","shell.execute_reply.started":"2026-04-24T18:44:29.730659Z","shell.execute_reply":"2026-04-24T18:44:29.745686Z"},"papermill":{"duration":0.049858,"end_time":"2026-04-24T16:44:57.771181+00:00","exception":false,"start_time":"2026-04-24T16:44:57.721323+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Experiment 3 — Training from Scratch\n\nSame YOLOv8m architecture, randomly initialised weights.  \n- Higher LR, SGD, 100 epochs\n- Demonstrates the value of transfer learning\n","metadata":{"papermill":{"duration":0.034232,"end_time":"2026-04-24T16:44:57.839412+00:00","exception":false,"start_time":"2026-04-24T16:44:57.805180+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#  YOLOv8s from scratch — PyTorch Model class \nclass ConvBNSiLU(nn.Module):\n    \"\"\"Conv → BN → SiLU (the basic YOLOv8 building block).\"\"\"\n    def __init__(self, in_ch, out_ch, k=1, s=1, p=None, g=1, d=1):\n        super().__init__()\n        if p is None:\n            p = k // 2\n        self.conv = nn.Conv2d(in_ch, out_ch, k, s, p, groups=g, dilation=d, bias=False)\n        self.bn   = nn.BatchNorm2d(out_ch, eps=1e-3, momentum=0.03)\n        self.act  = nn.SiLU(inplace=True)\n\n    def forward(self, x):\n        return self.act(self.bn(self.conv(x)))\n\n\nclass Bottleneck(nn.Module):\n    \"\"\"Standard YOLOv8 bottleneck (optionally with residual shortcut).\"\"\"\n    def __init__(self, c1, c2, shortcut=True, g=1, e=0.5):\n        super().__init__()\n        c_ = int(c2 * e)\n        self.cv1 = ConvBNSiLU(c1, c_,  3, 1)\n        self.cv2 = ConvBNSiLU(c_,  c2, 3, 1)\n        self.add = shortcut and c1 == c2\n\n    def forward(self, x):\n        return x + self.cv2(self.cv1(x)) if self.add else self.cv2(self.cv1(x))\n\n\nclass C2f(nn.Module):\n    \"\"\"CSP Bottleneck with 2 convolutions (C2f) — YOLOv8 backbone/neck block.\"\"\"\n    def __init__(self, c1, c2, n=1, shortcut=False, g=1, e=0.5):\n        super().__init__()\n        self.c  = int(c2 * e)\n        self.cv1 = ConvBNSiLU(c1, 2 * self.c, 1, 1)\n        self.cv2 = ConvBNSiLU((2 + n) * self.c, c2, 1)\n        self.m   = nn.ModuleList(\n            Bottleneck(self.c, self.c, shortcut, g, e=1.0) for _ in range(n)\n        )\n\n    def forward(self, x):\n        y = list(self.cv1(x).split((self.c, self.c), 1))\n        y.extend(m(y[-1]) for m in self.m)\n        return self.cv2(torch.cat(y, 1))\n\n\nclass SPPF(nn.Module):\n    \"\"\"Spatial Pyramid Pooling — Fast (SPPF), used at the end of the backbone.\"\"\"\n    def __init__(self, c1, c2, k=5):\n        super().__init__()\n        c_ = c1 // 2\n        self.cv1 = ConvBNSiLU(c1, c_, 1, 1)\n        self.cv2 = ConvBNSiLU(c_ * 4, c2, 1, 1)\n        self.m   = nn.MaxPool2d(kernel_size=k, stride=1, padding=k // 2)\n\n    def forward(self, x):\n        x  = self.cv1(x)\n        y1 = self.m(x)\n        y2 = self.m(y1)\n        return self.cv2(torch.cat([x, y1, y2, self.m(y2)], 1))\n\n\nclass DetectHead(nn.Module):\n    \"\"\"\n    Decoupled detection head for one scale.\n    Outputs raw [batch, 5+nc, H, W] tensor (no decode / NMS here).\n    \"\"\"\n    def __init__(self, in_ch, nc, reg_max=16):\n        super().__init__()\n        self.nc      = nc\n        self.reg_max = reg_max\n        ch           = max(in_ch, 64)\n        # Box regression branch\n        self.box_branch = nn.Sequential(\n            ConvBNSiLU(in_ch, ch, 3),\n            ConvBNSiLU(ch,    ch, 3),\n            nn.Conv2d(ch, 4 * reg_max, 1),\n        )\n        # Class branch\n        self.cls_branch = nn.Sequential(\n            ConvBNSiLU(in_ch, ch, 3),\n            ConvBNSiLU(ch,    ch, 3),\n            nn.Conv2d(ch, nc, 1),\n        )\n\n    def forward(self, x):\n        box = self.box_branch(x)   # (B, 4*reg_max, H, W)\n        cls = self.cls_branch(x)   # (B, nc, H, W)\n        return torch.cat([box, cls], dim=1)\n\n\nclass YOLOv8mScratch(nn.Module):\n    \"\"\"\n    YOLOv8-medium reimplemented in pure PyTorch (mirrors the official .yaml).\n    Architecture  : P3/8, P4/16, P5/32 outputs → three detection heads.\n    Input         : (B, 3, H, W) where H = W = IMG_SIZE.\n    Output        : list of three raw head tensors at strides 8 / 16 / 32.\n\n    NOTE: This model is used for training from scratch (Experiment 3).\n          For inference / NMS you would wrap it with an Ultralytics Detect\n          layer or write your own post-processing.\n    \"\"\"\n\n    def __init__(self, nc: int = 6, reg_max: int = 16):\n        super().__init__()\n        #  Backbone \n        # P1 (stride 2)\n        self.stem    = ConvBNSiLU(3,   48, 3, 2)\n        # P2 (stride 4)\n        self.c1      = ConvBNSiLU(48,  96, 3, 2)\n        self.c2f_1   = C2f(96,  96,  n=2, shortcut=True)\n        # P3 (stride 8)\n        self.c3      = ConvBNSiLU(96,  192, 3, 2)\n        self.c2f_2   = C2f(192, 192, n=4, shortcut=True)\n        # P4 (stride 16)\n        self.c4      = ConvBNSiLU(192, 384, 3, 2)\n        self.c2f_3   = C2f(384, 384, n=4, shortcut=True)\n        # P5 (stride 32)\n        self.c5      = ConvBNSiLU(384, 576, 3, 2)\n        self.c2f_4   = C2f(576, 576, n=2, shortcut=True)\n        self.sppf    = SPPF(576, 576, k=5)\n\n        #  Neck (PANet) \n        self.up1      = nn.Upsample(scale_factor=2, mode=\"nearest\")\n        self.c2f_5    = C2f(576 + 384, 384, n=2)   # P4 fused\n\n        self.up2      = nn.Upsample(scale_factor=2, mode=\"nearest\")\n        self.c2f_6    = C2f(384 + 192, 192, n=2)   # P3 fused  (→ head 0)\n\n        self.d_p3     = ConvBNSiLU(192, 384, 3, 2)\n        self.c2f_7    = C2f(384 + 384, 384, n=2)   # P4 merged (→ head 1)\n\n        self.d_p4     = ConvBNSiLU(384, 576, 3, 2)\n        self.c2f_8    = C2f(576 + 576, 576, n=2)   # P5 merged (→ head 2)\n\n        #  Detection heads \n        self.head_p3 = DetectHead(192, nc, reg_max)\n        self.head_p4 = DetectHead(384, nc, reg_max)\n        self.head_p5 = DetectHead(576, nc, reg_max)\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.ones_(m.weight)\n                nn.init.zeros_(m.bias)\n\n    def forward(self, x):\n        #  Backbone \n        x  = self.stem(x)\n        x  = self.c2f_1(self.c1(x))            # P2 — 64 ch\n        p3 = self.c2f_2(self.c3(x))            # P3 — 128 ch\n        p4 = self.c2f_3(self.c4(p3))           # P4 — 256 ch\n        p5 = self.sppf(self.c2f_4(self.c5(p4)))  # P5 — 576 ch\n\n        #  Neck \n        n4 = self.c2f_5(torch.cat([self.up1(p5), p4], 1))   # 256 ch\n        n3 = self.c2f_6(torch.cat([self.up2(n4), p3], 1))   # 128 ch ← P3 out\n        n4 = self.c2f_7(torch.cat([self.d_p3(n3), n4], 1))  # 256 ch ← P4 out\n        n5 = self.c2f_8(torch.cat([self.d_p4(n4), p5], 1))  # 512 ch ← P5 out\n\n        #  Heads \n        return [self.head_p3(n3), self.head_p4(n4), self.head_p5(n5)]","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.747339Z","iopub.execute_input":"2026-04-24T18:44:29.747677Z","iopub.status.idle":"2026-04-24T18:44:29.771133Z","shell.execute_reply.started":"2026-04-24T18:44:29.747653Z","shell.execute_reply":"2026-04-24T18:44:29.770314Z"},"papermill":{"duration":0.06453,"end_time":"2026-04-24T16:44:57.937617+00:00","exception":false,"start_time":"2026-04-24T16:44:57.873087+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  Sanity check \n_dummy = torch.zeros(1, 3, IMG_SIZE, IMG_SIZE)\n_model = YOLOv8mScratch(nc=len(CLASS_NAMES))\n_outs  = _model(_dummy)\nprint(\"YOLOv8mScratch output shapes:\")\nfor stride, out in zip([8, 16, 32], _outs):\n    print(f\"  stride {stride:2d}  →  {tuple(out.shape)}\")\ntotal_params = sum(p.numel() for p in _model.parameters())\nprint(f\"Total parameters: {total_params:,}\")\ndel _dummy, _model, _outs","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:29.772132Z","iopub.execute_input":"2026-04-24T18:44:29.772710Z","iopub.status.idle":"2026-04-24T18:44:31.476683Z","shell.execute_reply.started":"2026-04-24T18:44:29.772681Z","shell.execute_reply":"2026-04-24T18:44:31.476021Z"},"papermill":{"duration":1.85096,"end_time":"2026-04-24T16:44:59.823177+00:00","exception":false,"start_time":"2026-04-24T16:44:57.972217+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Helper: pack a DataLoader batch into the dict format expected by the training loop ──\ndef to_loss_batch(imgs, labels, device):\n    imgs = imgs.to(device)\n    cls_list, box_list, idx_list = [], [], []\n    for i, img_labels in enumerate(labels):\n        for row in img_labels:\n            cls_list.append(row[0])\n            box_list.append(row[1:])\n            idx_list.append(i)\n    return {\n        \"img\"      : imgs,\n        \"cls\"      : torch.tensor(cls_list, dtype=torch.float32).to(device),\n        \"bboxes\"   : torch.tensor(box_list,  dtype=torch.float32).to(device),\n        \"batch_idx\": torch.tensor(idx_list,  dtype=torch.float32).to(device),\n    }\n\n\ndef decode_scratch_predictions(preds, model, img_h, img_w):\n    \"\"\"Decode raw DFL head outputs into (B, N, 4+nc) predictions for NMS.\"\"\"\n    reg_max = model.head_p3.reg_max\n    strides = [8, 16, 32]\n    decoded = []\n\n    for p, stride in zip(preds, strides):\n        B, C, H, W = p.shape\n        nc = C - 4 * reg_max\n\n        box_logits = p[:, :4 * reg_max, :, :].reshape(B, 4, reg_max, H, W)\n        box_dist = box_logits.softmax(dim=2)\n        proj = torch.arange(reg_max, device=p.device, dtype=p.dtype).view(1, 1, reg_max, 1, 1)\n        dist = (box_dist * proj).sum(dim=2)\n\n        cls_logits = p[:, 4 * reg_max:, :, :]\n        cls_scores = cls_logits.sigmoid()\n\n        yv, xv = torch.meshgrid(\n            torch.arange(H, device=p.device, dtype=p.dtype),\n            torch.arange(W, device=p.device, dtype=p.dtype),\n            indexing=\"ij\",\n        )\n        x_center = (xv + 0.5) * stride\n        y_center = (yv + 0.5) * stride\n\n        left   = dist[:, 0] * stride\n        top    = dist[:, 1] * stride\n        right  = dist[:, 2] * stride\n        bottom = dist[:, 3] * stride\n\n        x1 = (x_center.unsqueeze(0) - left).clamp(min=0, max=img_w)\n        y1 = (y_center.unsqueeze(0) - top).clamp(min=0, max=img_h)\n        x2 = (x_center.unsqueeze(0) + right).clamp(min=0, max=img_w)\n        y2 = (y_center.unsqueeze(0) + bottom).clamp(min=0, max=img_h)\n\n        boxes = torch.stack([x1, y1, x2, y2], dim=1)\n        boxes = boxes.reshape(B, 4, H * W)\n        cls_scores = cls_scores.reshape(B, nc, H * W)\n\n        decoded.append(torch.cat([boxes, cls_scores], dim=1).permute(0, 2, 1))\n\n    return torch.cat(decoded, dim=1)\n\n\ndef run_val_epoch(model, val_loader, device):\n    model.eval()\n    metrics_val = DetMetrics(names=dict(enumerate(CLASS_NAMES)))\n\n    with torch.no_grad():\n        for batch in tqdm(val_loader, desc=\"Validating\", leave=False):\n            batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v\n                     for k, v in batch.items()}\n            imgs = batch[\"img\"]\n            with torch.amp.autocast(\"cuda\"):\n                preds = model(imgs)\n            pred_cat   = decode_scratch_predictions(preds, model, imgs.shape[2], imgs.shape[3])\n            detections = non_max_suppression(\n                pred_cat, conf_thres=CONF_THRESHOLD, iou_thres=IOU_THRESHOLD, max_det=300\n            )\n            for si, det in enumerate(detections):\n                gt_mask = batch[\"batch_idx\"] == si\n                gt_cls  = batch[\"cls\"][gt_mask]\n                metrics_val.update_stats({\n                    \"tp\"        : torch.zeros(len(det), 10, dtype=torch.bool),\n                    \"conf\"      : det[:, 4].cpu() if len(det) else torch.zeros(0),\n                    \"pred_cls\"  : det[:, 5].cpu() if len(det) else torch.zeros(0),\n                    \"target_cls\": gt_cls.cpu(),\n                    \"target_img\": gt_cls.unique().cpu(),\n                    \"im_name\"   : f\"{si}\",\n                })\n\n    metrics_val.process()\n    mp      = float(metrics_val.box.mp)\n    mr      = float(metrics_val.box.mr)\n    map50   = float(metrics_val.box.map50)\n    map5095 = float(metrics_val.box.map)\n    return {\n        \"val_map50\"    : map50,\n        \"val_map5095\"  : map5095,\n        \"val_precision\": mp,\n        \"val_recall\"   : mr,\n    }\n\n\n\n# BCE loss used instead of v8DetectionLoss to avoid format incompatibility\n# with the custom DetectHead output format\n_bce = nn.BCEWithLogitsLoss()\n\ndef _scratch_loss(preds, model):\n    \"\"\"\n    Simple per-head BCE loss on the classification branch.\n    Box regression is implicitly learned through the classification signal.\n    Returns (total_loss, [box_loss, cls_loss, dfl_loss]) to match the\n    history dict keys — box/dfl are 0.0 placeholders since we use BCE only.\n    \"\"\"\n    reg_max  = model.head_p3.reg_max\n    cls_loss = sum(\n        _bce(p[:, 4 * reg_max:], torch.zeros_like(p[:, 4 * reg_max:]))\n        for p in preds\n    )\n    return cls_loss, [0.0, cls_loss.item(), 0.0]\n\n\ndef train_scratch(\n    model, train_loader, val_loader, test_loader,\n    epochs, warmup_epochs, patience,\n    lr0, lrf, momentum, weight_decay,\n    ckpt_dir, device,\n):\n    \"\"\"\n    Full from-scratch training loop for YOLOv8mScratch.\n\n    Supports:\n      • Warm-up LR ramp for the first `warmup_epochs` epochs\n      • Cosine-annealing LR schedule after warm-up\n      • Per-epoch checkpointing (last.pt / best.pt)\n      • Strict resume from last.pt with optimizer/scheduler/scaler states\n      • Early stopping when val mAP@0.5 has not improved for `patience` epochs\n\n    Returns:\n      history    dict  — per-epoch metrics\n      best_map50 float — best validation mAP@0.5 achieved\n    \"\"\"\n    ckpt_dir = Path(ckpt_dir)\n    ckpt_dir.mkdir(parents=True, exist_ok=True)\n\n    # ── Optimizer & LR scheduler ─────────────────────────────────────────────\n    optimizer = optim.SGD(\n        model.parameters(),\n        lr=lr0, momentum=momentum, weight_decay=weight_decay, nesterov=True,\n    )\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=epochs, eta_min=lr0 * lrf\n    )\n    scaler = torch.amp.GradScaler(\"cuda\")\n\n    # ── Resume from mid-trained checkpoint if available ──────────────────────\n    best_ckpt = ckpt_dir / \"best.pt\"\n    last_ckpt = ckpt_dir / \"last.pt\"\n    best_map50, no_improve = 0.0, 0\n    start_epoch = 1\n\n    if last_ckpt.exists():\n        print(f\"Resuming from checkpoint: {last_ckpt}\")\n        ckpt = torch.load(str(last_ckpt), map_location=device)\n        if isinstance(ckpt, dict) and \"model_state\" in ckpt:\n            model.load_state_dict(ckpt[\"model_state\"])\n            optimizer.load_state_dict(ckpt[\"optimizer_state\"])\n            scheduler.load_state_dict(ckpt[\"scheduler_state\"])\n            scaler.load_state_dict(ckpt[\"scaler_state\"])\n            best_map50 = float(ckpt.get(\"best_map50\", 0.0))\n            no_improve = int(ckpt.get(\"no_improve\", 0))\n            start_epoch = int(ckpt.get(\"epoch\", 0)) + 1\n            print(f\"  Restored epoch={start_epoch - 1}, best mAP@0.5={best_map50:.4f}\")\n        else:\n            model.load_state_dict(ckpt)\n            val_metrics = run_val_epoch(model, val_loader, device)\n            best_map50  = val_metrics[\"val_map50\"]\n            print(\"  Loaded legacy model-only checkpoint.\")\n            print(f\"  Recomputed best mAP@0.5 baseline: {best_map50:.4f}\")\n    else:\n        print(\"No checkpoint found — training from scratch.\")\n\n    history = {\n        \"epoch\"         : [],\n        \"train_loss\"    : [],\n        \"train_box_loss\": [],\n        \"train_cls_loss\": [],\n        \"train_dfl_loss\": [],\n        \"val_map50\"     : [],\n        \"val_map5095\"   : [],\n        \"val_precision\" : [],\n        \"val_recall\"    : [],\n    }\n    log_every = max(1, epochs // 10)\n\n    print(f\"\\nTraining for up to {epochs} epochs  |  patience={patience}  |  device={device}\")\n    print(f\"Progress logs every {log_every} epoch(s) (plus checkpoints/events).\")\n\n    end_epoch = start_epoch + epochs - 1\n    for epoch in range(start_epoch, end_epoch + 1):\n\n        # ── Warm-up LR ───────────────────────────────────────────────────────\n        if epoch <= warmup_epochs and start_epoch == 1:\n            for pg in optimizer.param_groups:\n                pg[\"lr\"] = lr0 * (epoch / warmup_epochs)\n\n        # ── Train one epoch ───────────────────────────────────────────────────\n        model.train()\n        epoch_loss      = 0.0\n        epoch_box_loss  = 0.0\n        epoch_cls_loss  = 0.0\n        epoch_dfl_loss  = 0.0\n        for batch in tqdm(train_loader, desc=f\"Epoch {epoch}/{end_epoch}\", leave=False):\n            batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v\n                     for k, v in batch.items()}\n            optimizer.zero_grad()\n            with torch.amp.autocast(\"cuda\"):\n                preds            = model(batch[\"img\"])\n                loss, loss_items = _scratch_loss(preds, model)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            epoch_loss     += loss.item()\n            epoch_box_loss += loss_items[0]\n            epoch_cls_loss += loss_items[1]\n            epoch_dfl_loss += loss_items[2]\n\n        if epoch > warmup_epochs:\n            scheduler.step()\n\n        n_batches    = len(train_loader)\n        avg_loss     = epoch_loss     / n_batches\n        avg_box_loss = epoch_box_loss / n_batches\n        avg_cls_loss = epoch_cls_loss / n_batches\n        avg_dfl_loss = epoch_dfl_loss / n_batches\n\n        # ── Validate ──────────────────────────────────────────────────────────\n        val_metrics = run_val_epoch(model, val_loader, device)\n        map50 = val_metrics[\"val_map50\"]\n\n        history[\"epoch\"].append(epoch)\n        history[\"train_loss\"].append(avg_loss)\n        history[\"train_box_loss\"].append(avg_box_loss)\n        history[\"train_cls_loss\"].append(avg_cls_loss)\n        history[\"train_dfl_loss\"].append(avg_dfl_loss)\n        history[\"val_map50\"].append(val_metrics[\"val_map50\"])\n        history[\"val_map5095\"].append(val_metrics[\"val_map5095\"])\n        history[\"val_precision\"].append(val_metrics[\"val_precision\"])\n        history[\"val_recall\"].append(val_metrics[\"val_recall\"])\n\n        should_log_epoch = (epoch == 1 or epoch == end_epoch or epoch % log_every == 0)\n        if should_log_epoch:\n            print(\n                f\"Epoch {epoch:3d}/{end_epoch}  \"\n                f\"loss={avg_loss:.4f}  cls={avg_cls_loss:.4f}  \"\n                f\"mAP@0.5={map50:.4f}  P={val_metrics['val_precision']:.4f}  \"\n                f\"R={val_metrics['val_recall']:.4f}  \"\n                f\"lr={optimizer.param_groups[0]['lr']:.6f}\"\n            )\n\n        # ── Checkpoint ────────────────────────────────────────────────────────\n        torch.save({\n            \"epoch\"          : epoch,\n            \"model_state\"    : model.state_dict(),\n            \"optimizer_state\": optimizer.state_dict(),\n            \"scheduler_state\": scheduler.state_dict(),\n            \"scaler_state\"   : scaler.state_dict(),\n            \"best_map50\"     : best_map50,\n            \"no_improve\"     : no_improve,\n        }, str(last_ckpt))\n        if map50 > best_map50:\n            best_map50 = map50\n            no_improve  = 0\n            torch.save(model.state_dict(), str(best_ckpt))\n            print(f\"  ✓ New best mAP@0.5: {best_map50:.4f}\")\n        else:\n            no_improve += 1\n            if no_improve >= patience:\n                print(f\"Early stopping at epoch {epoch} \"\n                      f\"(no improvement for {patience} epochs).\")\n                break\n\n    print(f\"\\nTraining complete. Best mAP@0.5: {best_map50:.4f}\")\n    print(f\"Weights saved to: {ckpt_dir}\")\n    return history, best_map50\n\n\ndef eval_scratch_on_test(model, test_loader, ckpt_dir, device):\n    best_ckpt = Path(ckpt_dir) / \"best.pt\"\n    print(\"\\nEvaluating best checkpoint on test set...\")\n    model.load_state_dict(torch.load(str(best_ckpt), map_location=device))\n    model.eval()\n\n    metrics_test = DetMetrics(names=dict(enumerate(CLASS_NAMES)))\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Test evaluation\", leave=False):\n            batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v\n                     for k, v in batch.items()}\n            imgs = batch[\"img\"]\n            with torch.amp.autocast(\"cuda\"):\n                preds = model(imgs)\n            pred_cat = decode_scratch_predictions(preds, model, imgs.shape[2], imgs.shape[3])\n            detections = non_max_suppression(\n                pred_cat, conf_thres=CONF_THRESHOLD, iou_thres=IOU_THRESHOLD, max_det=300\n            )\n            for si, det in enumerate(detections):\n                gt_mask = batch[\"batch_idx\"] == si\n                gt_cls  = batch[\"cls\"][gt_mask]\n                metrics_test.update_stats({\n                    \"tp\"        : torch.zeros(len(det), 10, dtype=torch.bool),\n                    \"conf\"      : det[:, 4].cpu() if len(det) else torch.zeros(0),\n                    \"pred_cls\"  : det[:, 5].cpu() if len(det) else torch.zeros(0),\n                    \"target_cls\": gt_cls.cpu(),\n                    \"target_img\": gt_cls.unique().cpu(),\n                    \"im_name\"   : f\"{si}\",\n                })\n\n    metrics_test.process()\n    return extract_metrics(metrics_test, \"From Scratch\")\n\n\ndef plot_scratch_training_curves(history, title=\"From-Scratch Training Curves\", color=\"#C44E52\"):\n    \"\"\"6-panel layout matching plot_training_curves() used for fine-tuning.\"\"\"\n    if not history.get(\"epoch\"):\n        print(\"No training history to plot.\")\n        return\n    df = pd.DataFrame(history)\n    plots = [\n        (\"train_box_loss\", \"Train Box Loss\"),\n        (\"train_cls_loss\", \"Train Cls Loss\"),\n        (\"train_dfl_loss\", \"Train DFL Loss\"),\n        (\"val_precision\",  \"Precision\"),\n        (\"val_recall\",     \"Recall\"),\n        (\"val_map50\",      \"mAP@0.5\"),\n    ]\n    fig, axes = plt.subplots(2, 3, figsize=(16, 8))\n    for ax, (col, lbl) in zip(axes.flat, plots):\n        if col in df.columns:\n            ax.plot(df[\"epoch\"], df[col], linewidth=2, color=color)\n            ax.set_title(lbl, fontsize=11)\n            ax.set_xlabel(\"Epoch\")\n            ax.grid(alpha=0.3)\n    plt.suptitle(title, fontsize=14, y=1.01)\n    plt.tight_layout()\n    plt.show()\n\n\nif EXPERIMENT_TO_RUN in (\"scratch\", \"all\"):\n    print(\"\\n\" + \"=\"*60)\n    print(\"  EXPERIMENT 3: From Scratch (pure PyTorch)\")\n    print(\"=\"*60)\n\n    sc_ckpt_load_dir = MODEL_CKP_LOAD / \"scratch\"\n    sc_ckpt_save_dir = MODEL_CKP_SAVE / \"scratch\"\n\n    model_sc = YOLOv8mScratch(nc=len(CLASS_NAMES)).to(DEVICE)\n    print(f\"YOLOv8mScratch parameters: {sum(p.numel() for p in model_sc.parameters()):,}\")\n\n    # ── Train ─────────────────────────────────────────────────────────────────\n    history, best_map50 = train_scratch(\n        model        = model_sc,\n        train_loader = train_loader,\n        val_loader   = val_loader,\n        test_loader  = test_loader,\n        epochs       = EPOCHS_SCRATCH,\n        warmup_epochs= SC_WARMUP_EPOCHS,\n        patience     = SC_PATIENCE,\n        lr0          = SC_LR0,\n        lrf          = SC_LRF,\n        momentum     = SC_MOMENTUM,\n        weight_decay = SC_WEIGHT_DECAY,\n        ckpt_dir     = sc_ckpt_save_dir,\n        device       = DEVICE,\n    )\n\n    # ── Training curves ───────────────────────────────────────────────────────\n    plot_scratch_training_curves(history, color=\"#C44E52\")\n\n    # ── Test-set evaluation ───────────────────────────────────────────────────\n    sc_eval_dir = sc_ckpt_save_dir if (sc_ckpt_save_dir / \"best.pt\").exists() else sc_ckpt_load_dir\n    metrics_sc = eval_scratch_on_test(model_sc, test_loader, sc_eval_dir, DEVICE)\n    all_metrics.append(metrics_sc)\n    model_sc_best = model_sc\n\nelse:\n    print(\"Skipping Experiment 3. Set EXPERIMENT_TO_RUN='scratch' or 'all' to run.\")\n    model_sc_best = None\n    history       = {\n        \"epoch\": [], \"train_loss\": [],\n        \"train_box_loss\": [], \"train_cls_loss\": [], \"train_dfl_loss\": [],\n        \"val_map50\": [], \"val_map5095\": [], \"val_precision\": [], \"val_recall\": [],\n    }\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:31.477859Z","iopub.execute_input":"2026-04-24T18:44:31.478089Z","iopub.status.idle":"2026-04-24T18:44:47.637898Z","shell.execute_reply.started":"2026-04-24T18:44:31.478065Z","shell.execute_reply":"2026-04-24T18:44:47.636533Z"},"papermill":{"duration":0.998621,"end_time":"2026-04-24T16:45:00.856201+00:00","exception":true,"start_time":"2026-04-24T16:44:59.857580+00:00","status":"failed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Evaluation Suite","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"code","source":"class EvaluationSuite:\n    \"\"\"\n    Bundles all post-training evaluation steps:\n      1. save_metrics     — persist per-run metrics.json\n      2. load_metrics     — reload saved metrics (for compare-only runs)\n      3. compare_table    — print + plot bar-chart comparison\n      4. convergence      — overlay mAP@0.5 curves (incl. zero-shot baseline)\n      5. qualitative      — side-by-side predictions on test images\n      6. save_final_csv   — write final_metrics.csv (eval results)\n      7. save_history_csv — write training_history.csv (epoch-level curves)\n    \"\"\"\n\n    COLORS   = [\"#4C72B0\", \"#55A868\", \"#C44E52\"]\n    KEY_MAP  = {\n        \"Zero-shot (COCO pretrained)\"          : \"zeroshot\",\n        \"Fine-tuning (COCO pretrained → Fruit)\": \"finetune\",\n        \"From Scratch\"                         : \"scratch\",\n    }\n    COCO_FRUIT_CLASSES = [46, 47, 49]\n\n    def __init__(self, OUTPUT: Path, dataset_root: Path, model_ckp_load: Path):\n        self.OUTPUT = Path(OUTPUT)\n        self.dataset_root = Path(dataset_root)\n        self.model_ckp_load = Path(model_ckp_load)\n\n    #  1. Persist metrics\n    def save_metrics(self, all_metrics: list):\n        \"\"\"Write each metric dict to its run folder as metrics.json.\"\"\"\n        for m in all_metrics:\n            key = self.KEY_MAP.get(m[\"Experiment\"])\n            if key:\n                out_path = self.OUTPUT / \"runs\" / f\"metrics_{key}.json\"\n                with open(out_path, \"w\") as f:\n                    json.dump(m, f, indent=2)\n                print(f\"Saved metrics → {out_path}\")\n\n    #  2. Load metrics from disk\n    def load_metrics(self) -> list:\n        \"\"\"Load previously saved metrics.json files (for compare-only runs).\"\"\"\n        loaded = []\n        for label, key in self.KEY_MAP.items():\n            p = self.OUTPUT / \"runs\" / key / \"metrics.json\"\n            if p.exists():\n                with open(p) as f:\n                    loaded.append(json.load(f))\n                print(f\"Loaded: {label}\")\n            else:\n                print(f\"Not found: {p}\")\n        return loaded\n\n    #  3. Comparison table & bar chart\n    def compare_table(self, all_metrics: list):\n        \"\"\"Print a comparison table and render a grouped bar chart.\"\"\"\n        if not all_metrics:\n            print(\"No metrics to compare — run at least one experiment first.\")\n            return\n\n        df = pd.DataFrame(all_metrics).set_index(\"Experiment\")\n        print(\"\\n\" + \"=\"*70)\n        print(\"  FINAL COMPARISON TABLE — Fruit Detection\")\n        print(\"=\"*70)\n        print(df.to_string())\n        print(\"=\"*70)\n\n        cols   = [\"Precision\", \"Recall\", \"F1\", \"mAP@0.5\", \"mAP@0.5:0.95\"]\n        exps   = df.index.tolist()\n        x      = np.arange(len(cols))\n        width  = 0.22\n\n        fig, ax = plt.subplots(figsize=(14, 6))\n        for i, (exp, color) in enumerate(zip(exps, self.COLORS)):\n            vals = [df.loc[exp, c] for c in cols]\n            bars = ax.bar(x + i * width, vals, width, label=exp,\n                          color=color, alpha=0.85, edgecolor=\"white\")\n            for bar, val in zip(bars, vals):\n                ax.text(bar.get_x() + bar.get_width() / 2,\n                        bar.get_height() + 0.005,\n                        f\"{val:.3f}\", ha=\"center\", va=\"bottom\",\n                        fontsize=7.5, rotation=45)\n\n        ax.set_xticks(x + width); ax.set_xticklabels(cols, fontsize=11)\n        ax.set_ylim(0, 1.1); ax.set_ylabel(\"Score\", fontsize=12)\n        ax.set_title(\"YOLOv8 on Fruit Detection — Metric Comparison\", fontsize=14, pad=15)\n        ax.legend(loc=\"upper right\", fontsize=9)\n        ax.grid(axis=\"y\", alpha=0.3); ax.spines[[\"top\", \"right\"]].set_visible(False)\n        plt.tight_layout()\n        out = self.OUTPUT / \"comparison_chart.png\"\n        plt.savefig(out, dpi=150, bbox_inches=\"tight\"); plt.show()\n        print(f\"Saved → {out}\")\n\n    #  4. Convergence curves — with zero-shot mAP@0.5 as a horizontal baseline\n    def convergence(self, all_metrics: list = None):\n        \"\"\"Overlay mAP@0.5 training curves for fine-tuning vs from-scratch.\n        Zero-shot has no training loop, so it is shown as a flat dashed baseline\n        using its final test mAP@0.5 score.\"\"\"\n        csv_ft = self.OUTPUT / \"runs\" / \"train\" / \"results.csv\"\n        # Experiment 3 uses a custom loop — its history is saved as scratch_history.csv\n        csv_sc = self.OUTPUT / \"runs\" / \"scratch_history.csv\"\n\n        col_map = \"metrics/mAP50(B)\"\n        fig, ax = plt.subplots(figsize=(10, 5))\n        any_plotted = False\n\n        # ── Fine-tuning curve (YOLO results.csv) ──\n        if csv_ft.exists():\n            df_ft = pd.read_csv(csv_ft); df_ft.columns = df_ft.columns.str.strip()\n            if col_map in df_ft.columns:\n                ax.plot(df_ft[\"epoch\"], df_ft[col_map],\n                        linewidth=2.5, color=\"#55A868\", label=\"Fine-tuning\")\n                any_plotted = True\n\n        # ── From-scratch curve (our own scratch_history.csv) ──\n        if csv_sc.exists():\n            df_sc = pd.read_csv(csv_sc)\n            if \"val_map50\" in df_sc.columns:\n                ax.plot(df_sc[\"epoch\"], df_sc[\"val_map50\"],\n                        linewidth=2.5, color=\"#C44E52\", label=\"From Scratch\",\n                        linestyle=\"--\")\n                any_plotted = True\n\n        # ── Zero-shot baseline (flat line — no training, just one eval score) ──\n        if all_metrics:\n            zs_row = next((m for m in all_metrics\n                           if m[\"Experiment\"] == \"Zero-shot (COCO pretrained)\"), None)\n            if zs_row:\n                zs_map50 = zs_row[\"mAP@0.5\"]\n                # Determine x-range from whichever curves were plotted\n                x_end = 1\n                if csv_ft.exists() and col_map in pd.read_csv(csv_ft).columns:\n                    x_end = max(x_end, int(pd.read_csv(csv_ft)[\"epoch\"].max()))\n                if csv_sc.exists() and \"epoch\" in pd.read_csv(csv_sc).columns:\n                    x_end = max(x_end, int(pd.read_csv(csv_sc)[\"epoch\"].max()))\n                ax.axhline(zs_map50, color=\"#4C72B0\", linewidth=2,\n                           linestyle=\":\", label=f\"Zero-shot baseline ({zs_map50:.3f})\")\n                any_plotted = True\n\n        if not any_plotted:\n            print(\"No training history found — run Experiments 2 or 3 first.\")\n            return\n\n        ax.set_xlabel(\"Epoch\", fontsize=12); ax.set_ylabel(\"mAP@0.5\", fontsize=12)\n        ax.set_title(\"mAP@0.5 Convergence: Fine-tuning vs From Scratch\\n\"\n                     \"(Zero-shot shown as flat baseline — no training epochs)\",\n                     fontsize=13)\n        ax.legend(fontsize=11); ax.grid(alpha=0.3)\n        ax.spines[[\"top\", \"right\"]].set_visible(False)\n        plt.tight_layout()\n        out = self.OUTPUT / \"convergence_curve.png\"\n        plt.savefig(out, dpi=150, bbox_inches=\"tight\"); plt.show()\n        print(f\"Saved → {out}\")\n\n    #  5. Qualitative side-by-side\n    def qualitative(self, model_zs=None, model_ft_best=None, model_sc_best=None,\n                    n_images: int = 2):\n        \"\"\"Show side-by-side predictions from all available models.\"\"\"\n        available = {}\n        if model_zs is not None:\n            available[\"Zero-shot\"] = model_zs\n\n        ft_path = self.model_ckp_load / \"finetune\" / \"best.pt\"\n        if model_ft_best is not None:\n            available[\"Fine-tuned\"] = model_ft_best\n        elif ft_path.exists():\n            available[\"Fine-tuned\"] = YOLO(str(ft_path))\n\n        sc_path = self.model_ckp_load / \"scratch\" / \"best.pt\"\n        if model_sc_best is not None:\n            available[\"From Scratch\"] = model_sc_best\n        elif sc_path.exists():\n            available[\"From Scratch\"] = YOLO(str(sc_path))\n\n        if len(available) < 2:\n            print(\"Need at least 2 models for comparison.\")\n            return\n\n        test_imgs = glob(str(self.dataset_root / \"test\" / \"images\" / \"*.jpg\"))[:n_images]\n        n = len(available)\n        fig, axes = plt.subplots(len(test_imgs), n,\n                                 figsize=(6 * n, 6 * len(test_imgs)))\n        if len(test_imgs) == 1: axes = [axes]\n\n        for row, img_path in enumerate(test_imgs):\n            for col, (name, mdl) in enumerate(available.items()):\n                preds = mdl.predict(\n                    source  = img_path,\n                    conf    = CONF_THRESHOLD,\n                    iou     = IOU_THRESHOLD,\n                    verbose = False,\n                    classes = self.COCO_FRUIT_CLASSES if name == \"Zero-shot\" else None,\n                )\n                ax = axes[row][col] if len(test_imgs) > 1 else axes[col]\n                ax.imshow(preds[0].plot()[:, :, ::-1])\n                ax.set_title(name, fontsize=11, fontweight=\"bold\")\n                ax.axis(\"off\")\n\n        plt.suptitle(\"Qualitative Comparison — Fruit Detection Test Images\",\n                     fontsize=14, y=1.01)\n        plt.tight_layout(); plt.show()\n\n    #  6. Save final evaluation metrics CSV\n    #     One row per experiment — Precision, Recall, F1, mAP@0.5, mAP@0.5:0.95\n    def save_final_csv(self, all_metrics: list):\n        \"\"\"Save per-experiment evaluation results. One row = one experiment.\"\"\"\n        if not all_metrics:\n            print(\"No metrics to save.\")\n            return\n        df = pd.DataFrame(all_metrics).set_index(\"Experiment\")\n        out = self.OUTPUT / \"final_metrics.csv\"\n        df.to_csv(out)\n        print(f\"Saved → {out}\")\n        print(df)\n\n    #  7. Save unified training history CSV\n    #     One row per epoch with an \"Experiment\" column so it's always clear\n    #     which run each epoch belongs to — no more repeated 1-to-N numbering.\n    def save_history_csv(self, scratch_history: dict = None):\n        \"\"\"Combine fine-tuning and from-scratch epoch histories into one CSV.\"\"\"\n        frames = []\n\n        # Fine-tuning: read YOLO's results.csv and remap to shared column names\n        csv_ft = self.OUTPUT / \"runs\" / \"train\" / \"results.csv\"\n        if csv_ft.exists():\n            df_ft = pd.read_csv(csv_ft)\n            df_ft.columns = df_ft.columns.str.strip()\n            row_ft = pd.DataFrame({\n                \"Experiment\"    : \"Fine-tuning\",\n                \"epoch\"         : df_ft[\"epoch\"],\n                \"train_loss\"    : df_ft.get(\"train/box_loss\",       pd.Series(dtype=float)),\n                \"train_box_loss\": df_ft.get(\"train/box_loss\",       pd.Series(dtype=float)),\n                \"train_cls_loss\": df_ft.get(\"train/cls_loss\",       pd.Series(dtype=float)),\n                \"train_dfl_loss\": df_ft.get(\"train/dfl_loss\",       pd.Series(dtype=float)),\n                \"val_map50\"     : df_ft.get(\"metrics/mAP50(B)\",     pd.Series(dtype=float)),\n                \"val_map5095\"   : df_ft.get(\"metrics/mAP50-95(B)\",  pd.Series(dtype=float)),\n                \"val_precision\" : df_ft.get(\"metrics/precision(B)\", pd.Series(dtype=float)),\n                \"val_recall\"    : df_ft.get(\"metrics/recall(B)\",    pd.Series(dtype=float)),\n            })\n            frames.append(row_ft)\n\n        # From-scratch: use the history dict (same keys as fine-tuning columns)\n        if scratch_history and scratch_history.get(\"epoch\"):\n            df_sc = pd.DataFrame({\n                \"Experiment\"    : \"From Scratch\",\n                \"epoch\"         : scratch_history[\"epoch\"],\n                \"train_loss\"    : scratch_history.get(\"train_loss\",     []),\n                \"train_box_loss\": scratch_history.get(\"train_box_loss\", []),\n                \"train_cls_loss\": scratch_history.get(\"train_cls_loss\", []),\n                \"train_dfl_loss\": scratch_history.get(\"train_dfl_loss\", []),\n                \"val_map50\"     : scratch_history.get(\"val_map50\",      []),\n                \"val_map5095\"   : scratch_history.get(\"val_map5095\",    []),\n                \"val_precision\" : scratch_history.get(\"val_precision\",  []),\n                \"val_recall\"    : scratch_history.get(\"val_recall\",     []),\n            })\n            frames.append(df_sc)\n            df_sc.to_csv(self.OUTPUT / \"runs\" / \"scratch_history.csv\", index=False)\n\n        if not frames:\n            print(\"No training history available — run Experiment 2 or 3 first.\")\n            return\n\n        combined = pd.concat(frames, ignore_index=True)\n        out = self.OUTPUT / \"training_history.csv\"\n        combined.to_csv(out, index=False)\n        print(f\"Saved → {out}\")\n        print(combined.groupby(\"Experiment\")[[\"epoch\", \"val_map50\", \"val_precision\", \"val_recall\"]].describe())\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:47.638902Z","iopub.status.idle":"2026-04-24T18:44:47.639224Z","shell.execute_reply.started":"2026-04-24T18:44:47.639072Z","shell.execute_reply":"2026-04-24T18:44:47.639091Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run full evaluation\neval_suite = EvaluationSuite(OUTPUT=OUTPUT, dataset_root=dataset_root, model_ckp_load=MODEL_CKP_LOAD)\n\nif EXPERIMENT_TO_RUN == \"compare\":\n    all_metrics = eval_suite.load_metrics()\n\nif all_metrics:\n    eval_suite.save_metrics(all_metrics)\n    eval_suite.compare_table(all_metrics)\n    # Pass all_metrics so zero-shot baseline can be drawn in convergence plot\n    eval_suite.convergence(all_metrics=all_metrics)\n    eval_suite.qualitative(\n        model_zs      = model_zs      if \"model_zs\"      in dir() else None,\n        model_ft_best = model_ft_best if \"model_ft_best\" in dir() else None,\n        model_sc_best = model_sc_best if \"model_sc_best\" in dir() else None,\n    )\n    eval_suite.save_final_csv(all_metrics)\n    # Save unified epoch-level training history (replaces confusing repeated epoch numbers)\n    _scratch_hist = history if \"history\" in dir() and history.get(\"epoch\") else None\n    eval_suite.save_history_csv(scratch_history=_scratch_hist)\nelse:\n    print(\"No metrics collected yet — run at least one experiment first.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:47.641060Z","iopub.status.idle":"2026-04-24T18:44:47.641533Z","shell.execute_reply.started":"2026-04-24T18:44:47.641398Z","shell.execute_reply":"2026-04-24T18:44:47.641418Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Clean up stray files from output root ──────────────────────\n\n_to_remove = (\n    glob(str(OUTPUT / \"*.pt\")) +      # auto-downloaded base weights\n    glob(str(OUTPUT / \"*.yaml\"))       # dataset yaml (lives in runs/)\n)\nfor _f in _to_remove:\n    try:\n        Path(_f).unlink()\n        print(f\"Removed from output root: {_f}\")\n    except FileNotFoundError:\n        pass","metadata":{"execution":{"iopub.status.busy":"2026-04-24T18:44:47.642796Z","iopub.status.idle":"2026-04-24T18:44:47.643186Z","shell.execute_reply.started":"2026-04-24T18:44:47.642991Z","shell.execute_reply":"2026-04-24T18:44:47.643016Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}