{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":1462296,"sourceType":"datasetVersion","datasetId":857191},{"sourceId":604441,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":403514,"modelId":421446}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# Kaggle: Server-centric FL+KD (WBF) for YOLO Student (FIXED)\n# - Prepares dataset from COCO + optional ImageNet LOC\n# - Server-centric: client-forward -> WBF -> mix with Teacher -> pseudo-label -> train Student (1 epoch/round)\n# - Visualization of Teacher vs Student vs WBF\n# - Fixes:\n#   * Removed invalid \"with cv2.imread(...)\" (use PIL.Image.open)\n#   * Replaced truthiness checks on arrays/lists with len(...) > 0\n# ============================================================","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:43:16.19436Z","iopub.execute_input":"2025-09-22T02:43:16.194821Z","iopub.status.idle":"2025-09-22T02:43:16.198962Z","shell.execute_reply.started":"2025-09-22T02:43:16.194793Z","shell.execute_reply":"2025-09-22T02:43:16.198231Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Install dependencies","metadata":{}},{"cell_type":"code","source":"!pip -q install \"ultralytics>=8.3.0\" ensemble-boxes pycocotools opencv-python-headless\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:43:16.200975Z","iopub.execute_input":"2025-09-22T02:43:16.201587Z","iopub.status.idle":"2025-09-22T02:44:33.200332Z","shell.execute_reply.started":"2025-09-22T02:43:16.201554Z","shell.execute_reply":"2025-09-22T02:44:33.199546Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports & setup","metadata":{}},{"cell_type":"code","source":"import os, sys, random, time, shutil, glob, math, json, warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport cv2\nimport yaml\nimport matplotlib.pyplot as plt\nfrom PIL import Image  # <-- needed for ImageNet W,H\n\nimport torch\nfrom ultralytics import YOLO\nfrom ensemble_boxes import weighted_boxes_fusion\n\ntry:\n    from pycocotools.coco import COCO\nexcept Exception as e:\n    print(\"pycocotools import failed:\", e)\n    raise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:33.201318Z","iopub.execute_input":"2025-09-22T02:44:33.201516Z","iopub.status.idle":"2025-09-22T02:44:38.637366Z","shell.execute_reply.started":"2025-09-22T02:44:33.201491Z","shell.execute_reply":"2025-09-22T02:44:38.636585Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CONFIG (edit if needed)","metadata":{}},{"cell_type":"code","source":"SEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {DEVICE}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n\nDATASET_DIR = '/kaggle/working/dataset'\nIMAGES_DIR  = os.path.join(DATASET_DIR, 'images')\nLABELS_DIR  = os.path.join(DATASET_DIR, 'labels')\nos.makedirs(IMAGES_DIR, exist_ok=True)\nos.makedirs(LABELS_DIR, exist_ok=True)\n\nNAMES = ['person', 'phone', 'reflex_camera', 'polaroid_camera']\nNC    = len(NAMES)\n\nTEACHER_PATH = '/kaggle/input/finetuneyoloteacher/pytorch/default/6/teacher_best.pt'  # <-- update if needed\nSTUDENT_VARIANT = 'n'  # 'n'|'s'|'m'\n\n# FL+KD schedule\nNUM_ROUNDS = 3\nNUM_CLIENTS = 5\nIMAGES_PER_ROUND = 1200\nIMG_SIZE = 640\nALPHA_START = 1.0\nALPHA_END   = 0.5\n\nFLKD_OUTDIR = '/kaggle/working/flkd'\nos.makedirs(FLKD_OUTDIR, exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.638898Z","iopub.execute_input":"2025-09-22T02:44:38.639325Z","iopub.status.idle":"2025-09-22T02:44:38.721376Z","shell.execute_reply.started":"2025-09-22T02:44:38.639305Z","shell.execute_reply":"2025-09-22T02:44:38.720805Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Auto-detect inputs (COCO & ImageNet)","metadata":{}},{"cell_type":"code","source":"def detect_coco_root():\n    candidates = [\n        '/kaggle/input/coco-2017-dataset/coco2017',\n        '/kaggle/input/coco-2017-dataset',\n        '/kaggle/input/coco2017',\n        '/kaggle/input/coco-2017'\n    ]\n    for root in candidates:\n        ann = os.path.join(root, 'annotations', 'instances_train2017.json')\n        imgd = os.path.join(root, 'train2017')\n        if os.path.exists(ann) and os.path.isdir(imgd):\n            return root\n    for dirpath, dirnames, filenames in os.walk('/kaggle/input'):\n        if 'instances_train2017.json' in filenames:\n            root = os.path.dirname(dirpath) if os.path.basename(dirpath) == 'annotations' else dirpath\n            if os.path.isdir(os.path.join(root, 'train2017')):\n                return root\n    return None\n\ndef detect_imagenet_root():\n    root = '/kaggle/input/imagenet-object-localization-challenge'\n    return root if os.path.isdir(os.path.join(root, 'ILSVRC')) else None\n\nCOCO_ROOT = detect_coco_root()\nIMAGENET_ROOT = detect_imagenet_root()\nprint(f\"COCO root: {COCO_ROOT}\")\nprint(f\"ImageNet root: {IMAGENET_ROOT}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.722118Z","iopub.execute_input":"2025-09-22T02:44:38.722339Z","iopub.status.idle":"2025-09-22T02:44:38.734329Z","shell.execute_reply.started":"2025-09-22T02:44:38.722322Z","shell.execute_reply":"2025-09-22T02:44:38.733633Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare dataset (COCO + optional ImageNet LOC)","metadata":{}},{"cell_type":"code","source":"def prepare_dataset():\n    \"\"\"\n    Build YOLO-style dataset under /kaggle/working/dataset\n    images/ and labels/ with 4 classes: person, phone, reflex_camera, polaroid_camera\n    \"\"\"\n    class_counts = {0:0, 1:0, 2:0, 3:0}\n\n    # ---- ImageNet LOC -> classes 2 & 3 (optional)\n    if IMAGENET_ROOT:\n        print(\"\\n=== Processing ImageNet LOC ===\")\n        try:\n            annotations_file = os.path.join(IMAGENET_ROOT, 'LOC_train_solution.csv')\n            if os.path.exists(annotations_file):\n                annotations = pd.read_csv(annotations_file)\n                imagenet_classes = {\n                    'n02992529': 1, # mobile phone -> phone\n                    'n04069434': 2, # reflex camera\n                    'n03976467': 3  # polaroid camera\n                }\n                desired = list(imagenet_classes.keys())\n\n                filtered = annotations[annotations['PredictionString'].str.contains('|'.join(desired))]\n                print(f\"Found {len(filtered)} candidate records from ImageNet\")\n\n                max_per_class = {1: 1500, 2: 1000, 3: 1000}\n\n                for _, row in tqdm(filtered.iterrows(), total=len(filtered)):\n                    image_id = row['ImageId']\n                    predictions = row['PredictionString'].split()\n                    i = 0\n                    wrote_any = False\n\n                    label_path = os.path.join(LABELS_DIR, f'imagenet_{image_id}.txt')\n                    lines = []\n                    while i < len(predictions):\n                        if predictions[i] in desired:\n                            class_id = predictions[i]\n                            xmin, ymin, xmax, ymax = map(float, predictions[i+1:i+5])\n\n                            img_path = os.path.join(IMAGENET_ROOT, 'ILSVRC', 'Data', 'CLS-LOC', 'train', class_id, f'{image_id}.JPEG')\n                            if not os.path.exists(img_path):\n                                i += 5\n                                continue\n\n                            yolo_cls = imagenet_classes[class_id]\n                            if class_counts[yolo_cls] >= max_per_class.get(yolo_cls, 10**9):\n                                i += 5\n                                continue\n\n                            # Copy image once\n                            out_img = os.path.join(IMAGES_DIR, f'imagenet_{image_id}.jpg')\n                            if not os.path.exists(out_img):\n                                shutil.copy(img_path, out_img)\n\n                            # Get width/height using PIL (NO cv2 context manager!)\n                            with Image.open(img_path) as im:\n                                W, H = im.size\n\n                            x_center = ((xmin + xmax) / 2) / W\n                            y_center = ((ymin + ymax) / 2) / H\n                            w_norm   = (xmax - xmin) / W\n                            h_norm   = (ymax - ymin) / H\n\n                            lines.append(f\"{yolo_cls} {x_center:.6f} {y_center:.6f} {w_norm:.6f} {h_norm:.6f}\\n\")\n                            class_counts[yolo_cls] += 1\n                            wrote_any = True\n                            i += 5\n                        else:\n                            i += 1\n\n                    if wrote_any:\n                        mode = 'w' if not os.path.exists(label_path) else 'a'\n                        with open(label_path, mode) as f:\n                            f.writelines(lines)\n            else:\n                print(\"ImageNet LOC csv not found -> skip.\")\n        except Exception as e:\n            print(\"ImageNet processing failed:\", e)\n\n    # ---- COCO -> classes 0 & 1\n    if COCO_ROOT:\n        print(\"\\n=== Processing COCO ===\")\n        try:\n            coco = COCO(os.path.join(COCO_ROOT, 'annotations', 'instances_train2017.json'))\n            person_ids = coco.getImgIds(catIds=[1])   # person\n            phone_ids  = coco.getImgIds(catIds=[77])  # cell phone\n            max_person_images = 2000\n            if len(person_ids) > max_person_images:\n                person_ids = random.sample(person_ids, max_person_images)\n            all_ids = list(set(person_ids + phone_ids))\n            print(f\"Person imgs: {len(person_ids)} | Phone imgs: {len(phone_ids)} | Unique: {len(all_ids)}\")\n\n            for img_id in tqdm(all_ids):\n                img_info = coco.loadImgs(img_id)[0]\n                ann_ids = coco.getAnnIds(imgIds=img_id, catIds=[1,77])\n                anns = coco.loadAnns(ann_ids)\n                if not anns:\n                    continue\n\n                src_img = os.path.join(COCO_ROOT, 'train2017', img_info['file_name'])\n                if not os.path.exists(src_img):\n                    continue\n\n                dst_img = os.path.join(IMAGES_DIR, f\"coco_{img_info['file_name']}\")\n                if not os.path.exists(dst_img):\n                    shutil.copy(src_img, dst_img)\n\n                lbl_path = os.path.join(LABELS_DIR, f\"coco_{Path(img_info['file_name']).stem}.txt\")\n                with open(lbl_path, 'w') as f:\n                    for ann in anns:\n                        if ann['category_id'] == 1:\n                            ycls = 0\n                        elif ann['category_id'] == 77:\n                            ycls = 1\n                        else:\n                            continue\n                        x,y,w,h = ann['bbox']\n                        if w <= 0 or h <= 0:\n                            continue\n                        x_c = (x + w/2) / img_info['width']\n                        y_c = (y + h/2) / img_info['height']\n                        w_n = w / img_info['width']\n                        h_n = h / img_info['height']\n                        x_c = max(0,min(1,x_c)); y_c = max(0,min(1,y_c))\n                        w_n = max(0,min(1,w_n)); h_n = max(0,min(1,h_n))\n                        f.write(f\"{ycls} {x_c:.6f} {y_c:.6f} {w_n:.6f} {h_n:.6f}\\n\")\n                        class_counts[ycls] += 1\n        except Exception as e:\n            print(\"COCO processing failed:\", e)\n    else:\n        print(\"COCO dataset not found -> skipping COCO stage.\")\n\n    # ---- Stats\n    print(\"\\n=== Dataset Stats ===\")\n    for i, cname in enumerate(NAMES):\n        print(f\"{cname}: {class_counts.get(i,0)} instances\")\n    if class_counts[2] == 0 or class_counts[3] == 0:\n        print(\"⚠️ Note: No samples for reflex/polaroid cameras (ImageNet stage failed or not present).\")\n\n    # ---- Train/Val split\n    all_images = [f for f in os.listdir(IMAGES_DIR) if f.lower().endswith(('.jpg','.jpeg','.png'))]\n    random.shuffle(all_images)\n    split = int(0.8 * len(all_images))\n    train_images = all_images[:split]\n    val_images   = all_images[split:]\n\n    with open(os.path.join(DATASET_DIR, 'train.txt'), 'w') as f:\n        for img in train_images:\n            f.write(os.path.join(IMAGES_DIR, img) + '\\n')\n    with open(os.path.join(DATASET_DIR, 'val.txt'), 'w') as f:\n        for img in val_images:\n            f.write(os.path.join(IMAGES_DIR, img) + '\\n')\n\n    data_yaml = os.path.join(DATASET_DIR, 'data.yaml')\n    with open(data_yaml, 'w') as f:\n        f.write(f\"train: {os.path.join(DATASET_DIR, 'train.txt')}\\n\")\n        f.write(f\"val: {os.path.join(DATASET_DIR, 'val.txt')}\\n\")\n        f.write(f\"nc: {NC}\\n\")\n        f.write(f\"names: {NAMES}\\n\")\n    print(f\"\\nTotal images: {len(all_images)} | train: {len(train_images)} | val: {len(val_images)}\")\n    print(f\"data.yaml -> {data_yaml}\")\n\n    return data_yaml, os.path.join(DATASET_DIR, 'train.txt'), os.path.join(DATASET_DIR, 'val.txt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.735161Z","iopub.execute_input":"2025-09-22T02:44:38.735423Z","iopub.status.idle":"2025-09-22T02:44:38.75477Z","shell.execute_reply.started":"2025-09-22T02:44:38.735389Z","shell.execute_reply":"2025-09-22T02:44:38.754098Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Teacher path auto-detect fallback","metadata":{}},{"cell_type":"code","source":"def autodetect_teacher_path(default_path):\n    if os.path.exists(default_path):\n        return default_path\n    print(f\"Teacher path not found: {default_path}. Scanning /kaggle/input for *.pt ...\")\n    pts = glob.glob('/kaggle/input/**/*.pt', recursive=True)\n    if pts:\n        print(\"Found candidate:\", pts[0])\n        return pts[0]\n    raise FileNotFoundError(\"No teacher weights (.pt) found under /kaggle/input. Please set TEACHER_PATH correctly.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.755592Z","iopub.execute_input":"2025-09-22T02:44:38.756269Z","iopub.status.idle":"2025-09-22T02:44:38.773468Z","shell.execute_reply.started":"2025-09-22T02:44:38.756241Z","shell.execute_reply":"2025-09-22T02:44:38.772793Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# WBF helpers","metadata":{}},{"cell_type":"code","source":"def results_to_norm_lists(res, W, H):\n    \"\"\"Ultralytics Results -> XYXY normalized [0,1] lists for WBF.\"\"\"\n    if res.boxes is None or len(res.boxes) == 0:\n        return [], [], []\n    xyxy = res.boxes.xyxy.detach().cpu().numpy()\n    conf = res.boxes.conf.detach().cpu().numpy()\n    cls  = res.boxes.cls.detach().cpu().numpy().astype(int)\n    boxes = [[x1/W, y1/H, x2/W, y2/H] for (x1,y1,x2,y2) in xyxy]\n    return boxes, conf.tolist(), cls.tolist()\n\ndef xyxy_to_xywhn(boxes_xyxy):\n    out = []\n    for x1,y1,x2,y2 in boxes_xyxy:\n        w = max(1e-6, x2-x1); h = max(1e-6, y2-y1)\n        cx = x1 + w/2; cy = y1 + h/2\n        out.append([cx, cy, w, h])\n    return out\n\ndef photometric_augment(img, blur=None, brightness=None):\n    im = img.copy()\n    if blur and blur > 0:\n        k = int(blur*2) | 1\n        im = cv2.GaussianBlur(im, (k,k), blur)\n    if brightness:\n        hsv = cv2.cvtColor(im, cv2.COLOR_RGB2HSV).astype(np.float32)\n        hsv[...,2] = np.clip(hsv[...,2]*brightness, 0, 255)\n        im = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB)\n    return im\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.7742Z","iopub.execute_input":"2025-09-22T02:44:38.774431Z","iopub.status.idle":"2025-09-22T02:44:38.790021Z","shell.execute_reply.started":"2025-09-22T02:44:38.774404Z","shell.execute_reply":"2025-09-22T02:44:38.789356Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Server-centric FL+KD Trainer (patched truthiness checks)","metadata":{}},{"cell_type":"code","source":"class ServerFLKDTrainer:\n    \"\"\"\n    Server-centric pipeline:\n      - n client forward (student) on augmented views -> per-client detections\n      - WBF across clients -> y_agg\n      - teacher forward on same images -> y_T\n      - WBF([y_T, y_agg], weights=[alpha, 1-alpha]) -> y_star\n      - write pseudo labels (copy subset images) -> train student (1 epoch)\n    \"\"\"\n    def __init__(self, teacher_path, student_variant='n', names=None, imgsz=640, out_dir='/kaggle/working/flkd'):\n        self.names = names or ['person','phone','reflex_camera','polaroid_camera']\n        self.nc = len(self.names)\n        self.imgsz = imgsz\n        self.out_dir = out_dir\n        os.makedirs(out_dir, exist_ok=True)\n\n        print(\"Loading teacher ...\")\n        self.teacher = YOLO(teacher_path)\n        self.teacher.model.to(DEVICE).eval()\n\n        print(f\"Creating student ({student_variant}) ...\")\n        student_ckpts = {'n':'yolo11n.pt','s':'yolo11s.pt','m':'yolo11m.pt'}\n        self.student = YOLO(student_ckpts[student_variant])\n        self.student.model.to(DEVICE)\n\n    def _client_forward_batch(self, rgbs, domain):\n        aug_rgbs = [photometric_augment(im, domain.get('blur'), domain.get('brightness')) for im in rgbs]\n        res = self.student(aug_rgbs, imgsz=self.imgsz, conf=0.001, iou=0.7, verbose=False)\n        return res\n\n    def _teacher_forward_batch(self, rgbs):\n        res = self.teacher(rgbs, imgsz=self.imgsz, conf=0.001, iou=0.7, verbose=False)\n        return res\n\n    def _aggregate_clients_wbf(self, multi_client_results, sizes, weights=None,\n                               iou_thr=0.55, skip_box_thr=0.001):\n        C = len(multi_client_results)\n        B = len(multi_client_results[0])\n        fused = []\n        for b in range(B):\n            boxes_list, scores_list, labels_list = [], [], []\n            H,W = sizes[b]\n            for c in range(C):\n                bxs, scs, lbs = results_to_norm_lists(multi_client_results[c][b], W, H)\n                if len(bxs) > 0:  # <-- fixed truthiness check\n                    boxes_list.append(bxs); scores_list.append(scs); labels_list.append(lbs)\n            if len(boxes_list) > 0:  # <-- fixed\n                ws = weights if (weights and len(weights)==len(boxes_list)) else None\n                f_boxes, f_scores, f_labels = weighted_boxes_fusion(\n                    boxes_list, scores_list, labels_list,\n                    weights=ws, iou_thr=iou_thr, skip_box_thr=skip_box_thr\n                )\n            else:\n                f_boxes, f_scores, f_labels = [], [], []\n            fused.append({'boxes_norm': f_boxes, 'scores': f_scores, 'labels': f_labels})\n        return fused\n\n    def _mix_wbf(self, yT, yAgg, alpha=0.7, iou_thr=0.55):\n        mixed = []\n        for b in range(len(yT)):\n            boxes_list, scores_list, labels_list = [], [], []\n            if len(yT[b]['boxes_norm']) > 0:  # <-- fixed\n                boxes_list.append(yT[b]['boxes_norm']); scores_list.append(yT[b]['scores']); labels_list.append(yT[b]['labels'])\n            if len(yAgg[b]['boxes_norm']) > 0:  # <-- fixed\n                boxes_list.append(yAgg[b]['boxes_norm']); scores_list.append(yAgg[b]['scores']); labels_list.append(yAgg[b]['labels'])\n            if len(boxes_list) > 0:  # <-- fixed\n                w = [alpha, 1.0-alpha][:len(boxes_list)]\n                f_boxes, f_scores, f_labels = weighted_boxes_fusion(\n                    boxes_list, scores_list, labels_list,\n                    weights=w, iou_thr=iou_thr, skip_box_thr=0.001\n                )\n            else:\n                f_boxes, f_scores, f_labels = [], [], []\n            mixed.append({'boxes_norm': f_boxes, 'scores': f_scores, 'labels': f_labels})\n        return mixed\n\n    def _write_round_dataset(self, image_paths, fused_targets, save_root):\n        \"\"\"\n        Copy selected images to save_root/images and write labels in save_root/labels\n        so Ultralytics can find paired labels by replacing 'images' -> 'labels'.\n        \"\"\"\n        img_dir = os.path.join(save_root, 'images'); os.makedirs(img_dir, exist_ok=True)\n        lbl_dir = os.path.join(save_root, 'labels'); os.makedirs(lbl_dir, exist_ok=True)\n\n        for p, tgt in zip(image_paths, fused_targets):\n            # copy image\n            dst_img = os.path.join(img_dir, os.path.basename(p))\n            if not os.path.exists(dst_img):\n                shutil.copy(p, dst_img)\n\n            # write label\n            stem = Path(p).stem\n            lp = os.path.join(lbl_dir, f\"{stem}.txt\")\n            keep = [(b,s,l) for b,s,l in zip(tgt['boxes_norm'], tgt['scores'], tgt['labels']) if s >= 0.25]\n            if len(keep) == 0:\n                open(lp, 'w').close()\n                continue\n            boxes_xywh = xyxy_to_xywhn([k[0] for k in keep])\n            with open(lp, 'w') as f:\n                for (cx,cy,w,h), (_,sc,lb) in zip(boxes_xywh, keep):\n                    f.write(f\"{int(lb)} {cx:.6f} {cy:.6f} {w:.6f} {h:.6f}\\n\")\n\n        # YAML\n        data_yaml = os.path.join(save_root, 'data.yaml')\n        with open(data_yaml, 'w') as f:\n            f.write(f\"path: {save_root}\\n\")\n            f.write(f\"train: images\\n\")\n            f.write(f\"val: images\\n\")\n            f.write(f\"nc: {self.nc}\\n\")\n            f.write(f\"names: {self.names}\\n\")\n        return data_yaml\n\n    def run(self, proxy_txt, num_rounds=3, num_clients=5, images_per_round=1200,\n            alpha_start=1.0, alpha_end=0.5):\n        with open(proxy_txt, 'r') as f:\n            pool = [l.strip() for l in f.readlines()]\n        print(f\"[Server] Proxy pool size: {len(pool)}\")\n    \n        domains = [{'blur': random.uniform(0.2, 1.5), 'brightness': random.uniform(0.7, 1.3)}\n                   for _ in range(num_clients)]\n    \n        for r in range(num_rounds):\n            alpha = alpha_start + (alpha_end - alpha_start) * (r / max(1, num_rounds-1))\n            subset = random.sample(pool, min(images_per_round, len(pool)))\n            print(f\"\\n=== Round {r+1}/{num_rounds} | alpha={alpha:.2f} | subset={len(subset)} ===\")\n    \n            fused_all = []\n            BATCH = 16\n            for i in range(0, len(subset), BATCH):\n                paths = subset[i:i+BATCH]\n                rgbs, sizes = [], []\n                for p in paths:\n                    bgr = cv2.imread(p)\n                    if bgr is None:  # skip silently\n                        continue\n                    rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)\n                    rgbs.append(rgb); sizes.append(rgb.shape[:2])\n    \n                if len(rgbs) == 0:\n                    fused_all.extend([{'boxes_norm':[], 'scores':[], 'labels':[]}] * len(paths))\n                    continue\n    \n                # client forwards\n                per_client_results = []\n                for c in range(num_clients):\n                    res_c = self._client_forward_batch(rgbs, domains[c])\n                    per_client_results.append(res_c)\n    \n                # WBF across clients\n                y_agg = self._aggregate_clients_wbf(per_client_results, sizes)\n    \n                # teacher forward\n                yT_raw = self._teacher_forward_batch(rgbs)\n                y_T = []\n                for b_idx, res in enumerate(yT_raw):\n                    H, W = sizes[b_idx]\n                    bxs, scs, lbs = results_to_norm_lists(res, W, H)\n                    y_T.append({'boxes_norm': bxs, 'scores': scs, 'labels': lbs})\n    \n                # mix\n                y_star = self._mix_wbf(y_T, y_agg, alpha=alpha)\n                fused_all.extend(y_star)\n    \n            # write round dataset & train student (1 epoch)\n            save_root = os.path.join(self.out_dir, f\"round_{r+1}\")\n            os.makedirs(save_root, exist_ok=True)\n            data_yaml = self._write_round_dataset(subset, fused_all, save_root)\n    \n            print(f\"[Server] Train student on pseudo labels (round {r+1}) ...\")\n            self.student.train(\n                data=data_yaml, imgsz=self.imgsz, epochs=1, batch=16,\n                lr0=0.001, patience=0,\n                device=(0 if DEVICE.type=='cuda' else 'cpu'),\n                verbose=False, cache=False, save=True,\n                project=self.out_dir, name=f\"student_round_{r+1}\"\n            )\n    \n            # =========================\n            # [FIX] Re-init student từ weight của vòng vừa train\n            # để tránh KeyError: 'model' khi gọi .train() ở vòng tiếp theo\n            weights_dir = os.path.join(self.out_dir, f\"student_round_{r+1}\", \"weights\")\n            resume_ckpt = None\n            for nm in (\"best.pt\", \"last.pt\"):\n                cand = os.path.join(weights_dir, nm)\n                if os.path.exists(cand):\n                    resume_ckpt = cand\n                    break\n            if resume_ckpt:\n                print(f\"[Server] Re-init student from {resume_ckpt}\")\n                self.student = YOLO(resume_ckpt)\n                self.student.model.to(DEVICE)\n            else:\n                print(f\"[Server][WARN] No weights in {weights_dir}; keep current student\")\n            # =========================\n    \n        print(\"\\n[Server] FL+KD complete.\")\n        try:\n            self.student.save(os.path.join(self.out_dir, 'student_flkd.pt'))\n        except Exception:\n            pass\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.791824Z","iopub.execute_input":"2025-09-22T02:44:38.792019Z","iopub.status.idle":"2025-09-22T02:44:38.814911Z","shell.execute_reply.started":"2025-09-22T02:44:38.792004Z","shell.execute_reply":"2025-09-22T02:44:38.814243Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization (Teacher vs Student vs WBF)\n","metadata":{}},{"cell_type":"code","source":"def find_latest_student_ckpt(root=FLKD_OUTDIR):\n    runs = sorted(glob.glob(os.path.join(root, \"student_round_*\")), key=lambda p: (len(p), p))\n    runs = [r for r in runs if os.path.isdir(r)]\n    if not runs:\n        return None\n    last = runs[-1]\n    for name in [\"best.pt\", \"last.pt\"]:\n        cand = os.path.join(last, \"weights\", name)\n        if os.path.exists(cand):\n            return cand\n    return None\n\ndef visualize_wbf(teacher_path, student_path, data_yaml_path, num_samples=6, alpha_wbf=0.7):\n    teacher = YOLO(teacher_path)\n    student = YOLO(student_path)\n\n    with open(data_yaml_path, 'r') as f:\n        cfg = yaml.safe_load(f)\n    val_file = cfg['val']\n    if os.path.isdir(val_file):\n        val_imgs = sorted([os.path.join(val_file, p) for p in os.listdir(val_file)])\n    else:\n        with open(val_file, 'r') as f:\n            val_imgs = [l.strip() for l in f.readlines()]\n\n    if len(val_imgs) == 0:\n        print(\"No val images for visualization.\")\n        return\n\n    samples = random.sample(val_imgs, min(num_samples, len(val_imgs)))\n    fig, axes = plt.subplots(len(samples), 4, figsize=(22, 5*len(samples)))\n    if len(samples) == 1:\n        axes = axes.reshape(1, -1)\n\n    for r, p in enumerate(samples):\n        bgr = cv2.imread(p); rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)\n        H, W = rgb.shape[:2]\n        t_res = teacher(p, verbose=False)[0]\n        s_res = student(p, verbose=False)[0]\n        t_img = t_res.plot(); s_img = s_res.plot()\n\n        b1,s1,l1 = results_to_norm_lists(t_res, W, H)\n        b2,s2,l2 = results_to_norm_lists(s_res, W, H)\n        if (len(b1) > 0) or (len(b2) > 0):  # <-- fixed\n            lists_b = [b1]; lists_s = [s1]; lists_l = [l1]; weights = [alpha_wbf]\n            if len(b2) > 0:\n                lists_b.append(b2); lists_s.append(s2); lists_l.append(l2); weights.append(1-alpha_wbf)\n            f_boxes, f_scores, f_labels = weighted_boxes_fusion(\n                lists_b, lists_s, lists_l, weights=weights, iou_thr=0.55, skip_box_thr=0.001\n            )\n            fused = rgb.copy()\n            for (x1,y1,x2,y2),sc,lb in zip(f_boxes,f_scores,f_labels):\n                x1,y1,x2,y2 = int(x1*W), int(y1*H), int(x2*W), int(y2*H)\n                cv2.rectangle(fused,(x1,y1),(x2,y2),(0,255,0),2)\n                cv2.putText(fused, f\"{NAMES[int(lb)]}:{sc:.2f}\", (x1,max(0,y1-4)),\n                            cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0,255,0), 2)\n        else:\n            fused = rgb\n\n        axes[r,0].imshow(rgb);   axes[r,0].set_title(\"Original\"); axes[r,0].axis('off')\n        axes[r,1].imshow(t_img); axes[r,1].set_title(\"Teacher\");  axes[r,1].axis('off')\n        axes[r,2].imshow(s_img); axes[r,2].set_title(\"Student\");  axes[r,2].axis('off')\n        axes[r,3].imshow(fused); axes[r,3].set_title(\"WBF(T+S)\"); axes[r,3].axis('off')\n\n    plt.tight_layout()\n    out_png = '/kaggle/working/detection_comparison_wbf.png'\n    plt.savefig(out_png, dpi=150, bbox_inches='tight')\n    print(\"Saved visualization ->\", out_png)\n    plt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.815765Z","iopub.execute_input":"2025-09-22T02:44:38.815998Z","iopub.status.idle":"2025-09-22T02:44:38.834775Z","shell.execute_reply.started":"2025-09-22T02:44:38.815977Z","shell.execute_reply":"2025-09-22T02:44:38.834267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MAIN","metadata":{}},{"cell_type":"code","source":"def main():\n    print(\"🚀 Server-centric FL+KD (WBF) — End-to-End on Kaggle\")\n    data_yaml, train_txt, val_txt = prepare_dataset()\n    teacher_path = autodetect_teacher_path(TEACHER_PATH)\n\n    print(\"\\n✅ Dataset ready at:\", data_yaml)\n    print(\"✅ Teacher weights:\", teacher_path)\n\n    server = ServerFLKDTrainer(\n        teacher_path=teacher_path,\n        student_variant=STUDENT_VARIANT,\n        names=NAMES,\n        imgsz=IMG_SIZE,\n        out_dir=FLKD_OUTDIR\n    )\n    server.run(\n        proxy_txt=train_txt,\n        num_rounds=NUM_ROUNDS,\n        num_clients=NUM_CLIENTS,\n        images_per_round=IMAGES_PER_ROUND,\n        alpha_start=ALPHA_START,\n        alpha_end=ALPHA_END\n    )\n\n    student_ckpt = find_latest_student_ckpt(FLKD_OUTDIR)\n    if student_ckpt is not None:\n        visualize_wbf(teacher_path, student_ckpt, data_yaml, num_samples=6, alpha_wbf=0.7)\n    else:\n        print(\"Student checkpoint not found for visualization.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.835468Z","iopub.execute_input":"2025-09-22T02:44:38.835797Z","iopub.status.idle":"2025-09-22T02:44:38.853248Z","shell.execute_reply.started":"2025-09-22T02:44:38.83578Z","shell.execute_reply":"2025-09-22T02:44:38.852632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T02:44:38.853864Z","iopub.execute_input":"2025-09-22T02:44:38.854127Z","iopub.status.idle":"2025-09-22T03:03:37.5921Z","shell.execute_reply.started":"2025-09-22T02:44:38.8541Z","shell.execute_reply":"2025-09-22T03:03:37.591103Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Add Metrics","metadata":{}},{"cell_type":"code","source":"# ==== CONFIG CHO KHỐI ĐÁNH GIÁ ====\nimport os, glob, json, math, random, yaml\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom PIL import Image\nfrom pathlib import Path\n\nimport torch\nfrom ultralytics import YOLO\nfrom pycocotools.coco import COCO\nfrom pycocotools.cocoeval import COCOeval\nfrom ensemble_boxes import weighted_boxes_fusion\n\n# Tận dụng các biến đã có từ notebook của bạn:\n# - DEVICE, NAMES, IMG_SIZE, DATASET_DIR, LABELS_DIR, FLKD_OUTDIR, TEACHER_PATH\n# - các hàm autodetect_teacher_path(), find_latest_student_ckpt() đã định nghĩa ở trên\n\n# Thư mục xuất hình\nMETRICS_DIR = '/kaggle/working/metrics_figs'\nos.makedirs(METRICS_DIR, exist_ok=True)\n\n# data.yaml đã được write trong prepare_dataset()\nDATA_YAML_PATH = os.path.join(DATASET_DIR, 'data.yaml')\nassert os.path.exists(DATA_YAML_PATH), \"Không tìm thấy data.yaml; hãy chạy chuẩn bị dataset trước.\"\n\nwith open(DATA_YAML_PATH, 'r') as f:\n    data_cfg = yaml.safe_load(f)\n\n# Lấy danh sách ảnh val\nval_spec = data_cfg['val']\nif os.path.isdir(val_spec):\n    val_imgs = sorted([os.path.join(val_spec, p) for p in os.listdir(val_spec)\n                       if p.lower().endswith(('.jpg','.jpeg','.png'))])\nelse:\n    with open(val_spec, 'r') as f:\n        val_imgs = [l.strip() for l in f.readlines() if l.strip()]\n\nprint(f\"Số ảnh val: {len(val_imgs)}\")\n\n# Teacher & Student checkpoint\nteacher_path = autodetect_teacher_path(TEACHER_PATH)\nstudent_ckpt = find_latest_student_ckpt(FLKD_OUTDIR)\nprint(\"Teacher:\", teacher_path)\nprint(\"Student (latest):\", student_ckpt)\n\n# Giới hạn nhanh nếu cần (đánh giá nhanh)\nEVAL_LIMIT = None   # ví dụ đặt 600 để chạy nhanh hơn, hoặc None để dùng toàn bộ\nif EVAL_LIMIT is not None:\n    rng = np.random.RandomState(42)\n    val_imgs = list(rng.choice(val_imgs, size=min(EVAL_LIMIT, len(val_imgs)), replace=False))\nprint(\"Đánh giá trên:\", len(val_imgs), \"ảnh\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_ultra_val(model_path, tag):\n    model = YOLO(model_path)\n    # project/name để gom hình về thư mục metrics_figs/ultralytics/<tag>\n    results = model.val(\n        data=DATA_YAML_PATH,\n        imgsz=IMG_SIZE,\n        iou=0.55,           # khớp với WBF IoU bạn dùng\n        conf=0.001,\n        plots=True,         # -> sinh PR, P/R/F1, confusion matrix, ...\n        save_json=True,     # -> COCO-format JSON (nếu cần)\n        project=os.path.join(METRICS_DIR, 'ultralytics'),\n        name=tag,\n        verbose=False,\n        device=(0 if torch.cuda.is_available() else 'cpu'),\n    )\n    outdir = os.path.join(METRICS_DIR, 'ultralytics', tag)\n    print(f\"[Ultralytics] Đã lưu hình vào: {outdir}\")\n    return results\n\n_ = run_ultra_val(teacher_path, 'teacher')\nif student_ckpt is not None:\n    _ = run_ultra_val(student_ckpt, 'student')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_coco_gt_from_yolo(val_img_paths, labels_dir, class_names):\n    images, annotations, img_id_map = [], [], {}\n    ann_id = 1\n    for idx, p in enumerate(val_img_paths):\n        image_id = idx + 1\n        img_id_map[p] = image_id\n        # Kích thước ảnh\n        try:\n            with Image.open(p) as im:\n                W, H = im.size\n        except Exception:\n            bgr = cv2.imread(p)\n            if bgr is None:\n                # bỏ qua ảnh lỗi\n                continue\n            H, W = bgr.shape[:2]\n\n        images.append({\n            'id': image_id,\n            'file_name': os.path.basename(p),\n            'width': W, 'height': H\n        })\n\n        stem = Path(p).stem\n        # 2 ứng viên đường dẫn label: theo chuẩn bạn đã ghi\n        lp1 = os.path.join(labels_dir, f\"{stem}.txt\")\n        lp2 = os.path.splitext(p.replace('/images/', '/labels/'))[0] + '.txt'\n        lp = lp1 if os.path.exists(lp1) else (lp2 if os.path.exists(lp2) else None)\n\n        if lp and os.path.exists(lp):\n            with open(lp, 'r') as f:\n                for line in f:\n                    parts = line.strip().split()\n                    if len(parts) < 5:\n                        continue\n                    c = int(float(parts[0]))\n                    cx, cy, w, h = map(float, parts[1:5])\n                    x = (cx - w/2) * W\n                    y = (cy - h/2) * H\n                    bw = w * W\n                    bh = h * H\n                    annotations.append({\n                        'id': ann_id,\n                        'image_id': image_id,\n                        'category_id': c + 1,  # 1-based\n                        'bbox': [float(max(0, x)), float(max(0, y)),\n                                 float(max(1e-3, bw)), float(max(1e-3, bh))],\n                        'area': float(max(1e-3, bw * bh)),\n                        'iscrowd': 0\n                    })\n                    ann_id += 1\n\n    categories = [{'id': i+1, 'name': name} for i, name in enumerate(class_names)]\n    cocoGt = COCO()\n    cocoGt.dataset = {'images': images, 'annotations': annotations, 'categories': categories}\n    cocoGt.createIndex()\n    return cocoGt, img_id_map\n\ncocoGt, img_id_map = build_coco_gt_from_yolo(val_imgs, LABELS_DIR, NAMES)\nprint(\"COCO GT images:\", len(cocoGt.dataset['images']),\n      \"| annos:\", len(cocoGt.dataset['annotations']))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_to_coco_dets(model_path, img_paths, img_id_map,\n                         imgsz=640, conf=0.001, iou=0.55, batch=16):\n    model = YOLO(model_path)\n    # đẩy model lên GPU nếu có\n    if hasattr(model, 'model'):\n        model.model.to(torch.device('cuda' if torch.cuda.is_available() else 'cpu'))\n    dets = []\n    for i in range(0, len(img_paths), batch):\n        batch_paths = img_paths[i:i+batch]\n        results = model(batch_paths, imgsz=imgsz, conf=conf, iou=iou,\n                        device=(0 if torch.cuda.is_available() else 'cpu'),\n                        verbose=False)\n        for p, res in zip(batch_paths, results):\n            image_id = img_id_map[p]\n            H, W = res.orig_shape\n            if res.boxes is None or len(res.boxes) == 0:\n                continue\n            xyxy = res.boxes.xyxy.detach().cpu().numpy()\n            confs = res.boxes.conf.detach().cpu().numpy()\n            clss  = res.boxes.cls.detach().cpu().numpy().astype(int)\n            for (x1, y1, x2, y2), sc, c in zip(xyxy, confs, clss):\n                # clip vào ảnh\n                x1 = float(max(0.0, min(x1, W - 1)))\n                y1 = float(max(0.0, min(y1, H - 1)))\n                w  = float(max(1e-3, min(x2 - x1, W - 1 - x1)))\n                h  = float(max(1e-3, min(y2 - y1, H - 1 - y1)))\n                dets.append({\n                    'image_id': image_id,\n                    'category_id': int(c) + 1,\n                    'bbox': [x1, y1, w, h],\n                    'score': float(sc)\n                })\n    return dets\n\ndef eval_coco(cocoGt, dets):\n    # [FIX] Một số bản pycocotools yêu cầu 'info' trong GT dataset\n    if not hasattr(cocoGt, 'dataset') or 'info' not in cocoGt.dataset:\n        cocoGt.dataset['info'] = {'description': 'GT built from YOLO labels', 'version': '1.0'}\n\n    cocoDt = cocoGt.loadRes(dets)\n    E = COCOeval(cocoGt, cocoDt, iouType='bbox')\n    E.params.iouThrs = np.linspace(0.5, 0.95, 10)\n    E.params.maxDets = [100, 300, 300]\n\n    # (không bắt buộc nhưng nên có) đảm bảo catIds/imgIds khớp\n    E.params.catIds = [c['id'] for c in cocoGt.dataset['categories']]\n    E.params.imgIds = [im['id'] for im in cocoGt.dataset['images']]\n\n    E.evaluate(); E.accumulate(); E.summarize()\n\n    stats = {\n        'mAP_50_95': float(E.stats[0]),\n        'mAP_50':    float(E.stats[1]),\n        'mAP_75':    float(E.stats[2]),\n        'AR_1':      float(E.stats[6]),\n        'AR_10':     float(E.stats[7]),\n        'AR_100':    float(E.stats[8]),\n    }\n    # AP theo lớp\n    precisions = E.eval['precision']  # TxRxKxAxM\n    ap_per_class = {}\n    for idx, catId in enumerate(E.params.catIds):\n        p = precisions[:, :, idx, 0, -1]\n        p = p[p > -1]\n        ap_per_class[catId] = float(np.mean(p)) if p.size else float('nan')\n    return E, stats, ap_per_class\n\n\nteacher_dets = predict_to_coco_dets(teacher_path, val_imgs, img_id_map, imgsz=IMG_SIZE, iou=0.55)\nE_t, stats_t, apcls_t = eval_coco(cocoGt, teacher_dets)\n\nif student_ckpt is not None:\n    student_dets = predict_to_coco_dets(student_ckpt, val_imgs, img_id_map, imgsz=IMG_SIZE, iou=0.55)\n    E_s, stats_s, apcls_s = eval_coco(cocoGt, student_dets)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_ap_bar(ap_t, ap_s, class_names, out_path):\n    x = np.arange(len(class_names))\n    t = [ap_t.get(i+1, np.nan) for i in range(len(class_names))]\n    s = [ap_s.get(i+1, np.nan) for i in range(len(class_names))]\n\n    plt.figure(figsize=(10, 5))\n    width = 0.38\n    plt.bar(x - width/2, t, width, label='Teacher')\n    plt.bar(x + width/2, s, width, label='Student')\n    plt.xticks(x, class_names, rotation=0)\n    plt.ylim(0, 1.0)\n    plt.ylabel('AP (0.5:0.95)')\n    plt.title('Per-class AP — Teacher vs Student')\n    plt.legend()\n    plt.tight_layout()\n    plt.savefig(out_path, dpi=200, bbox_inches='tight')\n    plt.show()\n\nif student_ckpt is not None:\n    out_png = os.path.join(METRICS_DIR, 'ap_per_class_teacher_vs_student.png')\n    plot_ap_bar(apcls_t, apcls_s, NAMES, out_png)\n    print(\"Saved:\", out_png)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_pr_overlay(E_t, E_s, class_names, iou=0.5, out_dir=METRICS_DIR):\n    iou_thrs = E_t.params.iouThrs\n    i = int(np.argmin(np.abs(iou_thrs - iou)))\n    rec = E_t.params.recThrs\n    for k, name in enumerate(class_names):\n        p_t = E_t.eval['precision'][i, :, k, 0, -1]\n        p_s = E_s.eval['precision'][i, :, k, 0, -1]\n        p_t = np.where(p_t < 0, np.nan, p_t)\n        p_s = np.where(p_s < 0, np.nan, p_s)\n\n        plt.figure(figsize=(5.6, 5))\n        plt.plot(rec, p_t, label='Teacher')\n        plt.plot(rec, p_s, label='Student')\n        plt.xlabel('Recall')\n        plt.ylabel('Precision')\n        plt.title(f'PR Curve (IoU={iou:.2f}) — {name}')\n        plt.xlim(0, 1); plt.ylim(0, 1)\n        plt.legend()\n        plt.grid(True, linestyle='--', linewidth=0.5, alpha=0.5)\n        fn = os.path.join(out_dir, f'PR_{name}_IoU{iou:.2f}.png')\n        plt.tight_layout()\n        plt.savefig(fn, dpi=200, bbox_inches='tight')\n        plt.show()\n        print(\"Saved:\", fn)\n\nif student_ckpt is not None:\n    plot_pr_overlay(E_t, E_s, NAMES, iou=0.5, out_dir=METRICS_DIR)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_round_ckpts(root=FLKD_OUTDIR):\n    out = []\n    for rdir in sorted(glob.glob(os.path.join(root, 'student_round_*'))):\n        name = Path(rdir).name\n        best = os.path.join(rdir, 'weights', 'best.pt')\n        last = os.path.join(rdir, 'weights', 'last.pt')\n        if os.path.exists(best):\n            out.append((name, best))\n        elif os.path.exists(last):\n            out.append((name, last))\n    return out\n\ndef map_over_rounds(round_ckpts, img_paths, img_id_map, cocoGt):\n    rows = []\n    for tag, ckpt in round_ckpts:\n        dets = predict_to_coco_dets(ckpt, img_paths, img_id_map, imgsz=IMG_SIZE, iou=0.55)\n        _, stats, _ = eval_coco(cocoGt, dets)\n        rows.append({'round': tag, **stats})\n    df = pd.DataFrame(rows)\n    df['round_idx'] = df['round'].str.extract(r'(\\d+)').astype(int)\n    return df.sort_values('round_idx')\n\nround_ckpts = find_round_ckpts(FLKD_OUTDIR)\nif round_ckpts:\n    df_round = map_over_rounds(round_ckpts, val_imgs, img_id_map, cocoGt)\n    csv_path = os.path.join(METRICS_DIR, 'mAP_over_rounds.csv')\n    df_round.to_csv(csv_path, index=False)\n    print(\"Saved:\", csv_path)\n    # Vẽ mAP@[.5:.95] theo vòng\n    plt.figure(figsize=(6.2, 4.2))\n    plt.plot(df_round['round_idx'], df_round['mAP_50_95'], marker='o')\n    plt.xlabel('Round')\n    plt.ylabel('mAP (0.5:0.95)')\n    plt.title('Student mAP vs. FL+KD Round')\n    plt.grid(True, linestyle='--', linewidth=0.5, alpha=0.5)\n    plt.tight_layout()\n    out_png = os.path.join(METRICS_DIR, 'mAP_vs_round.png')\n    plt.savefig(out_png, dpi=200, bbox_inches='tight')\n    plt.show()\n    print(\"Saved:\", out_png)\nelse:\n    print(\"Không tìm thấy thư mục student_round_* trong\", FLKD_OUTDIR)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def results_to_norm_lists(res, W, H):\n    if res.boxes is None or len(res.boxes) == 0:\n        return [], [], []\n    xyxy = res.boxes.xyxy.detach().cpu().numpy()\n    conf = res.boxes.conf.detach().cpu().numpy()\n    cls  = res.boxes.cls.detach().cpu().numpy().astype(int)\n    boxes = [[x1/W, y1/H, x2/W, y2/H] for (x1,y1,x2,y2) in xyxy]\n    return boxes, conf.tolist(), cls.tolist()\n\ndef fuse_ts_wbf_for_image(res_t, res_s, alpha, W, H, iou_thr=0.55):\n    b1, s1, l1 = results_to_norm_lists(res_t, W, H)\n    b2, s2, l2 = results_to_norm_lists(res_s, W, H)\n\n    boxes_list, scores_list, labels_list, weights = [], [], [], []\n    if len(b1) > 0:\n        boxes_list.append(b1); scores_list.append(s1); labels_list.append(l1); weights.append(alpha)\n    if len(b2) > 0:\n        boxes_list.append(b2); scores_list.append(s2); labels_list.append(l2); weights.append(1 - alpha)\n\n    if len(boxes_list) == 0:\n        return []\n\n    f_boxes, f_scores, f_labels = weighted_boxes_fusion(\n        boxes_list, scores_list, labels_list,\n        weights=weights, iou_thr=iou_thr, skip_box_thr=0.001\n    )\n    dets = []\n    for (x1n, y1n, x2n, y2n), sc, lb in zip(f_boxes, f_scores, f_labels):\n        x1 = float(x1n * W); y1 = float(y1n * H)\n        w  = float((x2n - x1n) * W); h = float((y2n - y1n) * H)\n        dets.append({'bbox': [x1, y1, w, h], 'score': float(sc), 'category_id': int(lb) + 1})\n    return dets\n\ndef eval_wbf_alpha(teacher_path, student_path, alpha_list, img_paths, img_id_map, cocoGt, batch=16):\n    dets_per_alpha = {a: [] for a in alpha_list}\n    mT = YOLO(teacher_path); mS = YOLO(student_path)\n    if hasattr(mT, 'model'): mT.model.to(torch.device('cuda' if torch.cuda.is_available() else 'cpu'))\n    if hasattr(mS, 'model'): mS.model.to(torch.device('cuda' if torch.cuda.is_available() else 'cpu'))\n\n    for i in range(0, len(img_paths), batch):\n        batch_paths = img_paths[i:i+batch]\n        res_t = mT(batch_paths, imgsz=IMG_SIZE, conf=0.001, iou=0.55,\n                   device=(0 if torch.cuda.is_available() else 'cpu'), verbose=False)\n        res_s = mS(batch_paths, imgsz=IMG_SIZE, conf=0.001, iou=0.55,\n                   device=(0 if torch.cuda.is_available() else 'cpu'), verbose=False)\n        for p, rt, rs in zip(batch_paths, res_t, res_s):\n            H, W = rt.orig_shape\n            for a in alpha_list:\n                fused = fuse_ts_wbf_for_image(rt, rs, a, W, H, iou_thr=0.55)\n                for d in fused:\n                    det = d.copy()\n                    det['image_id'] = img_id_map[p]\n                    dets_per_alpha[a].append(det)\n\n    rows = []\n    for a in alpha_list:\n        E, stats, _ = eval_coco(cocoGt, dets_per_alpha[a])\n        rows.append({'alpha': float(a), **stats})\n    df = pd.DataFrame(rows).sort_values('alpha', ascending=False)\n    return df\n\nif student_ckpt is not None:\n    alphas = [1.0, 0.9, 0.7, 0.5, 0.3, 0.1, 0.0]\n    df_alpha = eval_wbf_alpha(teacher_path, student_ckpt, alphas, val_imgs, img_id_map, cocoGt, batch=16)\n    csv_path = os.path.join(METRICS_DIR, 'mAP_vs_alpha_TplusS.csv')\n    df_alpha.to_csv(csv_path, index=False)\n    print(\"Saved:\", csv_path)\n\n    plt.figure(figsize=(6.2, 4.2))\n    plt.plot(df_alpha['alpha'], df_alpha['mAP_50_95'], marker='o')\n    plt.xlabel('α (trọng số cho Teacher trong WBF)')\n    plt.ylabel('mAP (0.5:0.95)')\n    plt.title('WBF(T, S): mAP vs α')\n    plt.grid(True, linestyle='--', linewidth=0.5, alpha=0.5)\n    plt.tight_layout()\n    out_png = os.path.join(METRICS_DIR, 'mAP_vs_alpha_TplusS.png')\n    plt.savefig(out_png, dpi=200, bbox_inches='tight')\n    plt.show()\n    print(\"Saved:\", out_png)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rows = [{'model': 'Teacher', **stats_t}]\nif student_ckpt is not None:\n    rows.append({'model': 'Student (final)', **stats_s})\n\nsummary_df = pd.DataFrame(rows)\nsummary_csv = os.path.join(METRICS_DIR, 'summary_metrics.csv')\nsummary_df.to_csv(summary_csv, index=False)\nprint(\"Summary CSV:\", summary_csv)\nsummary_df\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}