{"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":"none","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":1462296,"sourceType":"datasetVersion","datasetId":857191},{"sourceId":604441,"sourceType":"modelInstanceVersion","modelInstanceId":403514,"modelId":421446}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-10-12T05:45:05.448275Z","iopub.execute_input":"2025-10-12T05:45:05.448609Z","iopub.status.idle":"2025-10-12T05:45:05.453194Z","shell.execute_reply.started":"2025-10-12T05:45:05.448584Z","shell.execute_reply":"2025-10-12T05:45:05.452624Z"}},"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-10-12T05:45:05.454575Z","iopub.execute_input":"2025-10-12T05:45:05.454876Z","iopub.status.idle":"2025-10-12T05:46:19.21365Z","shell.execute_reply.started":"2025-10-12T05:45:05.454859Z","shell.execute_reply":"2025-10-12T05:46:19.21297Z"}},"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-10-12T05:46:19.214584Z","iopub.execute_input":"2025-10-12T05:46:19.214886Z","iopub.status.idle":"2025-10-12T05:46:23.851705Z","shell.execute_reply.started":"2025-10-12T05:46:19.214846Z","shell.execute_reply":"2025-10-12T05:46:23.850956Z"}},"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-10-12T05:46:23.852545Z","iopub.execute_input":"2025-10-12T05:46:23.852925Z","iopub.status.idle":"2025-10-12T05:46:23.935695Z","shell.execute_reply.started":"2025-10-12T05:46:23.852894Z","shell.execute_reply":"2025-10-12T05:46:23.935046Z"}},"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-10-12T05:46:23.937819Z","iopub.execute_input":"2025-10-12T05:46:23.938276Z","iopub.status.idle":"2025-10-12T05:46:26.210503Z","shell.execute_reply.started":"2025-10-12T05:46:23.938256Z","shell.execute_reply":"2025-10-12T05:46:26.20978Z"}},"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-10-12T05:46:26.211409Z","iopub.execute_input":"2025-10-12T05:46:26.212081Z","iopub.status.idle":"2025-10-12T05:46:26.231654Z","shell.execute_reply.started":"2025-10-12T05:46:26.212058Z","shell.execute_reply":"2025-10-12T05:46:26.231054Z"}},"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-10-12T05:46:26.232425Z","iopub.execute_input":"2025-10-12T05:46:26.232669Z","iopub.status.idle":"2025-10-12T05:46:26.250115Z","shell.execute_reply.started":"2025-10-12T05:46:26.232648Z","shell.execute_reply":"2025-10-12T05:46:26.249518Z"}},"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-10-12T05:46:26.250747Z","iopub.execute_input":"2025-10-12T05:46:26.251006Z","iopub.status.idle":"2025-10-12T05:46:26.265333Z","shell.execute_reply.started":"2025-10-12T05:46:26.250988Z","shell.execute_reply":"2025-10-12T05:46:26.264844Z"}},"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            last_ckpt = os.path.join(self.out_dir, f\"student_round_{r+1}\", \"weights\", \"last.pt\")\n            if os.path.exists(last_ckpt):\n                self.student = YOLO(last_ckpt)\n            else:\n                raise FileNotFoundError(f\"Checkpoint not found: {last_ckpt}\")\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-10-12T05:46:26.26592Z","iopub.execute_input":"2025-10-12T05:46:26.266112Z","iopub.status.idle":"2025-10-12T05:46:26.290058Z","shell.execute_reply.started":"2025-10-12T05:46:26.266097Z","shell.execute_reply":"2025-10-12T05:46:26.289462Z"}},"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-10-12T05:46:26.290841Z","iopub.execute_input":"2025-10-12T05:46:26.291482Z","iopub.status.idle":"2025-10-12T05:46:26.310626Z","shell.execute_reply.started":"2025-10-12T05:46:26.29146Z","shell.execute_reply":"2025-10-12T05:46:26.309977Z"}},"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-10-12T05:46:26.311372Z","iopub.execute_input":"2025-10-12T05:46:26.31163Z","iopub.status.idle":"2025-10-12T05:46:26.327253Z","shell.execute_reply.started":"2025-10-12T05:46:26.311609Z","shell.execute_reply":"2025-10-12T05:46:26.326595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T05:46:26.327987Z","iopub.execute_input":"2025-10-12T05:46:26.32829Z","iopub.status.idle":"2025-10-12T06:05:35.787713Z","shell.execute_reply.started":"2025-10-12T05:46:26.328268Z","shell.execute_reply":"2025-10-12T06:05:35.786734Z"}},"outputs":[],"execution_count":null}]}