{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":103103,"databundleVersionId":13042974,"isSourceIdPinned":false}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"title-cell","cell_type":"markdown","source":"# AlphaDent — Segmentation with YOLO11m\n## Reduced Scope: Filling & Caries (2-Class) — ENHANCED\n\n**Enhancement Summary (applied from guide):**\n- ✅ **Model upgraded:** `yolo11s-seg.pt` → `yolo11m-seg.pt` (segmentation, not detection)\n- ✅ **Resolution increased:** `imgsz=640` → `imgsz=1024` (preserves small lesion detail in 15MP images)\n- ✅ **CLAHE preprocessing:** Applied softly on luminance channel (LAB space) only\n- ✅ **Image tiling:** 1024×1024 overlapping patches for high-res input preservation\n- ✅ **Augmentation tuned:** `mosaic=0.5`, `mixup=0`, `copy_paste=0`, `erasing=0.1`\n- ✅ **Learning rate corrected:** `lr0=0.0005` (medical imaging standard)\n- ✅ **Epochs extended:** 100 → 120 with patience=25\n- ✅ **Batch adjusted:** 6 (optimized for 1024px + yolo11m-seg + P100 16GB)\n\n**Classes used:**\n- `0` → **Filling** (original class 1)\n- `1` → **Caries** (original classes 3, 4, 5, 6, 7, 8 — all six caries types merged)\n\n**Dropped:** Abrasion (original 0) and Crown (original 2)  \n**Platform:** Colab / Kaggle\n\n---\n**Sections:**\n- Section 0 — Environment Setup\n- Section 1 — Data Discovery, Class Remapping & Validation\n- Section 1.5 — CLAHE Preprocessing & Image Tiling (NEW)\n- Section 2 — EDA (Filling + Caries only)\n- Section 3 — Pediatric Annotated Samples\n- Section 4 — Imbalance Strategy\n- Section 5 — Before/After Class Distribution\n- Section 6 — YAML Generation & Model Training\n- Section 7 — Evaluation Metrics\n- Section 8 — Predictions & Visualization\n- Section 9 — Save Outputs\n- Section 10 — Final Summary","metadata":{}},{"id":"sec0-hdr","cell_type":"markdown","source":"---\n## Section 0 — Environment Setup","metadata":{}},{"id":"58b9c9c1-7fed-43ff-b6f7-7ceff1d27406","cell_type":"code","source":"import shutil\n\nshutil.copytree(\n    \"/kaggle/input/competitions/alpha-dent/AlphaDent\",\n    \"/kaggle/working/AlphaDent\"\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec0-install","cell_type":"code","source":"import subprocess, sys\ndef install(pkg):\n    subprocess.check_call([sys.executable, '-m', 'pip', 'install', '-q', pkg])\n\ninstall('ultralytics>=8.0.0')\ninstall('opencv-python-headless')\ninstall('pandas')\ninstall('matplotlib')\ninstall('scikit-learn')\ninstall('seaborn')\ninstall('tqdm')\nprint('Packages installed.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec0-imports","cell_type":"code","source":"import os, sys, random, re, json, warnings\nimport numpy as np\nimport pandas as pd\nimport matplotlib\nmatplotlib.use('Agg')\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport cv2\nimport torch\nimport seaborn as sns\nfrom pathlib import Path\nfrom collections import defaultdict, Counter\nimport ultralytics\nfrom ultralytics import YOLO\n\nwarnings.filterwarnings('ignore')\n\nprint(f'Python     : {sys.version}')\nprint(f'PyTorch    : {torch.__version__}')\nprint(f'Ultralytics: {ultralytics.__version__}')\nprint(f'OpenCV     : {cv2.__version__}')\nif torch.cuda.is_available():\n    print(f'CUDA       : {torch.version.cuda}')\n    print(f'GPU        : {torch.cuda.get_device_name(0)}')\nelse:\n    print('CUDA       : Not available (CPU mode)')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec0-seeds","cell_type":"code","source":"SEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\nprint(f'Seed set to {SEED}.')\nprint('Note: full GPU reproducibility is best-effort; minor metric variance is expected.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec0-paths","cell_type":"code","source":"# ── Paths ──────────────────────────────────────────────────────────────────\nDATASET_ROOT = Path('/kaggle/input/competitions/alpha-dent/AlphaDent')   # <-- change if needed\n\nIMAGES_TRAIN = DATASET_ROOT / 'images' / 'train'\nIMAGES_VALID = DATASET_ROOT / 'images' / 'valid'\nIMAGES_TEST  = DATASET_ROOT / 'images' / 'test'\nLABELS_TRAIN = DATASET_ROOT / 'labels' / 'train'\nLABELS_VALID = DATASET_ROOT / 'labels' / 'valid'\n\n# All remapped labels go here (we never overwrite originals)\nREMAP_LABELS_TRAIN = Path('alphadent_2class/labels/train')\nREMAP_LABELS_VALID = Path('alphadent_2class/labels/valid')\n\n# ── ENHANCEMENT: Tiled dataset paths ───────────────────────────────────────\n# Image tiling: 1024×1024 patches with 15% overlap\n# Preserves fine lesion detail from 5000×3000 source images\nTILE_ROOT          = Path('alphadent_2class_tiled')\nTILE_IMAGES_TRAIN  = TILE_ROOT / 'images' / 'train'\nTILE_IMAGES_VALID  = TILE_ROOT / 'images' / 'valid'\nTILE_LABELS_TRAIN  = TILE_ROOT / 'labels' / 'train'\nTILE_LABELS_VALID  = TILE_ROOT / 'labels' / 'valid'\n\n# ── ENHANCEMENT: CLAHE preprocessed image paths ────────────────────────────\nCLAHE_ROOT         = Path('alphadent_2class_clahe')\nCLAHE_IMAGES_TRAIN = CLAHE_ROOT / 'images' / 'train'\nCLAHE_IMAGES_VALID = CLAHE_ROOT / 'images' / 'valid'\n\n# Output dirs\nRUNS_ROOT       = Path('runs/analysis')\nEDA_DIR         = RUNS_ROOT / 'eda_2class'\nPEDIATRIC_DIR   = RUNS_ROOT / 'pediatric_annotated_samples'\nIMBALANCE_DIR   = RUNS_ROOT / 'imbalance_strategy_stats'\nCURVES_DIR      = RUNS_ROOT / 'training_curves'\nEVAL_DIR        = RUNS_ROOT / 'eval_2class'\nPREDS_VALID_DIR = RUNS_ROOT / 'preds_valid'\nPREDS_TEST_DIR  = RUNS_ROOT / 'preds_test'\n\nfor d in [REMAP_LABELS_TRAIN, REMAP_LABELS_VALID,\n          TILE_IMAGES_TRAIN, TILE_IMAGES_VALID,\n          TILE_LABELS_TRAIN, TILE_LABELS_VALID,\n          CLAHE_IMAGES_TRAIN, CLAHE_IMAGES_VALID,\n          EDA_DIR, PEDIATRIC_DIR, IMBALANCE_DIR,\n          CURVES_DIR, EVAL_DIR, PREDS_VALID_DIR, PREDS_TEST_DIR]:\n    d.mkdir(parents=True, exist_ok=True)\n\n# ── Class remapping ────────────────────────────────────────────────────────\n# Original 9 classes:\n#   0=Abrasion  1=Filling  2=Crown\n#   3=CarC1  4=CarC2  5=CarC3  6=CarC4  7=CarC5  8=CarC6\n#\n# New 2 classes:\n#   0=Filling   (original 1)\n#   1=Caries    (original 3,4,5,6,7,8)\n#   DROPPED: original 0 (Abrasion), original 2 (Crown)\n\nORIG_TO_NEW = {\n    1: 0,   # Filling  -> new class 0\n    3: 1,   # Caries C1 -> new class 1\n    4: 1,   # Caries C2 -> new class 1\n    5: 1,   # Caries C3 -> new class 1\n    6: 1,   # Caries C4 -> new class 1\n    7: 1,   # Caries C5 -> new class 1\n    8: 1,   # Caries C6 -> new class 1\n}  # keys NOT in this dict (0=Abrasion, 2=Crown) are DROPPED\n\nDROPPED_CLASSES = {0, 2}   # Abrasion, Crown\nNUM_CLASSES     = 2\nCLASS_NAMES     = ['Filling', 'Caries']\n\nCLASS_COLORS = plt.cm.get_cmap('tab10', NUM_CLASSES)\ndef get_color(cls_id):\n    r, g, b, _ = CLASS_COLORS(cls_id)\n    return (int(b*255), int(g*255), int(r*255))\n\n# ── Enhancement flags ───────────────────────────────────────────────────────\nUSE_TILING = False          # Enable 1024×1024 overlapping tile strategy\nTILE_SIZE  = 1024          # Tile width/height in pixels\nTILE_OVERLAP = 0.15        # 15% overlap between tiles\nUSE_CLAHE  = False          # Enable soft CLAHE on LAB luminance channel\nCLAHE_CLIP = 2.0           # Clip limit (soft; higher = more contrast)\nCLAHE_GRID = (8, 8)        # Tile grid size for CLAHE\n\nprint('Class remapping:')\nprint('  Original -> New')\nfor orig, new in ORIG_TO_NEW.items():\n    print(f'    {orig} -> {new} ({CLASS_NAMES[new]})')\nprint(f'  Dropped: {DROPPED_CLASSES}')\nprint(f'Final classes: {CLASS_NAMES}')\nprint()\nprint(f'Enhancement flags:')\nprint(f'  USE_TILING : {USE_TILING}  (TILE_SIZE={TILE_SIZE}px, overlap={TILE_OVERLAP*100:.0f}%)')\nprint(f'  USE_CLAHE  : {USE_CLAHE}  (clip={CLAHE_CLIP}, grid={CLAHE_GRID})')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec1-hdr","cell_type":"markdown","source":"---\n## Section 1 — Data Discovery, Class Remapping & Validation","metadata":{}},{"id":"sec1-utils","cell_type":"code","source":"IMAGE_EXTS = {'.jpg', '.jpeg', '.png'}\n\ndef find_images(folder):\n    return sorted([p for p in folder.iterdir() if p.suffix.lower() in IMAGE_EXTS])\n\ndef parse_filename(fname):\n    stem = Path(fname).stem\n    m = re.match(r'^(p\\d+)_([MF])_(\\d+)_(\\d+)$', stem, re.IGNORECASE)\n    if m:\n        return dict(patient_id=m.group(1), gender=m.group(2).upper(),\n                    age=int(m.group(3)), img_idx=int(m.group(4)), parsed=True)\n    return dict(patient_id=None, gender=None, age=None, img_idx=None, parsed=False)\n\ndef parse_seg_label_file(path):\n    \"\"\"\n    Parse a YOLO segmentation label file.\n    Returns list of dicts: {orig_class_id, x_center, y_center, width, height}\n    Polygon coords are collapsed to tight bounding box.\n    \"\"\"\n    if not Path(path).exists():\n        return []\n    boxes = []\n    with open(path) as f:\n        for line in f:\n            parts = line.strip().split()\n            if len(parts) < 5:\n                continue\n            cls = int(parts[0])\n            coords = list(map(float, parts[1:]))\n            if len(coords) == 4:\n                xc, yc, w, h = coords\n            elif len(coords) >= 6 and len(coords) % 2 == 0:\n                xs = coords[0::2]; ys = coords[1::2]\n                x1, x2 = min(xs), max(xs)\n                y1, y2 = min(ys), max(ys)\n                xc = (x1+x2)/2; yc = (y1+y2)/2\n                w  = x2-x1;     h  = y2-y1\n            else:\n                continue\n            boxes.append(dict(orig_class_id=cls,\n                              x_center=xc, y_center=yc,\n                              width=max(w,1e-6), height=max(h,1e-6)))\n    return boxes\n\nprint('Utility functions defined.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec1-check-dirs","cell_type":"code","source":"# Verify directory structure\nfor name, d in [('images/train', IMAGES_TRAIN), ('images/valid', IMAGES_VALID),\n                ('images/test',  IMAGES_TEST),  ('labels/train', LABELS_TRAIN),\n                ('labels/valid', LABELS_VALID)]:\n    assert d.exists(), f'Missing: {d}'\n    print(f'OK  {name}: {d}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec1-remap","cell_type":"code","source":"# ── Remap & write new label files ─────────────────────────────────────────\n# For each original label file we:\n#   1. Skip boxes whose orig class is in DROPPED_CLASSES\n#   2. Remap kept classes via ORIG_TO_NEW\n#   3. Write a new .txt file in detection format (5 values per line)\n#   4. If no boxes remain after filtering → write empty file (image retained\n#      so EDA can reference it, but it contributes 0 GT instances)\n\ndef remap_label_file(src_label, dst_label, orig_to_new, dropped):\n    \"\"\"\n    Returns dict with stats: n_kept, n_dropped, n_remapped_classes\n    \"\"\"\n    boxes = parse_seg_label_file(src_label)\n    kept = []\n    n_dropped = 0\n    for b in boxes:\n        oc = b['orig_class_id']\n        if oc in dropped:\n            n_dropped += 1\n            continue\n        if oc not in orig_to_new:\n            # class not in mapping and not dropped -> skip silently\n            n_dropped += 1\n            continue\n        nc = orig_to_new[oc]\n        kept.append(f\"{nc} {b['x_center']:.6f} {b['y_center']:.6f} \"\n                    f\"{b['width']:.6f} {b['height']:.6f}\")\n    Path(dst_label).parent.mkdir(parents=True, exist_ok=True)\n    with open(dst_label, 'w') as f:\n        f.write('\\n'.join(kept) + ('\\n' if kept else ''))\n    return dict(n_kept=len(kept), n_dropped=n_dropped)\n\n\ndef remap_split(img_list, src_label_dir, dst_label_dir, split_name):\n    stats = []\n    for ip in img_list:\n        src = src_label_dir / (ip.stem + '.txt')\n        dst = dst_label_dir / (ip.stem + '.txt')\n        if src.exists():\n            s = remap_label_file(src, dst, ORIG_TO_NEW, DROPPED_CLASSES)\n        else:\n            # No label -> write empty\n            dst.parent.mkdir(parents=True, exist_ok=True)\n            open(dst, 'w').close()\n            s = dict(n_kept=0, n_dropped=0)\n        meta = parse_filename(ip.name)\n        stats.append(dict(filename=ip.name, split=split_name,\n                          n_kept=s['n_kept'], n_dropped=s['n_dropped'],\n                          **meta))\n    return pd.DataFrame(stats)\n\n\nprint('Scanning images...')\ntrain_imgs = find_images(IMAGES_TRAIN)\nvalid_imgs = find_images(IMAGES_VALID)\ntest_imgs  = find_images(IMAGES_TEST)\nprint(f'  Train: {len(train_imgs)}  Valid: {len(valid_imgs)}  Test: {len(test_imgs)}')\n\nprint('\\nRemapping train labels...')\ndf_train = remap_split(train_imgs, LABELS_TRAIN, REMAP_LABELS_TRAIN, 'train')\nprint('Remapping valid labels...')\ndf_valid = remap_split(valid_imgs, LABELS_VALID, REMAP_LABELS_VALID, 'valid')\n\nprint(f'\\n=== REMAP SUMMARY ===')\nfor split, df in [('train', df_train), ('valid', df_valid)]:\n    total_kept    = df['n_kept'].sum()\n    total_dropped = df['n_dropped'].sum()\n    empty_imgs    = (df['n_kept'] == 0).sum()\n    print(f'  {split}: {len(df)} images | '\n          f'{total_kept} boxes kept | '\n          f'{total_dropped} boxes dropped (Abrasion+Crown) | '\n          f'{empty_imgs} images with 0 remaining boxes')\n\n# Images with ZERO remaining boxes are kept in the dataset\n# (they are negative examples for Filling & Caries detection)\nprint('\\nImages with 0 remaining boxes are kept as hard negatives.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec1-load-boxes","cell_type":"code","source":"# ── Load all remapped boxes into flat DataFrames ──────────────────────────\ndef load_remapped_boxes(img_list, remap_label_dir, split_name):\n    rows = []\n    for ip in img_list:\n        lp = remap_label_dir / (ip.stem + '.txt')\n        meta = parse_filename(ip.name)\n        if not lp.exists():\n            continue\n        with open(lp) as f:\n            for line in f:\n                parts = line.strip().split()\n                if len(parts) != 5:\n                    continue\n                nc, xc, yc, w, h = int(parts[0]), *map(float, parts[1:])\n                rows.append(dict(\n                    filename=ip.name, split=split_name,\n                    class_id=nc, class_name=CLASS_NAMES[nc],\n                    x_center=xc, y_center=yc, width=w, height=h,\n                    area=w*h, aspect_ratio=w/h if h>0 else float('nan'),\n                    patient_id=meta['patient_id'],\n                    gender=meta['gender'], age=meta['age']\n                ))\n    return pd.DataFrame(rows)\n\ndf_train_boxes = load_remapped_boxes(train_imgs, REMAP_LABELS_TRAIN, 'train')\ndf_valid_boxes = load_remapped_boxes(valid_imgs, REMAP_LABELS_VALID, 'valid')\ndf_all_boxes   = pd.concat([df_train_boxes, df_valid_boxes], ignore_index=True)\n\nprint(f'Train boxes (Filling+Caries): {len(df_train_boxes)}')\nprint(f'Valid boxes (Filling+Caries): {len(df_valid_boxes)}')\nprint(f'Total boxes                 : {len(df_all_boxes)}')\nprint()\nprint('Train class distribution (remapped):')\nprint(df_train_boxes['class_name'].value_counts().to_string())\nprint()\nprint('Valid class distribution (remapped):')\nprint(df_valid_boxes['class_name'].value_counts().to_string())\n\n# Integrity check\nassert df_all_boxes['class_id'].between(0, NUM_CLASSES-1).all(), \\\n    'Class ID out of range after remapping!'\nassert ((df_all_boxes[['x_center','y_center','width','height']] >= 0) &\n        (df_all_boxes[['x_center','y_center','width','height']] <= 1)).all().all(), \\\n    'Normalised coords out of [0,1]!'\nprint('All integrity checks passed.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec15-hdr","cell_type":"markdown","source":"---\n## Section 1.5 — CLAHE Preprocessing & Image Tiling (ENHANCEMENT)\n\n### Why this matters\nYour source images are **5000×3000 pixels (15MP)** but training previously used `imgsz=640`.\nDownscaling 15MP → 640px destroys the fine texture details that distinguish dental caries.\n\n**Two complementary strategies applied:**\n\n1. **CLAHE (Contrast Limited Adaptive Histogram Equalization)**  \n   Applied *softly* on the LAB luminance channel only — enhances local contrast  \n   for subtle lesion boundaries without amplifying noise.\n\n2. **Image Tiling (1024×1024 with 15% overlap)**  \n   Slices each high-res image into overlapping patches before training.  \n   - Preserves fine lesion detail  \n   - Naturally augments dataset size  \n   - Dramatically improves small-object detection (expected +8–15% mAP)  \n   - Polygon masks are transformed correctly per-tile","metadata":{}},{"id":"sec15-clahe","cell_type":"code","source":"# ── CLAHE Preprocessing ───────────────────────────────────────────────────\n# Apply soft CLAHE on LAB luminance channel only.\n# Medical rule: subtle enhancement > aggressive augmentation.\n\ndef apply_clahe_lab(img_bgr, clip_limit=2.0, tile_grid=(8,8)):\n    \"\"\"\n    Apply CLAHE on the L (luminance) channel of LAB color space.\n    Preserves color hue/saturation while enhancing local contrast.\n    \"\"\"\n    lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB)\n    l, a, b_ch = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n    l_eq = clahe.apply(l)\n    lab_eq = cv2.merge([l_eq, a, b_ch])\n    return cv2.cvtColor(lab_eq, cv2.COLOR_LAB2BGR)\n\ndef preprocess_clahe_split(src_img_dir, dst_img_dir, clip=2.0, grid=(8,8)):\n    \"\"\"Apply CLAHE to all images in src_img_dir, save to dst_img_dir.\"\"\"\n    src_imgs = find_images(src_img_dir)\n    print(f'  Applying CLAHE to {len(src_imgs)} images in {src_img_dir.name}...')\n    for ip in src_imgs:\n        img = cv2.imread(str(ip))\n        if img is None:\n            print(f'    [WARN] Could not read: {ip.name}')\n            continue\n        out = apply_clahe_lab(img, clip_limit=clip, tile_grid=grid)\n        cv2.imwrite(str(dst_img_dir / ip.name), out)\n    print(f'  Saved {len(src_imgs)} CLAHE images -> {dst_img_dir}')\n\nif USE_CLAHE:\n    print('=== Applying CLAHE Preprocessing ===')\n    preprocess_clahe_split(IMAGES_TRAIN, CLAHE_IMAGES_TRAIN, CLAHE_CLIP, CLAHE_GRID)\n    preprocess_clahe_split(IMAGES_VALID, CLAHE_IMAGES_VALID, CLAHE_CLIP, CLAHE_GRID)\n    print('CLAHE preprocessing complete.')\n    \n    # Show before/after comparison for 2 sample images\n    sample_imgs = find_images(IMAGES_TRAIN)[:2]\n    fig, axes = plt.subplots(2, 2, figsize=(14, 8))\n    fig.suptitle('CLAHE Enhancement — Before vs After (LAB luminance)', fontsize=13, fontweight='bold')\n    for row, ip in enumerate(sample_imgs):\n        orig = cv2.imread(str(ip))\n        clahe_img = cv2.imread(str(CLAHE_IMAGES_TRAIN / ip.name))\n        if orig is None or clahe_img is None:\n            continue\n        # Show center crop for visibility\n        h, w = orig.shape[:2]\n        cy, cx = h//2, w//2\n        crop_size = min(800, h//2, w//2)\n        orig_crop  = orig[cy-crop_size//2:cy+crop_size//2, cx-crop_size//2:cx+crop_size//2]\n        clahe_crop = clahe_img[cy-crop_size//2:cy+crop_size//2, cx-crop_size//2:cx+crop_size//2]\n        axes[row,0].imshow(cv2.cvtColor(orig_crop, cv2.COLOR_BGR2RGB))\n        axes[row,0].set_title(f'Original: {ip.name[:25]}', fontsize=9)\n        axes[row,0].axis('off')\n        axes[row,1].imshow(cv2.cvtColor(clahe_crop, cv2.COLOR_BGR2RGB))\n        axes[row,1].set_title(f'After CLAHE (clip={CLAHE_CLIP})', fontsize=9)\n        axes[row,1].axis('off')\n    plt.tight_layout()\n    fig.savefig(EDA_DIR / 'clahe_comparison.png', dpi=100, bbox_inches='tight')\n    plt.show()\n    print(f'Saved: {EDA_DIR}/clahe_comparison.png')\nelse:\n    print('CLAHE disabled (USE_CLAHE=False). Using original images.')\n    CLAHE_IMAGES_TRAIN = IMAGES_TRAIN\n    CLAHE_IMAGES_VALID = IMAGES_VALID","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec15-tiling","cell_type":"code","source":"# ── Image Tiling — 1024×1024 patches with overlap ─────────────────────────\n# Rationale:\n#   Source images are 5000×3000px. Dental caries are small texture lesions.\n#   Tiling preserves full detail vs blindly downscaling to 640px.\n#   Polygon masks (YOLO seg format) are properly clipped per tile.\n#\n# Expected benefit: +8–15% mAP for small-object (caries) detection.\n\ndef tile_image_and_masks(img_path, label_path, out_img_dir, out_lbl_dir,\n                          tile_size=1024, overlap=0.15):\n    \"\"\"\n    Tile a single image + YOLO-format polygon label file into patches.\n    \n    Args:\n        img_path:    Path to source image\n        label_path:  Path to YOLO segmentation label (.txt) or None\n        out_img_dir: Directory for output tile images\n        out_lbl_dir: Directory for output tile labels\n        tile_size:   Tile size in pixels (square)\n        overlap:     Fractional overlap between adjacent tiles\n    \n    YOLO seg format per line: class_id x1 y1 x2 y2 ... xn yn\n    All coordinates are normalized [0,1] relative to image size.\n    \"\"\"\n    img = cv2.imread(str(img_path))\n    if img is None:\n        return 0\n    \n    H, W = img.shape[:2]\n    stem  = img_path.stem\n    \n    # Read original label lines\n    raw_lines = []\n    if label_path and label_path.exists():\n        with open(label_path) as f:\n            raw_lines = [l.strip() for l in f if l.strip()]\n    \n    stride = int(tile_size * (1 - overlap))\n    xs = list(range(0, max(W - tile_size, 0) + 1, stride))\n    ys = list(range(0, max(H - tile_size, 0) + 1, stride))\n    # Always include last position to cover right/bottom edge\n    if xs and xs[-1] + tile_size < W: xs.append(W - tile_size)\n    if ys and ys[-1] + tile_size < H: ys.append(H - tile_size)\n    if not xs: xs = [0]\n    if not ys: ys = [0]\n    \n    n_tiles = 0\n    for ty in ys:\n        for tx in xs:\n            x0 = max(0, tx)\n            y0 = max(0, ty)\n            x1 = min(W, tx + tile_size)\n            y1 = min(H, ty + tile_size)\n            tw = x1 - x0\n            th = y1 - y0\n            \n            tile_img = img[y0:y1, x0:x1]\n            tile_name = f'{stem}_tx{tx}_ty{ty}.jpg'\n            cv2.imwrite(str(out_img_dir / tile_name), tile_img,\n                        [cv2.IMWRITE_JPEG_QUALITY, 95])\n            \n            # Transform polygon coordinates to tile space\n            tile_lines = []\n            for line in raw_lines:\n                parts = line.split()\n                if len(parts) < 5:  # need class + at least 2 points\n                    continue\n                cls_id = parts[0]\n                coords = list(map(float, parts[1:]))\n                # Convert normalized -> absolute\n                pts = [(coords[i]*W, coords[i+1]*H)\n                       for i in range(0, len(coords)-1, 2)]\n                # Shift to tile coordinates (still absolute)\n                pts_tile = [(px - x0, py - y0) for px, py in pts]\n                # Clip to tile bounds\n                pts_clip = [(max(0, min(tw-1, px)), max(0, min(th-1, py)))\n                            for px, py in pts_tile]\n                # Check if polygon has meaningful area in this tile\n                valid = [p for p in pts_tile if 0 <= p[0] < tw and 0 <= p[1] < th]\n                if len(valid) < 3:\n                    continue\n                # Normalize back to [0,1] relative to tile\n                norm = [v for px, py in pts_clip\n                        for v in (px/tw, py/th)]\n                tile_lines.append(cls_id + ' ' + ' '.join(f'{v:.6f}' for v in norm))\n            \n            # Write label (even empty, so YOLO knows it's a background tile)\n            lbl_name = f'{stem}_tx{tx}_ty{ty}.txt'\n            with open(out_lbl_dir / lbl_name, 'w') as f:\n                f.write('\\n'.join(tile_lines))\n            n_tiles += 1\n    \n    return n_tiles\n\ndef tile_split(src_img_dir, src_lbl_dir, dst_img_dir, dst_lbl_dir,\n               tile_size=1024, overlap=0.15):\n    \"\"\"Tile all images+labels in a split directory.\"\"\"\n    imgs = find_images(src_img_dir)\n    total_tiles = 0\n    print(f'  Tiling {len(imgs)} images from {src_img_dir.name} ...')\n    for ip in imgs:\n        lp = (src_lbl_dir / ip.stem).with_suffix('.txt')\n        n = tile_image_and_masks(ip, lp if lp.exists() else None,\n                                  dst_img_dir, dst_lbl_dir,\n                                  tile_size=tile_size, overlap=overlap)\n        total_tiles += n\n    print(f'  Generated {total_tiles} tiles -> {dst_img_dir}')\n    return total_tiles\n\nif USE_TILING:\n    # Use CLAHE images as input to tiling if CLAHE was applied\n    src_train_imgs = CLAHE_IMAGES_TRAIN if USE_CLAHE else IMAGES_TRAIN\n    src_valid_imgs = CLAHE_IMAGES_VALID if USE_CLAHE else IMAGES_VALID\n    \n    print('=== Image Tiling (1024×1024, 15% overlap) ===')\n    n_train = tile_split(src_train_imgs, REMAP_LABELS_TRAIN,\n                          TILE_IMAGES_TRAIN, TILE_LABELS_TRAIN,\n                          TILE_SIZE, TILE_OVERLAP)\n    n_valid = tile_split(src_valid_imgs, REMAP_LABELS_VALID,\n                          TILE_IMAGES_VALID, TILE_LABELS_VALID,\n                          TILE_SIZE, TILE_OVERLAP)\n    \n    print(f'\\nTiling complete:')\n    print(f'  Train tiles: {n_train}')\n    print(f'  Valid tiles: {n_valid}')\n    \n    # Show sample tiles\n    tile_samples = find_images(TILE_IMAGES_TRAIN)[:4]\n    fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n    fig.suptitle('Sample 1024×1024 Tiles (from CLAHE-enhanced images)', fontsize=12, fontweight='bold')\n    for ax, tp in zip(axes, tile_samples):\n        tile = cv2.imread(str(tp))\n        if tile is not None:\n            ax.imshow(cv2.cvtColor(tile, cv2.COLOR_BGR2RGB))\n        ax.set_title(tp.name[:20], fontsize=7)\n        ax.axis('off')\n    plt.tight_layout()\n    fig.savefig(EDA_DIR / 'sample_tiles.png', dpi=100, bbox_inches='tight')\n    plt.show()\n    print(f'Saved: {EDA_DIR}/sample_tiles.png')\n    \n    # Effective training sources for YAML\n    FINAL_TRAIN_IMAGES = str(TILE_IMAGES_TRAIN)\n    FINAL_VALID_IMAGES = str(TILE_IMAGES_VALID)\n    FINAL_TRAIN_LABELS = str(TILE_LABELS_TRAIN)\n    FINAL_VALID_LABELS = str(TILE_LABELS_VALID)\n    print(f'\\nTraining will use tiled images at: {TILE_IMAGES_TRAIN}')\nelse:\n    print('Tiling disabled (USE_TILING=False). Using full-size images.')\n    src_train_imgs = CLAHE_IMAGES_TRAIN if USE_CLAHE else IMAGES_TRAIN\n    src_valid_imgs = CLAHE_IMAGES_VALID if USE_CLAHE else IMAGES_VALID\n    FINAL_TRAIN_IMAGES = str(src_train_imgs)\n    FINAL_VALID_IMAGES = str(src_valid_imgs)\n    FINAL_TRAIN_LABELS = str(REMAP_LABELS_TRAIN)\n    FINAL_VALID_LABELS = str(REMAP_LABELS_VALID)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec2-hdr","cell_type":"markdown","source":"---\n## Section 2 — EDA (Filling + Caries Only)","metadata":{}},{"id":"sec2-demo","cell_type":"code","source":"# ── 2.1 Demographics ──────────────────────────────────────────────────────\n%matplotlib inline\ndf_meta = pd.concat([df_train, df_valid], ignore_index=True)\ndf_meta_p = df_meta[df_meta['parsed']]\n\nfig, axes = plt.subplots(2, 2, figsize=(16, 11))\nfig.suptitle('AlphaDent 2-Class Dataset — Demographics & File Stats (Train+Valid)',\n             fontsize=14, fontweight='bold')\n\n# Age histogram\nax = axes[0, 0]\nages = df_meta_p['age'].dropna().astype(int)\nax.hist(ages, bins=range(ages.min(), ages.max()+2),\n        color='steelblue', edgecolor='white', alpha=0.85)\nax.axvspan(6, 11, alpha=0.18, color='orange', label='Pediatric (6-11)')\nax.set_xlabel('Patient Age'); ax.set_ylabel('Images')\nax.set_title('Age Distribution'); ax.legend()\n\n# Gender bar\nax = axes[0, 1]\ngc = df_meta_p['gender'].value_counts()\nbars = ax.bar(gc.index, gc.values, color=['#5B9BD5','#ED7D31'], edgecolor='white')\nfor bar, v in zip(bars, gc.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+2, str(v), ha='center')\nax.set_xlabel('Gender'); ax.set_ylabel('Images')\nax.set_title('Gender Distribution')\n\n# Images per patient\nax = axes[1, 0]\nipp = df_meta_p.groupby('patient_id').size()\nax.hist(ipp, bins=20, color='mediumseagreen', edgecolor='white', alpha=0.85)\nax.axvline(ipp.mean(), color='red', linestyle='--', label=f'Mean={ipp.mean():.1f}')\nax.set_xlabel('Images per Patient'); ax.set_ylabel('Patients')\nax.set_title('Images per Patient'); ax.legend()\n\n# Boxes remaining per image after remapping (train+valid)\nax = axes[1, 1]\nbpi = df_all_boxes.groupby('filename').size()\nax.hist(bpi, bins=30, color='mediumpurple', edgecolor='white', alpha=0.85)\nax.axvline(bpi.mean(), color='red', linestyle='--', label=f'Mean={bpi.mean():.1f}')\nax.set_xlabel('Boxes per Image (Filling+Caries only)')\nax.set_ylabel('Images')\nax.set_title('Boxes per Image (After Remapping)')\nax.legend()\n\nplt.tight_layout()\nfig.savefig(EDA_DIR / 'demographics.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {EDA_DIR}/demographics.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec2-label-stats","cell_type":"code","source":"# ── 2.2 Label Statistics ──────────────────────────────────────────────────\ncolors_bar = ['#4C72B0', '#DD8452']  # blue=Filling, orange=Caries\n\n# Boxes per class\nbox_counts = df_all_boxes['class_id'].value_counts().sort_index()\nfor i in range(NUM_CLASSES):\n    if i not in box_counts: box_counts[i] = 0\nbox_counts = box_counts.sort_index()\n\n# Images containing each class\nimgs_per_class = df_all_boxes.groupby('class_id')['filename'].nunique()\nfor i in range(NUM_CLASSES):\n    if i not in imgs_per_class: imgs_per_class[i] = 0\nimgs_per_class = imgs_per_class.sort_index()\n\nfig, axes = plt.subplots(2, 3, figsize=(20, 11))\nfig.suptitle('AlphaDent 2-Class — Label Statistics (Filling + Caries, Train+Valid)',\n             fontsize=13, fontweight='bold')\n\nx_labels = [f'{i}\\n{CLASS_NAMES[i]}' for i in range(NUM_CLASSES)]\n\n# 1. Boxes per class\nax = axes[0, 0]\nbars = ax.bar(range(NUM_CLASSES), box_counts.values, color=colors_bar, edgecolor='white')\nfor bar, v in zip(bars, box_counts.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+10, str(v), ha='center')\nax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(x_labels)\nax.set_ylabel('Box Count'); ax.set_title('Boxes per Class')\n\n# 2. Images per class\nax = axes[0, 1]\nbars = ax.bar(range(NUM_CLASSES), imgs_per_class.values, color=colors_bar, edgecolor='white')\nfor bar, v in zip(bars, imgs_per_class.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+2, str(v), ha='center')\nax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(x_labels)\nax.set_ylabel('Image Count'); ax.set_title('Images Containing Each Class')\n\n# 3. Caries sub-type breakdown BEFORE merging (use original label files)\nax = axes[0, 2]\ncaries_orig_ids  = [3, 4, 5, 6, 7, 8]\ncaries_names     = ['C1','C2','C3','C4','C5','C6']\ncaries_train_orig_counts = []\nfor c in caries_orig_ids:\n    cnt = 0\n    for ip in train_imgs:\n        lp = LABELS_TRAIN / (ip.stem + '.txt')\n        if not lp.exists(): continue\n        boxes = parse_seg_label_file(lp)\n        cnt += sum(1 for b in boxes if b['orig_class_id'] == c)\n    caries_train_orig_counts.append(cnt)\ncaries_colors = plt.cm.get_cmap('Reds', len(caries_orig_ids)+2)\nbar_colors = [caries_colors(i+2) for i in range(len(caries_orig_ids))]\nbars = ax.bar(caries_names, caries_train_orig_counts, color=bar_colors, edgecolor='white')\nfor bar, v in zip(bars, caries_train_orig_counts):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+2, str(v), ha='center', fontsize=9)\nax.set_xlabel('Caries Sub-type (original)')\nax.set_ylabel('Box Count (train)')\nax.set_title('Caries Sub-type Breakdown (before merging)')\n\n# 4. Box area distribution\nax = axes[1, 0]\nfor i, name in enumerate(CLASS_NAMES):\n    sub = df_all_boxes[df_all_boxes['class_id']==i]['area']\n    ax.hist(sub.clip(0, 0.3), bins=40, alpha=0.6, color=colors_bar[i], label=name, edgecolor='white')\nax.set_xlabel('Normalised Box Area (w*h)'); ax.set_ylabel('Count')\nax.set_title('Box Area Distribution'); ax.legend()\n\n# 5. Aspect ratio distribution\nax = axes[1, 1]\nfor i, name in enumerate(CLASS_NAMES):\n    ar = df_all_boxes[df_all_boxes['class_id']==i]['aspect_ratio'].replace([np.inf,-np.inf],np.nan).dropna()\n    ax.hist(ar.clip(0,5), bins=40, alpha=0.6, color=colors_bar[i], label=name, edgecolor='white')\nax.set_xlabel('Aspect Ratio (w/h)'); ax.set_ylabel('Count')\nax.set_title('Aspect Ratio Distribution'); ax.legend()\n\n# 6. Imbalance ratio\nax = axes[1, 2]\nmax_count = box_counts.max()\nratio = box_counts / max_count\nbars = ax.bar(range(NUM_CLASSES), ratio.values, color=colors_bar, edgecolor='white')\nfor bar, v in zip(bars, ratio.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+0.01, f'{v:.3f}', ha='center')\nax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(x_labels)\nax.set_ylabel('Ratio vs. majority class'); ax.set_title('Class Imbalance Ratio')\nax.axhline(1.0, color='gray', linestyle='--', alpha=0.5)\n\nplt.tight_layout()\nfig.savefig(EDA_DIR / 'label_stats.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {EDA_DIR}/label_stats.png')\n\nprint('\\n=== Class Distribution Summary (Train+Valid, 2-class) ===')\nsummary = pd.DataFrame({'class_id': range(NUM_CLASSES), 'class_name': CLASS_NAMES,\n                        'box_count': box_counts.values, 'img_count': imgs_per_class.values,\n                        'imbalance_ratio': ratio.values})\nprint(summary.to_string(index=False))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec3-hdr","cell_type":"markdown","source":"---\n## Section 3 — Pediatric Annotated Samples (Filling + Caries, age 6-11)","metadata":{}},{"id":"b744c6e0-0ccb-4694-9a20-ab1bb555b502","cell_type":"code","source":"def get_contrast_color(bgr):\n    # Convert BGR to perceived brightness\n    brightness = (0.299*bgr[2] + 0.587*bgr[1] + 0.114*bgr[0])\n    return (0,0,0) if brightness > 150 else (255,255,255)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec3-pediatric","cell_type":"code","source":"PED_MIN, PED_MAX = 6, 11\nN_SAMPLES = 3\n\nped_train = df_train[df_train['parsed'] & df_train['age'].between(PED_MIN, PED_MAX)]\nped_valid = df_valid[df_valid['parsed'] & df_valid['age'].between(PED_MIN, PED_MAX)]\nped_all   = pd.concat([ped_train, ped_valid], ignore_index=True)\n\nprint(f'Pediatric images (age {PED_MIN}-{PED_MAX}): {len(ped_all)}')\nif len(ped_all) == 0:\n    PED_MAX = 18\n    ped_all = pd.concat([df_train[df_train['parsed'] & df_train['age'].between(PED_MIN, PED_MAX)],\n                         df_valid[df_valid['parsed'] & df_valid['age'].between(PED_MIN, PED_MAX)]], ignore_index=True)\n    print(f'Extended to age {PED_MIN}-{PED_MAX}: {len(ped_all)} images')\n\nassert len(ped_all) > 0, 'No pediatric images found.'\n\nsample_ped = ped_all.sample(n=min(N_SAMPLES, len(ped_all)), random_state=SEED)\n\ndef draw_boxes_on_image(img_bgr, label_path, class_names, colors_fn):\n    img = img_bgr.copy()\n    H, W = img.shape[:2]\n\n    scale = min(W, H) / 800\n    thick = max(8, int(6 * scale))\n    fs = max(0.7, 1.2 * scale)\n    text_thick = max(2, thick - 2)\n\n    if Path(label_path).exists():\n        with open(label_path) as f:\n            for line in f:\n                parts = line.strip().split()\n                if len(parts) != 5:\n                    continue\n\n                nc, xc, yc, bw, bh = int(parts[0]), *map(float, parts[1:])\n                x1 = int((xc - bw/2) * W)\n                y1 = int((yc - bh/2) * H)\n                x2 = int((xc + bw/2) * W)\n                y2 = int((yc + bh/2) * H)\n\n                # ✅ Define box color first\n                box_color = colors_fn(nc)\n\n                # ✅ Now compute contrast text color\n                TEXT_COLOR = get_contrast_color(box_color)\n\n                # Draw bounding box\n                cv2.rectangle(img, (x1, y1), (x2, y2), box_color, thick)\n\n                lbl = class_names[nc] if 0 <= nc < len(class_names) else f'cls{nc}'\n                (tw, th), _ = cv2.getTextSize(lbl, cv2.FONT_HERSHEY_SIMPLEX, fs, text_thick)\n\n                # Background rectangle\n                cv2.rectangle(\n                    img,\n                    (x1, y1 - th - 10),\n                    (x1 + tw + 10, y1),\n                    box_color,\n                    -1\n                )\n\n                # Draw text with contrast color\n                cv2.putText(\n                    img,\n                    lbl,\n                    (x1 + 5, y1 - 5),\n                    cv2.FONT_HERSHEY_SIMPLEX,\n                    fs,\n                    TEXT_COLOR,\n                    text_thick,\n                    cv2.LINE_AA\n                )\n\n    return img\nimg_dir_map   = {'train': IMAGES_TRAIN, 'valid': IMAGES_VALID}\nlbl_dir_map_r = {'train': REMAP_LABELS_TRAIN, 'valid': REMAP_LABELS_VALID}\n\nn = len(sample_ped); n_cols = 4\nn_rows = (n + n_cols - 1) // n_cols\nfig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols*5, n_rows*4))\naxes = np.array(axes).flatten()\nfig.suptitle(f'Pediatric Samples (Age {PED_MIN}-{PED_MAX}) — Filling (blue) & Caries (orange)',\n             fontsize=12, fontweight='bold')\n\nfor idx, (_, row) in enumerate(sample_ped.iterrows()):\n    sp = row['split']\n    img_path = Path(img_dir_map[sp]) / row['filename']\n    lbl_path = Path(lbl_dir_map_r[sp]) / (Path(row['filename']).stem + '.txt')\n    img = cv2.imread(str(img_path))\n    if img is None:\n        axes[idx].set_visible(False); continue\n    annotated = draw_boxes_on_image(img, lbl_path, CLASS_NAMES, get_color)\n    axes[idx].imshow(cv2.cvtColor(annotated, cv2.COLOR_BGR2RGB))\n    axes[idx].set_title(f\"{row['patient_id']} | {row['gender']} | Age {row['age']}\", fontsize=8)\n    axes[idx].axis('off')\n\nfor idx in range(n, len(axes)): axes[idx].set_visible(False)\nplt.tight_layout()\nfig.savefig(PEDIATRIC_DIR / 'pediatric_annotated_samples.png', dpi=100, bbox_inches='tight')\nplt.show(); print(f'Saved: {PEDIATRIC_DIR}/pediatric_annotated_samples.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec4-hdr","cell_type":"markdown","source":"---\n## Section 4 — Imbalance Strategy\n\n### With only 2 classes the imbalance picture changes significantly:\n\n| Class   | Original sources                    | Train boxes |\n|---------|-------------------------------------|-------------|\n| Filling | orig class 1                        | 2,187       |\n| Caries  | orig classes 3+4+5+6+7+8 (merged)  | 736+1025+472+43+978+52 = **3,306** |\n\nAfter merging, **Caries is now the majority class** (~1.5× Filling).  \nThe imbalance is mild (ratio ≈ 0.66) compared to the original 9-class scenario.  \n\n**Strategy chosen:** Mild class-aware augmentation.  \n- Filling images receive a slightly higher sampling weight (≈1.51×) to equalise the two classes.\n- YOLOv8 augmentations (mosaic, HSV, scale, flip) are kept active for both classes.\n- No extreme oversampling needed — the two classes are already near-balanced.","metadata":{}},{"id":"sec4-weights","cell_type":"code","source":"# ── Compute train-only class counts ───────────────────────────────────────\ntrain_box_counts = df_train_boxes['class_id'].value_counts().sort_index()\nfor i in range(NUM_CLASSES):\n    if i not in train_box_counts: train_box_counts[i] = 0\ntrain_box_counts = train_box_counts.sort_index()\n\ntrain_imgs_per_class = df_train_boxes.groupby('class_id')['filename'].nunique()\nfor i in range(NUM_CLASSES):\n    if i not in train_imgs_per_class: train_imgs_per_class[i] = 0\ntrain_imgs_per_class = train_imgs_per_class.sort_index()\n\nprint('=== Train class box counts (2-class) ===')\nmax_boxes = train_box_counts.max()\nfor i in range(NUM_CLASSES):\n    ratio = train_box_counts[i] / max_boxes\n    tag   = ' <- MINORITY' if ratio < 0.9 else ''\n    print(f'  Class {i} ({CLASS_NAMES[i]:8s}): {train_box_counts[i]:5d} boxes | ratio={ratio:.3f}{tag}')\n\n# Inverse-frequency weights\nclass_weights = {i: max_boxes / max(train_box_counts[i], 1) for i in range(NUM_CLASSES)}\nprint(f'\\nClass weights (inverse frequency):')\nfor i, w in class_weights.items():\n    print(f'  Class {i} ({CLASS_NAMES[i]:8s}): {w:.3f}x')\n\n# Per-image weights\ntrain_file_to_classes = {}\nfor fname in df_train_boxes['filename'].unique():\n    train_file_to_classes[fname] = df_train_boxes[\n        df_train_boxes['filename']==fname]['class_id'].tolist()\n\n# Also include images with ZERO remaining boxes\nall_train_fnames = [ip.name for ip in train_imgs]\nfor fname in all_train_fnames:\n    if fname not in train_file_to_classes:\n        train_file_to_classes[fname] = []\n\nimage_weights = {}\nfor fname, classes in train_file_to_classes.items():\n    if classes:\n        image_weights[fname] = max(class_weights.get(c, 1.0) for c in classes)\n    else:\n        image_weights[fname] = 1.0  # hard negative\n\nw_vals = list(image_weights.values())\nprint(f'\\nImage weight stats:')\nprint(f'  Min:  {min(w_vals):.3f}')\nprint(f'  Max:  {max(w_vals):.3f}')\nprint(f'  Mean: {np.mean(w_vals):.3f}')\nprint(f'  Images with weight > 1.0: {sum(1 for v in w_vals if v > 1.0)}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec5-hdr","cell_type":"markdown","source":"---\n## Section 5 — Before/After Class Distribution","metadata":{}},{"id":"sec5-before","cell_type":"code","source":"# ── BEFORE ────────────────────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfig.suptitle('BEFORE Strategy — Train Class Distribution (2-class)',\n             fontsize=13, fontweight='bold')\n\nax = axes[0]\nbars = ax.bar(range(NUM_CLASSES), train_box_counts.values, color=colors_bar, edgecolor='white')\nfor bar, v in zip(bars, train_box_counts.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+10, str(v), ha='center')\nax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(x_labels)\nax.set_ylabel('Box Count'); ax.set_title('Boxes per Class (BEFORE)')\n\nax = axes[1]\nbars = ax.bar(range(NUM_CLASSES), train_imgs_per_class.values, color=colors_bar, edgecolor='white')\nfor bar, v in zip(bars, train_imgs_per_class.values):\n    ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+1, str(v), ha='center')\nax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(x_labels)\nax.set_ylabel('Image Count'); ax.set_title('Images per Class (BEFORE)')\n\nplt.tight_layout()\nfig.savefig(IMBALANCE_DIR / 'before_distribution.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {IMBALANCE_DIR}/before_distribution.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec5-simulate","cell_type":"code","source":"# ── Simulate effective distribution after mild weighted sampling ───────────\nall_train_files = list(train_file_to_classes.keys())\nsw = np.array([image_weights.get(f, 1.0) for f in all_train_files], dtype=float)\nsw /= sw.sum()\n\nK_SIM = 5\nN_DRAWS = K_SIM * len(all_train_files)\nrng = np.random.default_rng(SEED)\nsampled_idx = rng.choice(len(all_train_files), size=N_DRAWS, replace=True, p=sw)\n\neff_counts = Counter()\nfor idx in sampled_idx:\n    for c in train_file_to_classes[all_train_files[idx]]:\n        eff_counts[c] += 1\n\neff_vals = np.array([eff_counts.get(i, 0) for i in range(NUM_CLASSES)])\norig_total = train_box_counts.sum()\neff_norm   = (eff_vals / eff_vals.sum() * orig_total).astype(int)\n\nprint('Simulated effective distribution (normalised to original total):')\nfor i in range(NUM_CLASSES):\n    before = int(train_box_counts[i])\n    after  = int(eff_norm[i])\n    chg    = (after - before) / max(before, 1) * 100\n    print(f'  {CLASS_NAMES[i]:8s}: {before:5d} -> {after:5d} ({chg:+.0f}%)')\n\n# ── BEFORE vs AFTER plot ───────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfig.suptitle('Before vs After — Mild Weighted Sampling (2-class)',\n             fontsize=13, fontweight='bold')\n\nx = np.arange(NUM_CLASSES); width = 0.35\norig_v = np.array([int(train_box_counts[i]) for i in range(NUM_CLASSES)])\n\nax = axes[0]\nax.bar(x-width/2, orig_v,    width, label='Before', color='steelblue',  edgecolor='white')\nax.bar(x+width/2, eff_norm,  width, label='After',  color='darkorange', edgecolor='white')\nax.set_xticks(x); ax.set_xticklabels(x_labels)\nax.set_ylabel('Box Count'); ax.set_title('Box Counts Before/After (linear)')\nax.legend()\n\nax = axes[1]\nax.bar(x-width/2, orig_v+1,   width, label='Before', color='steelblue',  edgecolor='white')\nax.bar(x+width/2, eff_norm+1, width, label='After',  color='darkorange', edgecolor='white')\nax.set_xticks(x); ax.set_xticklabels(x_labels)\nax.set_ylabel('Box Count (log)'); ax.set_title('Box Counts Before/After (log)')\nax.set_yscale('log'); ax.legend()\n\nplt.tight_layout()\nfig.savefig(IMBALANCE_DIR / 'before_after_distribution.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {IMBALANCE_DIR}/before_after_distribution.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec6-hdr","cell_type":"markdown","source":"---\n## Section 6 — YAML Generation & Model Training (ENHANCED)\n\n### Key Changes from Enhancement Guide:\n| Setting | Before | After | Reason |\n|---------|--------|-------|--------|\n| Model | `yolo11s.pt` (detect) | `yolo11m-seg.pt` (segment) | Segmentation outperforms detection for dental caries; medium model improves accuracy |\n| `imgsz` | 640 | 1024 | Tiled images are already 1024px; avoids information loss |\n| `batch` | 4 | 6 | Balanced for yolo11m-seg + 1024px on P100 16GB |\n| `epochs` | 100 | 120 | More epochs for medium model convergence |\n| `lr0` | 0.001 | 0.0005 | Medical imaging standard (lower LR, more stable) |\n| `mosaic` | 1.0 | 0.5 | Realistic textures need less synthetic mixing |\n| `mixup` | 0.1 | 0.0 | Caries texture is subtle; mixup corrupts features |\n| `copy_paste` | 0.05 | 0.0 | Synthetic pasting harms lesion feature learning |\n| `erasing` | default | 0.1 | Minimal; helps robustness without over-augmenting |","metadata":{}},{"id":"sec6-yaml","cell_type":"code","source":"# ── Write the training YAML pointing to tiled/preprocessed images ─────────\n# YAML uses tiled images if USE_TILING=True, else CLAHE or original images.\n\nYAML_PATH = Path('alphadent_2class_enhanced.yaml')\n\nyaml_content = f\"\"\"# AlphaDent 2-Class Enhanced — Segmentation Training YAML\n# Auto-generated. Do not edit manually.\n\npath: /kaggle/working\n\ntrain: {FINAL_TRAIN_IMAGES}\nval:   {FINAL_VALID_IMAGES}\n\nnc: {NUM_CLASSES}\nnames: {CLASS_NAMES}\n\"\"\"\n\nYAML_PATH.write_text(yaml_content)\nprint(f'YAML written to: {YAML_PATH}')\nprint(yaml_content)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec6-config","cell_type":"code","source":"DEVICE = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\n# ── ENHANCED Training Configuration ────────────────────────────────────────\n# Changes from guide:\n#   model    : yolo11s.pt -> yolo11m-seg.pt  (segmentation + medium model)\n#   imgsz    : 640 -> 1024                   (1024px tiles, no info loss)\n#   batch    : 4 -> 6                        (yolo11m + 1024px on P100 16GB)\n#   epochs   : 100 -> 120                    (more epochs for convergence)\n#   lr0      : 0.001 -> 0.0005               (medical imaging standard LR)\n#   mosaic   : 1.0 -> 0.5                    (less synthetic mixing)\n#   mixup    : 0.1 -> 0.0                    (disabled: harms subtle textures)\n#   copy_paste: 0.05 -> 0.0                  (disabled: harms lesion features)\n#   erasing  : added 0.1                     (minimal random erasing)\n\nTRAIN_CONFIG = dict(\n    # ── Model & data ───────────────────────────────────────────────────────\n    model       = 'yolo11m-seg.pt',    # 🔥 UPGRADE: medium + segmentation\n    data        = str(YAML_PATH),\n    task        = 'segment',           # Explicit segmentation task\n    imgsz       = 1024,               # 🔥 UPGRADE: 640->1024 (tile size)\n    batch       = 6,                  # 🔥 ADJUSTED for yolo11m + 1024px\n    epochs      = 120,                # 🔥 EXTENDED: 100->120\n    patience    = 25,                 # Slightly more patience\n    device      = DEVICE,\n\n    # ── Optimizer ──────────────────────────────────────────────────────────\n    optimizer   = 'AdamW',\n    lr0         = 0.0005,             # 🔥 LOWERED: 0.001->0.0005 (medical)\n    lrf         = 0.01,\n    momentum    = 0.937,\n    weight_decay= 5e-4,\n    cos_lr      = True,\n\n    # ── Augmentation (tuned for medical imaging) ───────────────────────────\n    mosaic      = 0.5,    # 🔥 REDUCED: 1.0->0.5 (less synthetic mixing)\n    mixup       = 0.0,    # 🔥 DISABLED: harms subtle caries texture learning\n    copy_paste  = 0.0,    # 🔥 DISABLED: synthetic pasting harms features\n    erasing     = 0.1,    # 🔥 ADDED: minimal random erasing for robustness\n\n    # ── Standard spatial augmentations ────────────────────────────────────\n    hsv_h       = 0.015,\n    hsv_s       = 0.7,\n    hsv_v       = 0.4,\n    degrees     = 5.0,\n    translate   = 0.1,\n    scale       = 0.5,\n    shear       = 2.0,\n    flipud      = 0.1,\n    fliplr      = 0.5,\n\n    # ── System ────────────────────────────────────────────────────────────\n    amp         = True,\n    workers     = 4,\n    seed        = SEED,\n    project     = 'runs/train',\n    name        = 'alphadent_2class_enhanced',\n    exist_ok    = True,\n    save        = True,\n    save_period = 10,\n    verbose     = True,\n)\n\nprint('Enhanced Training config:')\nprint(f'  {\"Setting\":<15s}  {\"Value\":<20s}  Note')\nprint(f'  {\"-\"*55}')\nhighlights = {\n    \"model\": \"🔥 yolo11m-seg (upgraded)\",\n    \"imgsz\": \"🔥 1024px (was 640)\",\n    \"batch\": \"🔥 6 (optimized)\",\n    \"epochs\": \"🔥 120 (was 100)\",\n    \"lr0\": \"🔥 0.0005 (was 0.001)\",\n    \"mosaic\": \"🔥 0.5 (was 1.0)\",\n    \"mixup\": \"🔥 0.0 (disabled)\",\n    \"copy_paste\": \"🔥 0.0 (disabled)\",\n    \"erasing\": \"🔥 0.1 (added)\",\n}\nfor k, v in TRAIN_CONFIG.items():\n    note = highlights.get(k, '')\n    print(f'  {k:<15s}: {str(v):<20s} {note}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec6-train","cell_type":"code","source":"print('Loading yolo11m-seg (medium segmentation model)...')\nprint('NOTE: yolo11m-seg.pt will be auto-downloaded by Ultralytics on first run.')\nmodel = YOLO('yolo11m-seg.pt')\n\nprint('\\nStarting ENHANCED training...')\nprint(f'  Model   : yolo11m-seg.pt (segmentation)')\nprint(f'  imgsz   : {TRAIN_CONFIG[\"imgsz\"]}px')\nprint(f'  epochs  : {TRAIN_CONFIG[\"epochs\"]}')\nprint(f'  lr0     : {TRAIN_CONFIG[\"lr0\"]}')\nprint(f'  mixup   : {TRAIN_CONFIG[\"mixup\"]} (disabled for medical images)')\nprint()\n\ntry:\n    results = model.train(**TRAIN_CONFIG)\n    print('Training complete.')\nexcept RuntimeError as e:\n    if 'out of memory' in str(e).lower():\n        print('OOM — retrying with batch=4, imgsz=1024')\n        torch.cuda.empty_cache()\n        TRAIN_CONFIG.update(batch=4, imgsz=1024)\n        model = YOLO('yolo11m-seg.pt')\n        results = model.train(**TRAIN_CONFIG)\n    else:\n        raise","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec6-ckpt","cell_type":"code","source":"TRAIN_RUN_DIR = Path(TRAIN_CONFIG['project']) / TRAIN_CONFIG['name']\nBEST_CKPT     = TRAIN_RUN_DIR / 'weights' / 'best.pt'\nLAST_CKPT     = TRAIN_RUN_DIR / 'weights' / 'last.pt'\nprint(f'Training run dir : {TRAIN_RUN_DIR}')\nprint(f'Best checkpoint  : {BEST_CKPT}')\nprint(f'Last checkpoint  : {LAST_CKPT}')\nprint(f'Exists: best={BEST_CKPT.exists()}, last={LAST_CKPT.exists()}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"69f9d7c4-8ea1-4f6c-ba63-c98fba0507bc","cell_type":"code","source":"TRAIN_RUN_DIR","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec6-curves","cell_type":"code","source":"results_csv = str(TRAIN_RUN_DIR / 'results.csv')\n\ndf_log = pd.read_csv(results_csv)\ndf_log.columns = [c.strip() for c in df_log.columns]\nepochs = df_log.get('epoch', pd.Series(range(1, len(df_log)+1)))\n\ndef find_col(df, *kws):\n    for kw in kws:\n        m = [c for c in df.columns if kw.lower() in c.lower()]\n        if m: return m[0]\n    return None\n\ncol = dict(\n    tbox = find_col(df_log,'train/box_loss','box_loss'),\n    tcls = find_col(df_log,'train/cls_loss','cls_loss'),\n    tdfl = find_col(df_log,'train/dfl_loss','dfl_loss'),\n    tseg = find_col(df_log,'train/seg_loss'),\n    vbox = find_col(df_log,'val/box_loss'),\n    vcls = find_col(df_log,'val/cls_loss'),\n    vdfl = find_col(df_log,'val/dfl_loss'),\n    vseg = find_col(df_log,'val/seg_loss'),\n    prec = find_col(df_log,'metrics/precision','precision'),\n    rec  = find_col(df_log,'metrics/recall','recall'),\n    m50  = find_col(df_log,'metrics/mAP50(M)','metrics/mAP50(B)','mAP50'),\n    m595 = find_col(df_log,'metrics/mAP50-95(M)','metrics/mAP50-95(B)','mAP50-95'),\n)\n\nfig, axes = plt.subplots(2, 4, figsize=(24, 10))\nfig.suptitle('AlphaDent 2-Class ENHANCED Training Curves (yolo11m-seg, 1024px)', fontsize=14, fontweight='bold')\n\ndef sp(ax, c, color, label, title, ylabel='Loss'):\n    if c and c in df_log.columns:\n        ax.plot(epochs, pd.to_numeric(df_log[c], errors='coerce'),\n                color=color, lw=2, label=label)\n        ax.set_xlabel('Epoch'); ax.set_ylabel(ylabel)\n        ax.set_title(title); ax.legend(); ax.grid(True, alpha=0.3)\n\nsp(axes[0,0], col['tbox'], 'blue',  'Train', 'Box Loss')\nsp(axes[0,0], col['vbox'], 'red',   'Val',   'Box Loss')\nsp(axes[0,1], col['tcls'], 'blue',  'Train', 'Cls Loss')\nsp(axes[0,1], col['vcls'], 'red',   'Val',   'Cls Loss')\nsp(axes[0,2], col['tdfl'], 'blue',  'Train', 'DFL Loss')\nsp(axes[0,2], col['vdfl'], 'red',   'Val',   'DFL Loss')\nsp(axes[0,3], col['tseg'], 'blue',  'Train', 'Seg Loss')\nsp(axes[0,3], col['vseg'], 'red',   'Val',   'Seg Loss')\n\nif col['m50'] and col['m50'] in df_log.columns:\n    vals = pd.to_numeric(df_log[col['m50']], errors='coerce')\n    axes[1,0].plot(epochs, vals, 'g-', lw=2, label='mAP50')\n    best_ep = epochs.iloc[vals.idxmax()]\n    axes[1,0].axvline(best_ep, color='green', linestyle=':', label=f'Best@{best_ep}')\n    axes[1,0].set_xlabel('Epoch'); axes[1,0].set_ylabel('mAP50')\n    axes[1,0].set_title('mAP50 (Segmentation)'); axes[1,0].legend(); axes[1,0].grid(True, alpha=0.3)\n\nif col['m595'] and col['m595'] in df_log.columns:\n    axes[1,1].plot(epochs, pd.to_numeric(df_log[col['m595']], errors='coerce'),\n                   'm-', lw=2, label='mAP50-95')\n    axes[1,1].set_xlabel('Epoch'); axes[1,1].set_ylabel('mAP50-95')\n    axes[1,1].set_title('mAP50-95 (COCO, Segmentation)'); axes[1,1].legend(); axes[1,1].grid(True, alpha=0.3)\n\nif col['prec'] and col['rec']:\n    axes[1,2].plot(epochs, pd.to_numeric(df_log[col['prec']], errors='coerce'),\n                   'b-', lw=2, label='Precision')\n    axes[1,2].plot(epochs, pd.to_numeric(df_log[col['rec']], errors='coerce'),\n                   'r-', lw=2, label='Recall')\n    axes[1,2].set_xlabel('Epoch'); axes[1,2].set_ylabel('Score')\n    axes[1,2].set_title('Precision & Recall'); axes[1,2].legend(); axes[1,2].grid(True, alpha=0.3)\n\n# Summary text box\naxes[1,3].axis('off')\nsummary_text = (\n    'Enhancement Summary\\n'\n    '─────────────────\\n'\n    f'Model   : yolo11m-seg.pt\\n'\n    f'imgsz   : {TRAIN_CONFIG[\"imgsz\"]}px\\n'\n    f'epochs  : {TRAIN_CONFIG[\"epochs\"]}\\n'\n    f'lr0     : {TRAIN_CONFIG[\"lr0\"]}\\n'\n    f'mosaic  : {TRAIN_CONFIG[\"mosaic\"]}\\n'\n    f'mixup   : {TRAIN_CONFIG[\"mixup\"]} (off)\\n'\n    f'cp_paste: {TRAIN_CONFIG[\"copy_paste\"]} (off)\\n'\n    f'CLAHE   : {USE_CLAHE}\\n'\n    f'Tiling  : {USE_TILING}\\n'\n    '─────────────────\\n'\n    'Expected mAP50 ≥ 0.65'\n)\naxes[1,3].text(0.05, 0.95, summary_text, transform=axes[1,3].transAxes,\n               fontsize=10, va='top', fontfamily='monospace',\n               bbox=dict(boxstyle='round', facecolor='lightyellow', alpha=0.8))\n\nplt.tight_layout()\nfig.savefig(CURVES_DIR / 'training_curves_enhanced.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {CURVES_DIR}/training_curves_enhanced.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec7-hdr","cell_type":"markdown","source":"---\n## Section 7 — Evaluation Metrics","metadata":{}},{"id":"sec7-val","cell_type":"code","source":"print('Loading best checkpoint...')\nbest_model = YOLO(str(BEST_CKPT))\n\nprint('Running validation (segmentation model)...')\nval_results = best_model.val(\n    data     = str(YAML_PATH),\n    imgsz    = TRAIN_CONFIG['imgsz'],\n    batch    = TRAIN_CONFIG['batch'],\n    device   = DEVICE,\n    conf     = 0.001,\n    iou      = 0.60,\n    verbose  = True,\n    task     = 'segment',\n    project  = str(EVAL_DIR),\n    name     = 'val_eval',\n    exist_ok = True,\n)\n\n# Segmentation models expose both box and mask metrics\n# Prefer mask (segmentation) metrics; fall back to box if unavailable\ntry:\n    map50    = float(val_results.seg.map50)\n    map50_95 = float(val_results.seg.map)\n    prec     = float(val_results.seg.mp)\n    rec      = float(val_results.seg.mr)\n    metrics_source = 'seg (mask)'\nexcept Exception:\n    map50    = float(val_results.box.map50)\n    map50_95 = float(val_results.box.map)\n    prec     = float(val_results.box.mp)\n    rec      = float(val_results.box.mr)\n    metrics_source = 'box (fallback)'\n\nf1 = 2*prec*rec / max(prec+rec, 1e-9)\n\nprint(f'\\n=== Overall Validation Metrics (2-class, {metrics_source}) ===')\nprint(f'  mAP50      : {map50:.4f}')\nprint(f'  mAP50-95   : {map50_95:.4f}')\nprint(f'  Precision  : {prec:.4f}')\nprint(f'  Recall     : {rec:.4f}')\nprint(f'  F1         : {f1:.4f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec7-per-class","cell_type":"code","source":"# Extract per-class metrics (segmentation model)\ntry:\n    per_class_ap50 = val_results.seg.ap50\n    per_class_ap   = val_results.seg.ap\n    per_class_p    = val_results.seg.p\n    per_class_r    = val_results.seg.r\n    metrics_label  = 'Segmentation Mask'\nexcept Exception:\n    per_class_ap50 = val_results.box.ap50\n    per_class_ap   = val_results.box.ap\n    per_class_p    = val_results.box.p\n    per_class_r    = val_results.box.r\n    metrics_label  = 'Bounding Box'\n\nrows_pc = []\nfor ci in range(NUM_CLASSES):\n    ap50_ = float(per_class_ap50[ci]) if ci < len(per_class_ap50) else float('nan')\n    ap_   = float(per_class_ap[ci])   if ci < len(per_class_ap)   else float('nan')\n    p_    = float(per_class_p[ci])    if ci < len(per_class_p)    else float('nan')\n    r_    = float(per_class_r[ci])    if ci < len(per_class_r)    else float('nan')\n    f1_   = 2*p_*r_ / max(p_+r_, 1e-9)\n    rows_pc.append(dict(class_id=ci, class_name=CLASS_NAMES[ci],\n                        AP50=ap50_, AP50_95=ap_, Precision=p_, Recall=r_, F1=f1_))\n\ndf_pc = pd.DataFrame(rows_pc)\nprint(f'Per-Class Metrics ({metrics_label}):')\nprint(df_pc.to_string(index=False, float_format='{:.4f}'.format))\n\n# Bar chart\nfig, axes = plt.subplots(1, 3, figsize=(15, 4))\nfig.suptitle(f'Per-Class Metrics — AlphaDent 2-Class Enhanced ({metrics_label})',\n             fontsize=12, fontweight='bold')\nfor ax, col_name in zip(axes, ['AP50', 'AP50_95', 'F1']):\n    colors = [get_color(ci) for ci in range(NUM_CLASSES)]\n    colors_rgb = [(r/255, g/255, b/255) for b, g, r in colors]\n    bars = ax.bar(df_pc['class_name'], df_pc[col_name], color=colors_rgb, width=0.5)\n    ax.set_ylim(0, 1.05)\n    ax.set_ylabel(col_name); ax.set_title(col_name)\n    for bar, val in zip(bars, df_pc[col_name]):\n        ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+0.01,\n                f'{val:.3f}', ha='center', va='bottom', fontsize=11, fontweight='bold')\n    ax.grid(axis='y', alpha=0.3)\nplt.tight_layout()\nfig.savefig(EVAL_DIR / 'per_class_ap_enhanced.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {EVAL_DIR}/per_class_ap_enhanced.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec7-copy-plots","cell_type":"code","source":"import shutil\n\n# Copy auto-generated Ultralytics plots to eval dir\nfor fname in ['confusion_matrix.png', 'PR_curve.png',\n              'F1_curve.png', 'P_curve.png', 'R_curve.png']:\n    for search in [EVAL_DIR / 'val_eval', TRAIN_RUN_DIR]:\n        src = search / fname\n        if src.exists():\n            shutil.copy(src, EVAL_DIR / fname)\n            print(f'Copied {fname}')\n            img = cv2.cvtColor(cv2.imread(str(src)), cv2.COLOR_BGR2RGB)\n            plt.figure(figsize=(10, 8))\n            plt.imshow(img); plt.axis('off')\n            plt.title(fname.replace('_', ' ').replace('.png', ''))\n            plt.tight_layout()\n            plt.savefig(EVAL_DIR / f'{fname.replace(\".png\",\"_display.png\")}',\n                        dpi=100, bbox_inches='tight')\n            plt.show()\n            break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec7-iou-conf","cell_type":"code","source":"# ── IoU and confidence histograms ─────────────────────────────────────────\nPRED_CONF = 0.25\nIOU_THRESH = 0.50\n\nprint('Running inference on validation set...')\nval_preds = best_model.predict(\n    source  = str(IMAGES_VALID),\n    imgsz   = TRAIN_CONFIG['imgsz'],\n    conf    = PRED_CONF,\n    iou     = 0.60,\n    device  = DEVICE,\n    verbose = False,\n    save    = False,\n)\nprint(f'Predictions: {len(val_preds)} images')\n\ndef box_iou(b1, b2):\n    ix1 = max(b1[0],b2[0]); iy1 = max(b1[1],b2[1])\n    ix2 = min(b1[2],b2[2]); iy2 = min(b1[3],b2[3])\n    inter = max(0,ix2-ix1)*max(0,iy2-iy1)\n    a1 = (b1[2]-b1[0])*(b1[3]-b1[1])\n    a2 = (b2[2]-b2[0])*(b2[3]-b2[1])\n    return inter / (a1+a2-inter+1e-9)\n\nall_ious  = []; tp_confs = []; fp_confs = []\nfp_examples = []; fn_examples = []\n\nfor pred in val_preds:\n    ip  = Path(pred.path)\n    lp  = REMAP_LABELS_VALID / (ip.stem + '.txt')\n    H, W = pred.orig_shape\n    gt_abs = []\n    if lp.exists():\n        with open(lp) as f:\n            for line in f:\n                parts = line.strip().split()\n                if len(parts)!=5: continue\n                nc, xc, yc, bw, bh = int(parts[0]), *map(float, parts[1:])\n                gt_abs.append((nc, (xc-bw/2)*W, (yc-bh/2)*H,\n                                   (xc+bw/2)*W, (yc+bh/2)*H))\n\n    if pred.boxes is None or len(pred.boxes)==0:\n        for g in gt_abs:\n            fn_examples.append({'path':str(ip),'cls':g[0],'box':g[1:]})\n        continue\n\n    pboxes = pred.boxes.xyxy.cpu().numpy()\n    pconfs = pred.boxes.conf.cpu().numpy()\n    pcls   = pred.boxes.cls.cpu().numpy().astype(int)\n    matched_gt = [False]*len(gt_abs)\n\n    for pb, pc, conf in zip(pboxes, pcls, pconfs):\n        best_iou = 0; best_gi = -1\n        for gi, (gc,gx1,gy1,gx2,gy2) in enumerate(gt_abs):\n            if gc==pc and not matched_gt[gi]:\n                iou = box_iou(pb, (gx1,gy1,gx2,gy2))\n                if iou > best_iou: best_iou=iou; best_gi=gi\n        all_ious.append(best_iou)\n        if best_iou >= IOU_THRESH and best_gi>=0:\n            matched_gt[best_gi]=True; tp_confs.append(float(conf))\n        else:\n            fp_confs.append(float(conf))\n            fp_examples.append({'path':str(ip),'cls':int(pc),'box':tuple(pb),'conf':float(conf)})\n\n    for gi,(gc,gx1,gy1,gx2,gy2) in enumerate(gt_abs):\n        if not matched_gt[gi]:\n            fn_examples.append({'path':str(ip),'cls':gc,'box':(gx1,gy1,gx2,gy2)})\n\nprint(f'Predictions: {len(all_ious)} | TP: {len(tp_confs)} | FP: {len(fp_confs)}')\nprint(f'FP examples: {len(fp_examples)} | FN examples: {len(fn_examples)}')\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfig.suptitle('Prediction Quality — 2-Class Model', fontsize=13, fontweight='bold')\n\nax = axes[0]\nax.hist(all_ious, bins=50, color='teal', edgecolor='white', alpha=0.8)\nax.axvline(IOU_THRESH, color='red', linestyle='--', label=f'IoU={IOU_THRESH}')\nax.set_xlabel('IoU (pred vs best GT match, same class)')\nax.set_ylabel('Count'); ax.set_title('IoU Histogram'); ax.legend()\n\nax = axes[1]\nax.hist(tp_confs, bins=30, color='green',  alpha=0.7, label='TP', edgecolor='white')\nax.hist(fp_confs, bins=30, color='red',    alpha=0.7, label='FP', edgecolor='white')\nax.set_xlabel('Confidence'); ax.set_ylabel('Count')\nax.set_title('Confidence: TP vs FP'); ax.legend()\n\nplt.tight_layout()\nfig.savefig(EVAL_DIR / 'iou_confidence_histograms.png', dpi=120, bbox_inches='tight')\nplt.show(); print(f'Saved: {EVAL_DIR}/iou_confidence_histograms.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec7-fpfn","cell_type":"code","source":"def viz_errors(examples, label, color_bgr, n=6, save_path=None):\n    n = min(n, len(examples))\n    if n == 0:\n        print(f'No {label} examples.')\n        return\n\n    picked = random.sample(examples, n)\n\n    nc_ = min(3, n)\n    nr_ = (n + nc_ - 1) // nc_\n\n    fig, axes = plt.subplots(nr_, nc_, figsize=(nc_ * 6, nr_ * 5))\n    axes = np.array(axes).flatten()\n\n    fig.suptitle(f'{label} Examples (conf >= {PRED_CONF})',\n                 fontsize=18, fontweight='bold')\n\n    for i, ex in enumerate(picked):\n        img = cv2.imread(ex['path'])\n        if img is None:\n            axes[i].set_visible(False)\n            continue\n\n        h, w = img.shape[:2]\n        scale_factor = max(w, h) / 1000  # reference size = 1000px\n\n        # Dynamic scaling\n        box_thickness = max(2, int(3 * scale_factor))\n        font_scale = max(0.7, 0.8 * scale_factor)\n        text_thickness = max(1, int(2 * scale_factor))\n        padding = int(6 * scale_factor)\n\n        x1, y1, x2, y2 = [int(v) for v in ex['box']]\n\n        # Draw bounding box\n        cv2.rectangle(img, (x1, y1), (x2, y2),\n                      color_bgr, thickness=box_thickness)\n\n        # Prepare label\n        lbl_ = CLASS_NAMES[ex['cls']] if 0 <= ex['cls'] < NUM_CLASSES else '?'\n        cs = f\" {ex['conf']:.2f}\" if 'conf' in ex else ''\n        text = f\"{label}: {lbl_}{cs}\"\n\n        font = cv2.FONT_HERSHEY_SIMPLEX\n        (tw, th), baseline = cv2.getTextSize(\n            text, font, font_scale, text_thickness)\n\n        # Background rectangle\n        cv2.rectangle(\n            img,\n            (x1, y1 - th - baseline - padding),\n            (x1 + tw + padding, y1),\n            color_bgr,\n            -1\n        )\n\n        # White text\n        cv2.putText(\n            img,\n            text,\n            (x1 + padding // 2, y1 - padding // 2),\n            font,\n            font_scale,\n            (255, 255, 255),\n            text_thickness,\n            lineType=cv2.LINE_AA\n        )\n\n        axes[i].imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))\n        axes[i].set_title(Path(ex['path']).name[:30],\n                          fontsize=11,\n                          fontweight='bold')\n        axes[i].axis('off')\n\n    for i in range(n, len(axes)):\n        axes[i].set_visible(False)\n\n    plt.tight_layout()\n\n    if save_path:\n        fig.savefig(save_path, dpi=250, bbox_inches='tight')\n\n    plt.show()\n\n    if save_path:\n        print(f'Saved: {save_path}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"47542019-91f8-4e76-b9c1-9e85036d7875","cell_type":"code","source":"viz_errors(fp_examples, 'FP', (0,0,255),    n=6, save_path=EVAL_DIR/'fp_examples.png')\nviz_errors(fn_examples, 'FN', (0,165,255),  n=6, save_path=EVAL_DIR/'fn_examples.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec8-hdr","cell_type":"markdown","source":"---\n## Section 8 — Predictions & Visualization","metadata":{}},{"id":"sec8-valid","cell_type":"code","source":"# ── 8.1 Validation: predicted vs GT (colors matched correctly) ───────────\nimport random\nfrom pathlib import Path\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch\n\n# ====== SET YOUR CLASS IDS HERE ======\nFILLING_ID = 1\nCARIES_ID  = 3\n\n# ====== ONE SOURCE OF TRUTH FOR COLORS (BGR for OpenCV) ======\nGT_BGR      = (0, 255, 0)        # green\nFILLING_BGR = (255, 0, 0)        # blue\nCARIES_BGR  = (0, 165, 255)      # orange\n\ndef bgr_to_rgb01(bgr):\n    b, g, r = bgr\n    return (r/255, g/255, b/255)\n\nN_VIZ = 3\nVIZ_MAX_SIDE = 1400\nviz_preds = random.sample(val_preds, min(N_VIZ, len(val_preds)))\n\ndef _viz_scale_params(h, w):\n    sf = max(h, w) / 1000.0\n    box_th = max(2, int(round(3 * sf)))\n    gt_th  = max(2, int(round(2.5 * sf)))\n    font   = max(0.6, 0.8 * sf)\n    txt_th = max(1, int(round(2 * sf)))\n    pad    = max(4, int(round(6 * sf)))\n    return box_th, gt_th, font, txt_th, pad\n\ndef _draw_label(img, x1, y1, text, color_bgr, font_scale, text_thickness, pad):\n    font = cv2.FONT_HERSHEY_SIMPLEX\n    (tw, th), baseline = cv2.getTextSize(text, font, font_scale, text_thickness)\n\n    y_top = y1 - th - baseline - pad\n    if y_top < 0:\n        y_top = y1 + pad  # simple: put just below top if above is not possible\n\n    x2 = min(img.shape[1] - 1, x1 + tw + pad)\n    y2 = min(img.shape[0] - 1, y_top + th + baseline + pad)\n\n    cv2.rectangle(img, (x1, y_top), (x2, y2), color_bgr, -1)\n    cv2.putText(img, text, (x1 + pad // 2, y2 - baseline - pad // 2),\n                font, font_scale, (255, 255, 255), text_thickness, lineType=cv2.LINE_AA)\n\ndef cls_name(i):\n    return CLASS_NAMES[i] if 0 <= i < len(CLASS_NAMES) else f\"class_{i}\"\n\ndef _resize_for_viz(img, max_side):\n    h, w = img.shape[:2]\n    m = max(h, w)\n    if m <= max_side:\n        return img, 1.0\n    r = max_side / float(m)\n    new_w, new_h = int(round(w * r)), int(round(h * r))\n    return cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_AREA), r\n\nn_cols = 3\nn_rows = (len(viz_preds) + n_cols - 1) // n_cols\nfig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols * 6.5, n_rows * 5.5))\naxes = np.array(axes).flatten()\n\nfig.suptitle(\n    'Validation Predictions vs GT — Filling (blue) & Caries (orange)\\n'\n    'Solid colour = predicted | Green outline = GT',\n    fontsize=14, fontweight='bold'\n)\n\nfor idx, pred in enumerate(viz_preds):\n    ip = Path(pred.path)\n    lp = REMAP_LABELS_VALID / (ip.stem + '.txt')\n\n    img0 = cv2.imread(str(ip))\n    if img0 is None:\n        axes[idx].set_visible(False)\n        continue\n\n    H0, W0 = img0.shape[:2]\n\n    img, r = _resize_for_viz(img0, VIZ_MAX_SIDE)\n    H, W = img.shape[:2]\n    draw = img.copy()\n\n    box_th, gt_th, font_scale, txt_th, pad = _viz_scale_params(H, W)\n\n    # ---- GT (green) ----\n    if lp.exists():\n        with open(lp) as f:\n            for line in f:\n                parts = line.strip().split()\n                if len(parts) != 5:\n                    continue\n                _, xc, yc, bw, bh = int(parts[0]), *map(float, parts[1:])\n\n                x1o = int((xc - bw / 2) * W0); y1o = int((yc - bh / 2) * H0)\n                x2o = int((xc + bw / 2) * W0); y2o = int((yc + bh / 2) * H0)\n\n                x1 = int(round(x1o * r)); y1 = int(round(y1o * r))\n                x2 = int(round(x2o * r)); y2 = int(round(y2o * r))\n\n                cv2.rectangle(draw, (x1, y1), (x2, y2), GT_BGR, gt_th)\n\n    # ---- Predictions (only Filling + Caries to match legend) ----\n    if pred.boxes is not None and len(pred.boxes) > 0:\n        for pb, pc, conf in zip(\n            pred.boxes.xyxy.cpu().numpy(),\n            pred.boxes.cls.cpu().numpy().astype(int),\n            pred.boxes.conf.cpu().numpy()\n        ):\n            if pc not in (FILLING_ID, CARIES_ID):\n                continue\n\n            x1o, y1o, x2o, y2o = [int(v) for v in pb]\n            x1 = int(round(x1o * r)); y1 = int(round(y1o * r))\n            x2 = int(round(x2o * r)); y2 = int(round(y2o * r))\n\n            color = FILLING_BGR if pc == FILLING_ID else CARIES_BGR\n            cv2.rectangle(draw, (x1, y1), (x2, y2), color, box_th)\n\n            name = CLASS_NAMES[pc] if 0 <= pc < NUM_CLASSES else str(pc)\n            _draw_label(draw, x1, y1, f\"{name}: {conf:.2f}\", color, font_scale, txt_th, pad)\n\n    axes[idx].imshow(cv2.cvtColor(draw, cv2.COLOR_BGR2RGB))\n    axes[idx].set_title(ip.name[:35], fontsize=10, fontweight='bold')\n    axes[idx].axis('off')\n\n# ---- Legend (EXACT same colors as boxes) ----\nfig.legend(\n    handles=[\n        Patch(facecolor=bgr_to_rgb01(FILLING_BGR), label=f'Pred: {cls_name(FILLING_ID)}'),\n        Patch(facecolor=bgr_to_rgb01(CARIES_BGR),  label=f'Pred: {cls_name(CARIES_ID)}'),\n    ],\n    loc='upper right', fontsize=10\n)\n\nfor idx in range(len(viz_preds), len(axes)):\n    axes[idx].set_visible(False)\n\nplt.tight_layout()\nout_path = PREDS_VALID_DIR / 'val_predictions_vs_gt.png'\nfig.savefig(out_path, dpi=250, bbox_inches='tight')\nplt.show()\nprint(f'Saved: {out_path}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec8-test","cell_type":"code","source":"# ── 8.2 Test set predictions ──────────────────────────────────────────────\nprint('Predicting on test set...')\ntest_preds = best_model.predict(\n    source  = str(IMAGES_TEST),\n    imgsz   = TRAIN_CONFIG['imgsz'],\n    conf    = PRED_CONF,\n    iou     = 0.60,\n    device  = DEVICE,\n    verbose = False,\n    save    = False,\n)\nprint(f'Test predictions: {len(test_preds)} images')\n\ndef max_conf(p):\n    return float(p.boxes.conf.max()) if p.boxes is not None and len(p.boxes)>0 else 0.0\n\ntop6   = sorted(test_preds, key=max_conf, reverse=True)[:6]\nrand6  = random.sample(test_preds, min(6, len(test_preds)))\nviz_tp = list({id(p): p for p in top6+rand6}.values())[:12]\n\nn_cols=3; n_rows=(len(viz_tp)+n_cols-1)//n_cols\nfig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols*6, n_rows*5))\naxes = np.array(axes).flatten()\nfig.suptitle('Test Set Predictions (No GT) — Filling & Caries', fontsize=12, fontweight='bold')\n\nfor idx, pred in enumerate(viz_tp):\n    ip  = Path(pred.path)\n    img = cv2.imread(str(ip))\n    if img is None: axes[idx].set_visible(False); continue\n    draw = img.copy()\n    n_det = 0\n    if pred.boxes is not None and len(pred.boxes)>0:\n        n_det = len(pred.boxes)\n        for pb, pc, conf in zip(pred.boxes.xyxy.cpu().numpy(),\n                                 pred.boxes.cls.cpu().numpy().astype(int),\n                                 pred.boxes.conf.cpu().numpy()):\n            x1,y1,x2,y2 = [int(v) for v in pb]\n            color = get_color(int(pc))\n            cv2.rectangle(draw,(x1,y1),(x2,y2),color,3)\n            lbl_ = f\"{CLASS_NAMES[pc] if 0<=pc<NUM_CLASSES else pc}:{conf:.2f}\"\n            cv2.putText(draw,lbl_,(x1,max(y1-5,15)),cv2.FONT_HERSHEY_SIMPLEX,0.5,color,2)\n    axes[idx].imshow(cv2.cvtColor(draw,cv2.COLOR_BGR2RGB))\n    axes[idx].set_title(f'{ip.name[:25]} | {n_det} dets',fontsize=7)\n    axes[idx].axis('off')\n\nfor idx in range(len(viz_tp),len(axes)): axes[idx].set_visible(False)\nplt.tight_layout()\nfig.savefig(PREDS_TEST_DIR/'test_predictions.png', dpi=100, bbox_inches='tight')\nplt.show(); print(f'Saved: {PREDS_TEST_DIR}/test_predictions.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec9-hdr","cell_type":"markdown","source":"---\n## Section 9 — Save Outputs","metadata":{}},{"id":"sec9-save","cell_type":"code","source":"# ── results.csv ───────────────────────────────────────────────────────────\nsummary = dict(\n    model               = TRAIN_CONFIG['model'],\n    num_classes         = NUM_CLASSES,\n    classes             = str(CLASS_NAMES),\n    dropped_classes     = 'Abrasion (orig 0), Crown (orig 2)',\n    caries_merge        = 'orig 3,4,5,6,7,8 -> class 1',\n    imgsz               = TRAIN_CONFIG['imgsz'],\n    batch               = TRAIN_CONFIG['batch'],\n    epochs_config       = TRAIN_CONFIG['epochs'],\n    lr0                 = TRAIN_CONFIG['lr0'],\n    mosaic              = TRAIN_CONFIG['mosaic'],\n    mixup               = TRAIN_CONFIG['mixup'],\n    copy_paste          = TRAIN_CONFIG['copy_paste'],\n    clahe_applied       = USE_CLAHE,\n    tiling_applied      = USE_TILING,\n    tile_size           = TILE_SIZE if USE_TILING else 'N/A',\n    strategy            = 'mild_class_aware_augmentation + CLAHE + tiling',\n    mAP50               = round(map50, 4),\n    mAP50_95            = round(map50_95, 4),\n    precision           = round(prec, 4),\n    recall              = round(rec, 4),\n    f1                  = round(f1, 4),\n    best_checkpoint     = str(BEST_CKPT),\n)\npd.DataFrame([summary]).to_csv(RUNS_ROOT/'results.csv', index=False)\nprint(f'Saved: {RUNS_ROOT}/results.csv')\n\n# ── test predictions CSV ──────────────────────────────────────────────────\nrows = []\nfor pred in test_preds:\n    ip = Path(pred.path)\n    if pred.boxes is None or len(pred.boxes)==0:\n        rows.append(dict(image_id=ip.stem, filename=ip.name,\n                         box_x1=None,box_y1=None,box_x2=None,box_y2=None,\n                         confidence=None,class_id=None,class_name=None))\n        continue\n    for pb, pc, conf in zip(pred.boxes.xyxy.cpu().numpy(),\n                             pred.boxes.cls.cpu().numpy().astype(int),\n                             pred.boxes.conf.cpu().numpy()):\n        rows.append(dict(\n            image_id=ip.stem, filename=ip.name,\n            box_x1=round(float(pb[0]),2), box_y1=round(float(pb[1]),2),\n            box_x2=round(float(pb[2]),2), box_y2=round(float(pb[3]),2),\n            confidence=round(float(conf),4),\n            class_id=int(pc),\n            class_name=CLASS_NAMES[pc] if 0<=pc<NUM_CLASSES else f'cls{pc}'\n        ))\n\npd.DataFrame(rows).to_csv(PREDS_TEST_DIR/'test_predictions.csv', index=False)\nprint(f'Saved: {PREDS_TEST_DIR}/test_predictions.csv')\nprint(f'Total test predictions: {len(rows)}')\n\n# ── List all outputs ──────────────────────────────────────────────────────\nprint('\\nAll saved outputs:')\nfor folder in [EDA_DIR, PEDIATRIC_DIR, IMBALANCE_DIR,\n               CURVES_DIR, EVAL_DIR, PREDS_VALID_DIR, PREDS_TEST_DIR]:\n    files = [f for f in sorted(folder.glob('*')) if f.is_file()]\n    if files:\n        print(f'  {folder}/')\n        for f in files:\n            print(f'    {f.name:50s} ({f.stat().st_size/1024:.1f} KB)')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"sec10-hdr","cell_type":"markdown","source":"---\n## Section 10 — Final Summary","metadata":{}},{"id":"sec10-print","cell_type":"code","source":"print('='*72)\nprint('    ALPHADENT 2-CLASS ENHANCED MODEL — FINAL SUMMARY')\nprint('='*72)\n\nprint(\"\"\"\nSCOPE REDUCTION\n---------------\n  Original classes : 9\n  Dropped          : Abrasion (class 0), Crown (class 2)\n  Merging          : Caries C1..C6 (orig 3-8) -> single 'Caries' class\n  Final classes    : 2\n    0 = Filling   (orig class 1)\n    1 = Caries    (orig classes 3+4+5+6+7+8 merged)\n\nENHANCEMENT APPLIED\n-------------------\"\"\")\nprint(f'  Model           : yolo11m-seg.pt  (was: yolo11s.pt detect)')\nprint(f'  Task            : segmentation    (was: detection)')\nprint(f'  imgsz           : {TRAIN_CONFIG[\"imgsz\"]}px         (was: 640px)')\nprint(f'  batch           : {TRAIN_CONFIG[\"batch\"]}              (was: 4)')\nprint(f'  epochs          : {TRAIN_CONFIG[\"epochs\"]}            (was: 100)')\nprint(f'  lr0             : {TRAIN_CONFIG[\"lr0\"]}         (was: 0.001)')\nprint(f'  mosaic          : {TRAIN_CONFIG[\"mosaic\"]}            (was: 1.0)')\nprint(f'  mixup           : {TRAIN_CONFIG[\"mixup\"]}            (was: 0.1, now disabled)')\nprint(f'  copy_paste      : {TRAIN_CONFIG[\"copy_paste\"]}            (was: 0.05, now disabled)')\nprint(f'  CLAHE           : {USE_CLAHE}          (soft LAB luminance enhancement)')\nprint(f'  Tiling          : {USE_TILING}          ({TILE_SIZE}x{TILE_SIZE}px, {TILE_OVERLAP*100:.0f}% overlap)')\n\nprint(\"\"\"\nDATASET STATISTICS (after remapping + enhancement)\n---------------------------------------------------\"\"\")\nprint(f'  Train images (original) : {len(train_imgs)}')\nprint(f'  Valid images (original) : {len(valid_imgs)}')\nprint(f'  Test  images            : {len(test_imgs)} (labels withheld)')\nif USE_TILING:\n    n_train_tiles = len(list(TILE_IMAGES_TRAIN.glob('*.jpg')))\n    n_valid_tiles = len(list(TILE_IMAGES_VALID.glob('*.jpg')))\n    print(f'  Train tiles (1024px)    : {n_train_tiles}')\n    print(f'  Valid tiles (1024px)    : {n_valid_tiles}')\nfor i in range(NUM_CLASSES):\n    print(f'  {CLASS_NAMES[i]:8s} train boxes: {int(train_box_counts[i])}')\n\nprint(f\"\"\"\nIMBALANCE STRATEGY\n------------------\n  After merging all 6 Caries sub-types, the two classes are nearly\n  balanced (Caries ~3,318 vs Filling ~2,187 train masks; ratio ~1.52).\n  Strategy: mild inverse-frequency weighted sampling for Filling images.\n  Max image weight: {max(w_vals):.2f}x  (vs 138x in the 9-class setup)\n  No extreme augmentation needed.\n\nFINAL VALIDATION METRICS (Segmentation)\n-----------------------------------------\n  mAP50      : {map50:.4f}\n  mAP50-95   : {map50_95:.4f}\n  Precision  : {prec:.4f}\n  Recall     : {rec:.4f}\n  F1         : {f1:.4f}\n\nExpected range (from guide): mAP50 ≈ 0.65+, mAP50-95 ≈ 0.40+\n\"\"\")\n\nprint('PER-CLASS AP (Segmentation Masks)')\nprint('-----------------------------------')\nprint(df_pc[['class_id','class_name','AP50','AP50_95','Precision','Recall','F1']].to_string(\n    index=False, float_format='{:.4f}'.format))\n\nprint(f\"\"\"\nOUTPUT LOCATIONS\n----------------\n  EDA plots            : {EDA_DIR}\n  Pediatric samples    : {PEDIATRIC_DIR}\n  Imbalance stats      : {IMBALANCE_DIR}\n  Training curves      : {CURVES_DIR}\n  Evaluation metrics   : {EVAL_DIR}\n  Validation preds     : {PREDS_VALID_DIR}\n  Test predictions     : {PREDS_TEST_DIR}\n  Results CSV          : {RUNS_ROOT/'results.csv'}\n  Best weights         : {BEST_CKPT}\n\"\"\")\nprint('='*72)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"31d877ab-c2f2-480a-9ce1-6a54577baa02","cell_type":"code","source":"# predict_single_seg.py\n# YOLO11m Segmentation — Single Image Prediction (Filling & Caries)\n# Enhanced: supports segmentation masks + bounding boxes\n\nimport cv2\nimport numpy as np\nfrom pathlib import Path\nfrom ultralytics import YOLO\n\n# Class configuration (2 classes only)\nCLASS_NAMES = {\n    0: \"Filling\",\n    1: \"Caries\"\n}\n\nCLASS_COLORS = {\n    0: (0, 200, 255),   # Amber/Gold — Filling\n    1: (0, 60, 230)     # Red — Caries\n}\n\n\ndef apply_clahe_lab(img_bgr, clip_limit=2.0, tile_grid=(8,8)):\n    \"\"\"Apply soft CLAHE on LAB luminance channel (same as training).\"\"\"\n    lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB)\n    l, a, b_ch = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid)\n    l_eq = clahe.apply(l)\n    return cv2.cvtColor(cv2.merge([l_eq, a, b_ch]), cv2.COLOR_LAB2BGR)\n\n\ndef predict_single_image(\n    image_path: str,\n    model_path: str = \"best.pt\",\n    conf: float = 0.25,\n    imgsz: int = 1024,\n    apply_clahe: bool = True,\n    save: bool = True\n):\n    \"\"\"\n    Run YOLO11m-seg inference on a single dental image.\n    Renders both segmentation masks and bounding boxes.\n    \n    Args:\n        image_path:   Path to input image\n        model_path:   Path to best.pt (yolo11m-seg weights)\n        conf:         Confidence threshold\n        imgsz:        Inference image size (must match training: 1024)\n        apply_clahe:  Apply same CLAHE preprocessing as during training\n        save:         Save output image with annotations\n    \"\"\"\n    image_path = Path(image_path)\n\n    # Load model (segmentation)\n    model = YOLO(model_path)\n\n    # Read and optionally preprocess image\n    image = cv2.imread(str(image_path))\n    input_image = apply_clahe_lab(image) if apply_clahe else image\n\n    # Run inference\n    results = model.predict(\n        source=input_image,\n        conf=conf,\n        imgsz=imgsz,\n        verbose=False\n    )\n\n    result = results[0]\n    draw = image.copy()  # Draw on original (not CLAHE) for display\n\n    # Draw segmentation masks (if available)\n    if result.masks is not None:\n        for mask_data, cls_id in zip(result.masks.data.cpu().numpy(),\n                                      result.boxes.cls.cpu().numpy().astype(int)):\n            color = CLASS_COLORS.get(cls_id, (200, 200, 200))\n            # Resize mask to image size and apply colored overlay\n            mask_resized = cv2.resize(mask_data, (draw.shape[1], draw.shape[0]))\n            mask_bool = mask_resized > 0.5\n            overlay = draw.copy()\n            overlay[mask_bool] = color\n            draw = cv2.addWeighted(draw, 0.6, overlay, 0.4, 0)\n\n    # Draw bounding boxes and labels\n    if result.boxes is not None:\n        for pb, cls_id, score in zip(result.boxes.xyxy.cpu().numpy(),\n                                      result.boxes.cls.cpu().numpy().astype(int),\n                                      result.boxes.conf.cpu().numpy()):\n            x1, y1, x2, y2 = map(int, pb)\n            color = CLASS_COLORS.get(cls_id, (200, 200, 200))\n            label = f\"{CLASS_NAMES.get(cls_id, cls_id)} {score:.0%}\"\n            cv2.rectangle(draw, (x1, y1), (x2, y2), color, 2)\n            cv2.putText(draw, label, (x1, max(y1 - 10, 15)),\n                        cv2.FONT_HERSHEY_SIMPLEX, 0.6, color, 2, cv2.LINE_AA)\n\n    # Save output\n    if save:\n        output_path = image_path.with_name(image_path.stem + \"_pred_seg.jpg\")\n        cv2.imwrite(str(output_path), draw)\n        print(f\"Saved: {output_path}\")\n\n    return draw\n\n\n# Example usage\nif __name__ == \"__main__\":\n    predict_single_image(\"image.jpg\", \"best.pt\", apply_clahe=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}