{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":[{"sourceId":128792,"databundleVersionId":15494745,"sourceType":"competition"}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n================================================================================\n🏆 VISTA CODEFEST'26 - GRANDMASTER PIPELINE 🏆\n================================================================================\nOPTIMIZED FOR KAGGLE P100 FREE TIER (~6-7 HOURS)\n\nGRANDMASTER FIXES:\n✅ Trust COUNT MODEL more than YOLO (count model trained on real data)\n✅ ZERO HALLUCINATION - Never invent categories\n✅ Layout-aware synthetic generation\n✅ Class-aware arbitration\n✅ Reduced epochs (6+4+8 = 18 total vs 28)\n✅ CRASH RESILIENT - Atomic saves, model reuse, no deletion\n\n================================================================================\n\"\"\"\n\nimport subprocess, sys, pickle, time, zipfile\n\nfor pkg in [\"ultralytics\", \"timm\"]:\n    try: __import__(pkg.split(\"[\")[0])\n    except ImportError: subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pkg])\n\nimport os, json, shutil, yaml, glob, gc, cv2, random, warnings, math\nfrom pathlib import Path\nfrom collections import defaultdict, Counter\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport torchvision.transforms as T\nfrom tqdm.auto import tqdm\n\ntry:\n    import timm\n    HAS_TIMM = True\nexcept: HAS_TIMM = False\n\nwarnings.filterwarnings('ignore')\n\n# =============================================================================\n# CONFIGURATION - ELITE TIER + REDUCED TIME\n# =============================================================================\n\nclass Config:\n    ROOT_DIR = '/kaggle/input/vista26'\n    BASE_DIR = '/kaggle/input/vista26/Vistas Dataset Public/Vistas Dataset Public'\n    WORK_DIR = '/kaggle/working'\n    \n    USE_TEST_LABELS = True\n    \n    # Count Model - REDUCED EPOCHS\n    COUNT_BATCH_SIZE = 16\n    COUNT_LR = 1.5e-4\n    COUNT_IMGSZ = 288\n    MAX_COUNT = 50\n    \n    # REDUCED: 6+4 = 10 epochs total\n    ENSEMBLE_BACKBONES = ['efficientnet_b0', 'mobilenetv3_large_100']\n    ENSEMBLE_EPOCHS = [4, 3]\n    \n    # Trust COUNT MODEL more\n    MODEL_COUNT_WEIGHT = 0.65\n    YOLO_COUNT_WEIGHT = 0.35\n    \n    # Layout-aware synthetic\n    USE_SYNTHETIC = True\n    SYNTHETIC_COUNT = 150\n    SHELF_ROWS = [200, 450, 700, 950, 1200, 1450]\n    \n    # YOLO - REDUCED EPOCHS\n    YOLO_MODEL = 'yolov8s.pt'\n    YOLO_EPOCHS = 5\n    YOLO_IMGSZ = 640\n    YOLO_BATCH_SIZE = 8\n    YOLO_FREEZE_LAYERS = 0\n    \n    # Detection\n    CONF_THRESHOLD = 0.12\n    IOU_THRESHOLD = 0.45\n    MAX_DETECTIONS = 100\n    DEDUP_IOU_THRESHOLD = 0.5\n    \n    # TTA\n    USE_TTA = True\n    \n    # Time limits\n    MAX_TOTAL_TIME = 25000\n    MAX_YOLO_TIME = 10800\n    MAX_COUNT_TIME = 5400\n    \n    WORKERS = 2\n    SPLIT_RATIO = 0.9\n    SEED = 42\n\ncfg = Config()\ncfg.TRAIN_DIR = f\"{cfg.BASE_DIR}/train\"\ncfg.TEST_DIR = f\"{cfg.BASE_DIR}/test\"\ncfg.VAL_DIR = f\"{cfg.BASE_DIR}/validation\"\ncfg.BG_DIR = f\"{cfg.BASE_DIR}/background\"\ncfg.TRAIN_JSON = f\"{cfg.BASE_DIR}/instances_train.json\"\ncfg.TEST_JSON = f\"{cfg.BASE_DIR}/instances_test.json\"\ncfg.CATEGORIES_JSON = f\"{cfg.BASE_DIR}/Categories.json\"\ncfg.VAL_JSON = f\"{cfg.ROOT_DIR}/instances_val.json\"\ncfg.YOLO_DIR = f\"{cfg.WORK_DIR}/yolo_data\"\ncfg.SYNTHETIC_DIR = f\"{cfg.WORK_DIR}/synthetic\"\n\nrandom.seed(cfg.SEED)\nnp.random.seed(cfg.SEED)\ntorch.manual_seed(cfg.SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(cfg.SEED)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nPIPELINE_START = time.time()\n\n# =============================================================================\n# UTILITIES\n# =============================================================================\n\ndef cleanup():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\ndef safe_json(path):\n    try:\n        with open(path, 'r') as f:\n            data = json.load(f)\n        if 'root' in data and isinstance(data['root'], dict):\n            return data['root']\n        return data\n    except Exception as e:\n        print(f\"Error reading {path}: {e}\")\n        return None\n\ndef safe_imread(path, rgb=False):\n    try:\n        img = cv2.imread(path)\n        if img is None:\n            return None\n        return cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if rgb else img\n    except:\n        return None\n\ndef time_remaining():\n    return cfg.MAX_TOTAL_TIME - (time.time() - PIPELINE_START)\n\n# =============================================================================\n# 🔥 FIX 4: ATOMIC CHECKPOINT SYSTEM\n# =============================================================================\n\nclass Checkpoint:\n    def __init__(self):\n        self.file = f\"{cfg.WORK_DIR}/checkpoint.pkl\"\n        self.data = self._load()\n    \n    def _load(self):\n        if os.path.exists(self.file):\n            try:\n                with open(self.file, 'rb') as f:\n                    data = pickle.load(f)\n                print(f\"✅ Checkpoint loaded: Stage {data.get('stage', 0)}\")\n                return data\n            except:\n                # Try backup\n                backup = self.file + \".bak\"\n                if os.path.exists(backup):\n                    try:\n                        with open(backup, 'rb') as f:\n                            data = pickle.load(f)\n                        print(f\"✅ Checkpoint restored from backup: Stage {data.get('stage', 0)}\")\n                        return data\n                    except:\n                        pass\n        return {'stage': 0, 'data': {}}\n    \n    def save(self):\n        # 🔥 FIX 4: Atomic save - prevents corruption on disconnect\n        tmp = self.file + \".tmp\"\n        backup = self.file + \".bak\"\n        \n        with open(tmp, \"wb\") as f:\n            pickle.dump(self.data, f)\n        \n        # Backup existing checkpoint\n        if os.path.exists(self.file):\n            try:\n                shutil.copy2(self.file, backup)\n            except:\n                pass\n        \n        # Atomic replace\n        os.replace(tmp, self.file)\n    \n    def set(self, key, value):\n        self.data['data'][key] = value\n        self.save()\n    \n    def get(self, key, default=None):\n        return self.data['data'].get(key, default)\n    \n    def stage_done(self, stage):\n        return self.data.get('stage', 0) > stage\n    \n    def complete_stage(self, stage):\n        self.data['stage'] = stage + 1\n        self.save()\n    \n    def clear(self):\n        # 🔥 FIX 1: NEVER actually clear - keep for reuse\n        print(\"🛡️ Preserving checkpoint for potential reuse\")\n\nckpt = Checkpoint()\n\n# =============================================================================\n# HEADER\n# =============================================================================\n\nprint(\"=\" * 70)\nprint(\"🏆 VISTA CODEFEST'26 - GRANDMASTER PIPELINE\")\nprint(\"=\" * 70)\nprint(f\"Device: {DEVICE}\")\nprint(f\"⏱️  Target time: ~6-7 hours (18 total epochs)\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n\n# =============================================================================\n# STAGE 0: DATA LOADING\n# =============================================================================\n\nif not ckpt.stage_done(0):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"📍 STAGE 0: DATA LOADING\")\n    print(\"=\" * 70)\n    \n    val_data = safe_json(cfg.VAL_JSON)\n    if not val_data:\n        raise ValueError(\"Cannot read validation JSON!\")\n    \n    VAL_INFO = {}\n    VAL_IDS = []\n    for img in val_data.get('images', []):\n        img_id = int(img['id'])\n        VAL_IDS.append(img_id)\n        VAL_INFO[img_id] = {\n            'id': img_id,\n            'file_name': img['file_name'],\n            'width': img.get('width', 1800),\n            'height': img.get('height', 1800),\n            'level': img.get('level', 'medium'),\n            'path': f\"{cfg.VAL_DIR}/{img['file_name']}\"\n        }\n    VAL_IDS = sorted(VAL_IDS)\n    print(f\"✅ Validation: {len(VAL_IDS)} images\")\n    \n    cat_data = safe_json(cfg.CATEGORIES_JSON)\n    if not cat_data:\n        raise ValueError(\"Cannot read categories!\")\n    \n    cat_list = cat_data.get('categories', [])\n    categories = sorted(cat_list, key=lambda x: int(x['id']))\n    \n    coco_ids = [int(c['id']) for c in categories]\n    COCO_ID_SET = set(coco_ids)\n    coco_to_yolo = {cid: idx for idx, cid in enumerate(coco_ids)}\n    yolo_to_coco = {idx: cid for idx, cid in enumerate(coco_ids)}\n    class_names = [str(c.get('name', c.get('supercategory', f'class_{c[\"id\"]}'))) for c in categories]\n    NUM_CLASSES = len(categories)\n    \n    print(f\"✅ Categories: {NUM_CLASSES}\")\n    \n    def load_vista_json(json_path, check_categories=True):\n        data = safe_json(json_path)\n        if not data:\n            return {}, {}\n        \n        images = data.get('images', [])\n        img_dict = {}\n        ann_dict = {}\n        \n        for img in images:\n            img_id = img.get('id')\n            if img_id is None:\n                continue\n            \n            img_dict[img_id] = {\n                'id': img_id,\n                'file_name': img.get('file_name', ''),\n                'width': img.get('width', 1800),\n                'height': img.get('height', 1800),\n                'level': img.get('level', 'medium')\n            }\n            \n            ann_dict[img_id] = []\n            for ann in img.get('annotations', []):\n                cat_id = ann.get('category_id')\n                bbox = ann.get('bbox')\n                if cat_id is None or bbox is None:\n                    continue\n                if check_categories and cat_id not in COCO_ID_SET:\n                    continue\n                ann_dict[img_id].append({\n                    'category_id': int(cat_id),\n                    'bbox': bbox\n                })\n        \n        return img_dict, ann_dict\n    \n    train_img_dict, train_ann_dict = load_vista_json(cfg.TRAIN_JSON)\n    print(f\"✅ Train: {len(train_img_dict)} images (single-object)\")\n    \n    test_img_dict, test_ann_dict = load_vista_json(cfg.TEST_JSON)\n    print(f\"✅ Test: {len(test_img_dict)} images (multi-object)\")\n    \n    # Analyze\n    count_dist = Counter()\n    cat_freq = Counter()\n    level_dist = Counter()\n    \n    for img_id, anns in test_ann_dict.items():\n        count_dist[len(anns)] += 1\n        level = test_img_dict[img_id].get('level', 'medium')\n        level_dist[level] += 1\n        for ann in anns:\n            cat_freq[ann['category_id']] += 1\n    \n    RARE_CATEGORIES = set(c for c, cnt in cat_freq.items() if cnt < 15)\n    TOP_CATEGORIES = [c for c, _ in cat_freq.most_common(50)]\n    if not TOP_CATEGORIES:\n        TOP_CATEGORIES = coco_ids[:50]\n    \n    print(f\"📊 Rare categories (< 15 occurrences): {len(RARE_CATEGORIES)}\")\n    \n    background_paths = []\n    if os.path.exists(cfg.BG_DIR):\n        background_paths = [f for f in glob.glob(f\"{cfg.BG_DIR}/*\") \n                          if f.lower().endswith(('.jpg', '.jpeg', '.png'))]\n    print(f\"✅ Backgrounds: {len(background_paths)}\")\n    \n    by_level = defaultdict(list)\n    for img_id, info in test_img_dict.items():\n        by_level[info.get('level', 'medium')].append(img_id)\n    \n    TRAIN_IDS = set()\n    VAL_HOLD_IDS = set()\n    \n    for level, ids in by_level.items():\n        random.shuffle(ids)\n        split = int(len(ids) * cfg.SPLIT_RATIO)\n        TRAIN_IDS.update(ids[:split])\n        VAL_HOLD_IDS.update(ids[split:])\n    \n    print(f\"✅ Split: {len(TRAIN_IDS)} train, {len(VAL_HOLD_IDS)} holdout\")\n    \n    ckpt.set('VAL_INFO', VAL_INFO)\n    ckpt.set('VAL_IDS', VAL_IDS)\n    ckpt.set('coco_ids', coco_ids)\n    ckpt.set('COCO_ID_SET', COCO_ID_SET)\n    ckpt.set('coco_to_yolo', coco_to_yolo)\n    ckpt.set('yolo_to_coco', yolo_to_coco)\n    ckpt.set('class_names', class_names)\n    ckpt.set('NUM_CLASSES', NUM_CLASSES)\n    ckpt.set('train_img_dict', train_img_dict)\n    ckpt.set('train_ann_dict', train_ann_dict)\n    ckpt.set('test_img_dict', test_img_dict)\n    ckpt.set('test_ann_dict', test_ann_dict)\n    ckpt.set('TOP_CATEGORIES', TOP_CATEGORIES)\n    ckpt.set('RARE_CATEGORIES', RARE_CATEGORIES)\n    ckpt.set('TRAIN_IDS', TRAIN_IDS)\n    ckpt.set('VAL_HOLD_IDS', VAL_HOLD_IDS)\n    ckpt.set('count_dist', dict(count_dist))\n    ckpt.set('background_paths', background_paths)\n    \n    ckpt.complete_stage(0)\n    cleanup()\nelse:\n    print(\"\\n📍 Stage 0: Loading from checkpoint...\")\n    VAL_INFO = ckpt.get('VAL_INFO')\n    VAL_IDS = ckpt.get('VAL_IDS')\n    coco_ids = ckpt.get('coco_ids')\n    COCO_ID_SET = ckpt.get('COCO_ID_SET')\n    coco_to_yolo = ckpt.get('coco_to_yolo')\n    yolo_to_coco = ckpt.get('yolo_to_coco')\n    class_names = ckpt.get('class_names')\n    NUM_CLASSES = ckpt.get('NUM_CLASSES')\n    train_img_dict = ckpt.get('train_img_dict')\n    train_ann_dict = ckpt.get('train_ann_dict')\n    test_img_dict = ckpt.get('test_img_dict')\n    test_ann_dict = ckpt.get('test_ann_dict')\n    TOP_CATEGORIES = ckpt.get('TOP_CATEGORIES')\n    RARE_CATEGORIES = ckpt.get('RARE_CATEGORIES', set())\n    TRAIN_IDS = ckpt.get('TRAIN_IDS')\n    VAL_HOLD_IDS = ckpt.get('VAL_HOLD_IDS')\n    count_dist = ckpt.get('count_dist')\n    background_paths = ckpt.get('background_paths', [])\n\n# =============================================================================\n# STAGE 1: LAYOUT-AWARE SYNTHETIC GENERATION\n# =============================================================================\n\nif not ckpt.stage_done(1):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"📍 STAGE 1: LAYOUT-AWARE SYNTHETIC GENERATION\")\n    print(\"=\" * 70)\n    \n    os.makedirs(cfg.SYNTHETIC_DIR, exist_ok=True)\n    \n    object_crops = defaultdict(list)\n    for img_id in list(train_img_dict.keys())[:400]:\n        info = train_img_dict[img_id]\n        img = safe_imread(f\"{cfg.TRAIN_DIR}/{info['file_name']}\")\n        if img is None:\n            continue\n        H, W = img.shape[:2]\n        \n        for ann in train_ann_dict.get(img_id, []):\n            cat_id = ann['category_id']\n            if cat_id not in coco_to_yolo or len(object_crops[cat_id]) >= 40:\n                continue\n            \n            x, y, w, h = ann['bbox']\n            x1, y1 = max(0, int(x)), max(0, int(y))\n            x2, y2 = min(W, int(x + w)), min(H, int(y + h))\n            \n            if x2 - x1 < 30 or y2 - y1 < 30:\n                continue\n            \n            object_crops[cat_id].append(img[y1:y2, x1:x2].copy())\n    \n    print(f\"   ✅ Collected crops from {len(object_crops)} categories\")\n    \n    synthetic_samples = []\n    synthetic_yolo_data = []\n    \n    if background_paths and object_crops:\n        available_cats = list(object_crops.keys())\n        \n        for syn_idx in tqdm(range(cfg.SYNTHETIC_COUNT), desc=\"Synthetic\"):\n            bg = safe_imread(random.choice(background_paths))\n            if bg is None:\n                continue\n            \n            canvas = cv2.resize(bg, (1800, 1800))\n            H, W = 1800, 1800\n            \n            target_count = random.choices(\n                list(count_dist.keys()),\n                weights=list(count_dist.values()),\n                k=1\n            )[0] if count_dist else random.randint(3, 15)\n            \n            placed_boxes = []\n            labels = []\n            \n            for _ in range(target_count * 2):\n                if len(placed_boxes) >= target_count:\n                    break\n                \n                cat_id = random.choice(available_cats)\n                if not object_crops.get(cat_id):\n                    continue\n                \n                crop = random.choice(object_crops[cat_id]).copy()\n                \n                scale = random.uniform(0.4, 1.0)\n                new_w = min(int(crop.shape[1] * scale), W // 5)\n                new_h = min(int(crop.shape[0] * scale), H // 5)\n                if new_w < 25 or new_h < 25:\n                    continue\n                \n                crop_resized = cv2.resize(crop, (new_w, new_h))\n                \n                row_y = random.choice(cfg.SHELF_ROWS)\n                y = row_y + random.randint(-40, 40)\n                x = random.randint(30, W - new_w - 30)\n                \n                y = max(0, min(H - new_h, y))\n                \n                overlap = False\n                for bx, by, bw, bh in placed_boxes:\n                    if not (x + new_w < bx or x > bx + bw or y + new_h < by or y > by + bh):\n                        inter = max(0, min(x+new_w, bx+bw) - max(x, bx)) * max(0, min(y+new_h, by+bh) - max(y, by))\n                        if inter > 0.3 * new_w * new_h:\n                            overlap = True\n                            break\n                \n                if overlap:\n                    continue\n                \n                alpha = random.uniform(0.9, 1.0)\n                roi = canvas[y:y+new_h, x:x+new_w]\n                canvas[y:y+new_h, x:x+new_w] = cv2.addWeighted(crop_resized, alpha, roi, 1-alpha, 0)\n                \n                placed_boxes.append((x, y, new_w, new_h))\n                cx, cy = (x + new_w/2) / W, (y + new_h/2) / H\n                labels.append(f\"{coco_to_yolo[cat_id]} {cx:.6f} {cy:.6f} {new_w/W:.6f} {new_h/H:.6f}\")\n            \n            if len(placed_boxes) >= 1:\n                syn_path = f\"{cfg.SYNTHETIC_DIR}/syn_{syn_idx:04d}.jpg\"\n                cv2.imwrite(syn_path, canvas)\n                synthetic_samples.append((syn_path, len(placed_boxes)))\n                synthetic_yolo_data.append((syn_path, labels))\n    \n    print(f\"   ✅ Generated {len(synthetic_samples)} synthetic images\")\n    \n    ckpt.set('synthetic_samples', synthetic_samples)\n    ckpt.set('synthetic_yolo_data', synthetic_yolo_data)\n    ckpt.complete_stage(1)\n    cleanup()\nelse:\n    print(\"\\n📍 Stage 1: Loading from checkpoint...\")\n    synthetic_samples = ckpt.get('synthetic_samples', [])\n    synthetic_yolo_data = ckpt.get('synthetic_yolo_data', [])\n    \n    # 🔥 FIX 5: Verify synthetic files exist on load\n    synthetic_samples = [s for s in synthetic_samples if os.path.exists(s[0])]\n    synthetic_yolo_data = [s for s in synthetic_yolo_data if os.path.exists(s[0])]\n    print(f\"   ✅ Verified {len(synthetic_samples)} synthetic images exist\")\n\n# =============================================================================\n# STAGE 2: COUNT MODEL (REDUCED EPOCHS)\n# =============================================================================\n\nclass HybridCountModel(nn.Module):\n    def __init__(self, backbone='efficientnet_b0', max_count=50):\n        super().__init__()\n        self.max_count = max_count\n\n        if HAS_TIMM:\n            self.backbone = timm.create_model(backbone, pretrained=True, num_classes=0)\n        else:\n            import torchvision.models as models\n            resnet = models.resnet34(weights='IMAGENET1K_V1')\n            self.backbone = nn.Sequential(*list(resnet.children())[:-1])\n\n        # 🔥 DYNAMIC FEATURE SIZE DETECTION\n        with torch.no_grad():\n            dummy = torch.zeros(1, 3, cfg.COUNT_IMGSZ, cfg.COUNT_IMGSZ)\n            feat = self.backbone(dummy)\n            if len(feat.shape) > 2:\n                feat = feat.flatten(1)\n            feat_dim = feat.shape[1]\n\n        self.cls_head = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(feat_dim, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, max_count + 1)\n        )\n\n        self.reg_head = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(feat_dim, 128),\n            nn.ReLU(),\n            nn.Linear(128, 1)\n        )\n\n    def forward(self, x):\n        feat = self.backbone(x)\n        if len(feat.shape) > 2:\n            feat = feat.flatten(1)\n        return self.cls_head(feat), self.reg_head(feat).squeeze(-1)\n\n    def predict(self, x):\n        cls_out, reg_out = self.forward(x)\n        cls_probs = torch.softmax(cls_out, dim=1)\n        cls_pred = cls_probs.argmax(dim=1)\n        cls_conf = cls_probs.max(dim=1).values\n        reg_pred = torch.clamp(reg_out, 0, self.max_count).round().long()\n\n        final_pred = torch.where(\n            cls_conf > 0.5,\n            cls_pred,\n            ((cls_pred.float() + reg_pred.float()) / 2).round().long()\n        )\n        return final_pred, cls_conf, cls_probs\n\n    \n\nclass CountDataset(Dataset):\n    def __init__(self, samples, transform, max_count=50):\n        self.samples = [(p, min(c, max_count)) for p, c in samples]\n        self.transform = transform\n        self.max_count = max_count\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        path, count = self.samples[idx]\n        img = safe_imread(path, rgb=True)\n        if img is None:\n            img = np.zeros((cfg.COUNT_IMGSZ, cfg.COUNT_IMGSZ, 3), dtype=np.uint8)\n        return self.transform(img), count\n\ntrain_tf = T.Compose([\n    T.ToPILImage(),\n    T.Resize((cfg.COUNT_IMGSZ, cfg.COUNT_IMGSZ)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(p=0.2),\n    T.RandomRotation(10),\n    T.ColorJitter(0.2, 0.2, 0.15, 0.05),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\nval_tf = T.Compose([\n    T.ToPILImage(),\n    T.Resize((cfg.COUNT_IMGSZ, cfg.COUNT_IMGSZ)),\n    T.ToTensor(),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\nif not ckpt.stage_done(2):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"📍 STAGE 2: COUNT MODEL TRAINING (6+4 epochs)\")\n    print(\"=\" * 70)\n    \n    train_samples = []\n    val_samples = []\n    \n    for img_id in TRAIN_IDS:\n        if img_id in test_img_dict:\n            path = f\"{cfg.TEST_DIR}/{test_img_dict[img_id]['file_name']}\"\n            count = len(test_ann_dict.get(img_id, []))\n            if os.path.exists(path):\n                train_samples.append((path, count))\n    \n    for img_id in VAL_HOLD_IDS:\n        if img_id in test_img_dict:\n            path = f\"{cfg.TEST_DIR}/{test_img_dict[img_id]['file_name']}\"\n            count = len(test_ann_dict.get(img_id, []))\n            if os.path.exists(path):\n                val_samples.append((path, count))\n    \n    for img_id in list(train_img_dict.keys())[:200]:\n        path = f\"{cfg.TRAIN_DIR}/{train_img_dict[img_id]['file_name']}\"\n        if os.path.exists(path):\n            train_samples.append((path, 1))\n    \n    train_samples.extend(synthetic_samples)\n    random.shuffle(train_samples)\n    \n    print(f\"📊 Count samples: {len(train_samples)} train, {len(val_samples)} val\")\n    \n    if not val_samples:\n        val_samples = train_samples[:50]\n    \n    train_loader = DataLoader(\n        CountDataset(train_samples, train_tf, cfg.MAX_COUNT),\n        batch_size=cfg.COUNT_BATCH_SIZE, shuffle=True, num_workers=cfg.WORKERS\n    )\n    val_loader = DataLoader(\n        CountDataset(val_samples, val_tf, cfg.MAX_COUNT),\n        batch_size=cfg.COUNT_BATCH_SIZE, num_workers=cfg.WORKERS\n    )\n    \n    count_freq = Counter([s[1] for s in train_samples])\n    weights = torch.ones(cfg.MAX_COUNT + 1)\n    total = sum(count_freq.values())\n    for c, freq in count_freq.items():\n        if c <= cfg.MAX_COUNT:\n            weights[c] = math.sqrt(total / freq)\n    weights = weights.to(DEVICE)\n    \n    cls_criterion = nn.CrossEntropyLoss(weight=weights, label_smoothing=0.1)\n    reg_criterion = nn.SmoothL1Loss()\n    \n    model_paths = []\n    \n    for idx, (backbone, epochs) in enumerate(zip(cfg.ENSEMBLE_BACKBONES, cfg.ENSEMBLE_EPOCHS)):\n        model_path = f\"{cfg.WORK_DIR}/count_{idx}_{backbone}.pt\"\n        \n        # 🔥 FIX 2: Skip training if model already exists\n        if os.path.exists(model_path):\n            print(f\"\\n♻️ Reusing existing count model: {model_path}\")\n            model_paths.append(model_path)\n            continue\n        \n        print(f\"\\n🔥 Training count model {idx+1}: {backbone} ({epochs} epochs)\")\n        \n        if time_remaining() < cfg.MAX_COUNT_TIME:\n            print(\"⚠️ Time limit - skipping\")\n            break\n        \n        model = HybridCountModel(backbone, cfg.MAX_COUNT).to(DEVICE)\n        optimizer = optim.AdamW(model.parameters(), lr=cfg.COUNT_LR, weight_decay=1e-4)\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs)\n        scaler = GradScaler()\n        \n        best_acc = 0\n        \n        for epoch in range(epochs):\n            model.train()\n            epoch_loss = 0\n            \n            for imgs, labels in tqdm(train_loader, desc=f\"Ep{epoch+1}\", leave=False):\n                imgs = imgs.to(DEVICE)\n                labels = labels.to(DEVICE)\n                \n                optimizer.zero_grad()\n                \n                with autocast():\n                    cls_out, reg_out = model(imgs)\n                    cls_loss = cls_criterion(cls_out, labels)\n                    reg_loss = reg_criterion(reg_out, labels.float())\n                    loss = cls_loss + 0.5 * reg_loss\n                \n                if torch.isnan(loss):\n                    continue\n                \n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n                scaler.step(optimizer)\n                scaler.update()\n                epoch_loss += loss.item()\n            \n            scheduler.step()\n            \n            model.eval()\n            correct, total_n = 0, 0\n            with torch.no_grad():\n                for imgs, labels in val_loader:\n                    imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                    preds, _, _ = model.predict(imgs)\n                    correct += (preds == labels).sum().item()\n                    total_n += labels.size(0)\n            \n            acc = correct / max(1, total_n)\n            print(f\"   Ep{epoch+1}: Loss={epoch_loss/len(train_loader):.4f}, Acc={acc:.4f}\")\n            \n            if acc > best_acc:\n                best_acc = acc\n                torch.save(model.state_dict(), model_path)\n        \n        print(f\"   ✅ Best: {best_acc:.4f}\")\n        model_paths.append(model_path)\n        del model, optimizer, scheduler, scaler\n        cleanup()\n    \n    ckpt.set('count_model_paths', model_paths)\n    ckpt.complete_stage(2)\n    del train_loader, val_loader\n    cleanup()\nelse:\n    print(\"\\n📍 Stage 2: Loading from checkpoint...\")\n\n# =============================================================================\n# STAGE 3: YOLO TRAINING (8 EPOCHS)\n# =============================================================================\n\nif not ckpt.stage_done(3):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"📍 STAGE 3: YOLO TRAINING (8 epochs)\")\n    print(\"=\" * 70)\n    \n    # 🔥 FIX 3: Skip YOLO training if model exists\n    existing_yolo = glob.glob(f\"{cfg.WORK_DIR}/**/best.pt\", recursive=True)\n    if existing_yolo:\n        yolo_path = sorted(existing_yolo, key=os.path.getmtime)[-1]\n        print(f\"♻️ Reusing existing YOLO model: {yolo_path}\")\n        ckpt.set('yolo_path', yolo_path)\n        ckpt.complete_stage(3)\n    elif time_remaining() < cfg.MAX_YOLO_TIME:\n        print(\"⚠️ Insufficient time\")\n        ckpt.set('yolo_path', cfg.YOLO_MODEL)\n        ckpt.complete_stage(3)\n    else:\n        shutil.rmtree(cfg.YOLO_DIR, ignore_errors=True)\n        for split in ['train', 'val']:\n            os.makedirs(f\"{cfg.YOLO_DIR}/images/{split}\", exist_ok=True)\n            os.makedirs(f\"{cfg.YOLO_DIR}/labels/{split}\", exist_ok=True)\n        \n        def create_yolo_labels(img_dict, ann_dict, src_dir, split, id_filter=None):\n            count = 0\n            for img_id, info in img_dict.items():\n                if id_filter and img_id not in id_filter:\n                    continue\n                src = f\"{src_dir}/{info['file_name']}\"\n                if not os.path.exists(src):\n                    continue\n                \n                W, H = info.get('width', 1800), info.get('height', 1800)\n                name = f\"{img_id}_{Path(info['file_name']).stem}\"\n                dst = f\"{cfg.YOLO_DIR}/images/{split}/{name}.jpg\"\n                \n                try:\n                    os.symlink(src, dst)\n                except:\n                    try:\n                        shutil.copy2(src, dst)\n                    except:\n                        continue\n                \n                labels = []\n                for ann in ann_dict.get(img_id, []):\n                    cid = ann['category_id']\n                    if cid not in coco_to_yolo:\n                        continue\n                    x, y, w, h = ann['bbox']\n                    cx = max(0.001, min(0.999, (x + w/2) / W))\n                    cy = max(0.001, min(0.999, (y + h/2) / H))\n                    nw = max(0.001, min(0.999, w / W))\n                    nh = max(0.001, min(0.999, h / H))\n                    labels.append(f\"{coco_to_yolo[cid]} {cx:.6f} {cy:.6f} {nw:.6f} {nh:.6f}\")\n                \n                with open(f\"{cfg.YOLO_DIR}/labels/{split}/{name}.txt\", 'w') as f:\n                    f.write('\\n'.join(labels))\n                count += 1\n            return count\n        \n        n1 = create_yolo_labels(test_img_dict, test_ann_dict, cfg.TEST_DIR, 'train', TRAIN_IDS)\n        n2 = create_yolo_labels(test_img_dict, test_ann_dict, cfg.TEST_DIR, 'val', VAL_HOLD_IDS)\n        n3 = create_yolo_labels(train_img_dict, train_ann_dict, cfg.TRAIN_DIR, 'train')\n        \n        for syn_path, labels in synthetic_yolo_data:\n            if os.path.exists(syn_path):\n                name = Path(syn_path).stem\n                try:\n                    os.symlink(syn_path, f\"{cfg.YOLO_DIR}/images/train/{name}.jpg\")\n                except:\n                    pass\n                with open(f\"{cfg.YOLO_DIR}/labels/train/{name}.txt\", 'w') as f:\n                    f.write('\\n'.join(labels))\n        \n        print(f\"✅ YOLO data: {n1 + n3 + len(synthetic_yolo_data)} train, {n2} val\")\n        \n        with open(f\"{cfg.WORK_DIR}/dataset.yaml\", 'w') as f:\n            yaml.dump({\n                'path': cfg.WORK_DIR,\n                'train': 'yolo_data/images/train',\n                'val': 'yolo_data/images/val',\n                'nc': NUM_CLASSES,\n                'names': class_names\n            }, f)\n        \n        from ultralytics import YOLO\n        \n        try:\n            model = YOLO(cfg.YOLO_MODEL)\n            model.train(\n                data=f\"{cfg.WORK_DIR}/dataset.yaml\",\n                epochs=cfg.YOLO_EPOCHS,\n                imgsz=cfg.YOLO_IMGSZ,\n                batch=cfg.YOLO_BATCH_SIZE,\n                device=DEVICE,\n                workers=cfg.WORKERS,\n                project=cfg.WORK_DIR,\n                name='yolo',\n                exist_ok=True,\n                patience=4,\n                amp=True,\n                verbose=False,\n                freeze=cfg.YOLO_FREEZE_LAYERS\n            )\n            \n            best = glob.glob(f\"{cfg.WORK_DIR}/**/best.pt\", recursive=True)\n            yolo_path = sorted(best, key=os.path.getmtime)[-1] if best else cfg.YOLO_MODEL\n            del model\n        except Exception as e:\n            print(f\"⚠️ YOLO failed: {e}\")\n            yolo_path = cfg.YOLO_MODEL\n        \n        print(f\"✅ YOLO: {yolo_path}\")\n        ckpt.set('yolo_path', yolo_path)\n        ckpt.complete_stage(3)\n        cleanup()\nelse:\n    print(\"\\n📍 Stage 3: Loading from checkpoint...\")\n\n# =============================================================================\n# STAGE 4: INFERENCE WITH ELITE ARBITRATION\n# =============================================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"📍 STAGE 4: INFERENCE (ELITE)\")\nprint(\"=\" * 70)\n\nfrom ultralytics import YOLO\n\ncount_paths = ckpt.get('count_model_paths', [])\nyolo_path = ckpt.get('yolo_path', cfg.YOLO_MODEL)\n\ncount_models = []\nfor path in count_paths:\n    if os.path.exists(path):\n        backbone = 'efficientnet_b0' if 'b0' in path else 'mobilenetv3_large_100'\n        model = HybridCountModel(backbone, cfg.MAX_COUNT).to(DEVICE)\n        model.load_state_dict(torch.load(path, weights_only=True))\n        model.eval()\n        count_models.append(model)\n\nprint(f\"✅ Loaded {len(count_models)} count models\")\n\nyolo = YOLO(yolo_path)\nprint(f\"✅ Loaded YOLO: {yolo_path}\")\n\ndef predict_count_ensemble(img_np, models, transform):\n    all_preds, all_confs = [], []\n    \n    with torch.no_grad():\n        img_t = transform(img_np).unsqueeze(0).to(DEVICE)\n        for m in models:\n            pred, conf, _ = m.predict(img_t)\n            all_preds.append(pred.item())\n            all_confs.append(conf.item())\n        \n        if cfg.USE_TTA:\n            img_flip = np.fliplr(img_np).copy()\n            img_t = transform(img_flip).unsqueeze(0).to(DEVICE)\n            for m in models:\n                pred, conf, _ = m.predict(img_t)\n                all_preds.append(pred.item())\n                all_confs.append(conf.item())\n    \n    if sum(all_confs) > 0:\n        weighted = sum(p * c for p, c in zip(all_preds, all_confs)) / sum(all_confs)\n    else:\n        weighted = np.mean(all_preds)\n    \n    return int(round(weighted)), np.mean(all_confs)\n\ndef get_yolo_detections(yolo_model, img_path):\n    try:\n        results = yolo_model.predict(\n            img_path, conf=cfg.CONF_THRESHOLD, iou=cfg.IOU_THRESHOLD,\n            verbose=False, device=DEVICE, max_det=cfg.MAX_DETECTIONS\n        )\n        \n        if not results or results[0].boxes is None:\n            return []\n        \n        detections = []\n        boxes = results[0].boxes\n        \n        for cls, conf, box in zip(\n            boxes.cls.cpu().numpy().astype(int),\n            boxes.conf.cpu().numpy(),\n            boxes.xyxy.cpu().numpy()\n        ):\n            if cls in yolo_to_coco:\n                detections.append({\n                    'category': yolo_to_coco[cls],\n                    'score': float(conf),\n                    'box': box.tolist()\n                })\n        \n        return sorted(detections, key=lambda x: -x['score'])\n    except:\n        return []\n\ndef nms_by_category(detections, iou_threshold=0.5):\n    if not detections:\n        return []\n    \n    by_cat = defaultdict(list)\n    for d in detections:\n        by_cat[d['category']].append(d)\n    \n    final = []\n    for cat, dets in by_cat.items():\n        dets = sorted(dets, key=lambda x: -x['score'])\n        keep = []\n        \n        for d in dets:\n            overlap = False\n            for k in keep:\n                b1, b2 = d['box'], k['box']\n                x1, y1 = max(b1[0], b2[0]), max(b1[1], b2[1])\n                x2, y2 = min(b1[2], b2[2]), min(b1[3], b2[3])\n                inter = max(0, x2 - x1) * max(0, y2 - y1)\n                area1 = (b1[2] - b1[0]) * (b1[3] - b1[1])\n                area2 = (b2[2] - b2[0]) * (b2[3] - b2[1])\n                union = area1 + area2 - inter\n                if union > 0 and inter / union > iou_threshold:\n                    overlap = True\n                    break\n            if not overlap:\n                keep.append(d)\n        final.extend(keep)\n    \n    return sorted(final, key=lambda x: -x['score'])\n\ndef select_categories_zero_hallucination(detections, target_count):\n    if target_count == 0:\n        return []\n    \n    if not detections:\n        return []\n    \n    result = [d['category'] for d in detections[:target_count]]\n    \n    if len(result) < target_count and detections:\n        while len(result) < target_count:\n            for d in detections:\n                if len(result) >= target_count:\n                    break\n                result.append(d['category'])\n    \n    return sorted(result[:target_count])\n\ndef class_aware_arbitration(model_count, model_conf, yolo_count, detections):\n    if not detections:\n        return model_count\n    \n    detected_cats = set(d['category'] for d in detections)\n    has_rare = bool(detected_cats & RARE_CATEGORIES)\n    \n    if has_rare:\n        weight_yolo = 0.55\n        weight_model = 0.45\n    else:\n        weight_yolo = cfg.YOLO_COUNT_WEIGHT\n        weight_model = cfg.MODEL_COUNT_WEIGHT\n    \n    if model_conf > 0.7:\n        return model_count\n    \n    return int(round(weight_model * model_count + weight_yolo * yolo_count))\n\nresults = {}\nstats = Counter()\n\nfor img_id in tqdm(VAL_IDS, desc=\"Inference\"):\n    try:\n        info = VAL_INFO[img_id]\n        img_path = info['path']\n        \n        if not os.path.exists(img_path):\n            results[img_id] = []\n            stats['missing'] += 1\n            continue\n        \n        img = safe_imread(img_path, rgb=True)\n        if img is None:\n            results[img_id] = []\n            stats['read_error'] += 1\n            continue\n        \n        if count_models:\n            model_count, model_conf = predict_count_ensemble(img, count_models, val_tf)\n        else:\n            model_count, model_conf = 5, 0.0\n        \n        detections = get_yolo_detections(yolo, img_path)\n        detections = nms_by_category(detections, cfg.DEDUP_IOU_THRESHOLD)\n        yolo_count = len(detections)\n        \n        final_count = class_aware_arbitration(model_count, model_conf, yolo_count, detections)\n        final_count = max(0, min(final_count, cfg.MAX_COUNT))\n        \n        categories = select_categories_zero_hallucination(detections, final_count)\n        results[img_id] = categories\n        stats['ok'] += 1\n        \n    except Exception as e:\n        results[img_id] = []\n        stats['error'] += 1\n\nfor img_id in VAL_IDS:\n    if img_id not in results:\n        results[img_id] = []\n\nprint(f\"\\n📊 Stats: {dict(stats)}\")\n\ndel yolo\nfor m in count_models:\n    del m\ncleanup()\n\n# =============================================================================\n# STAGE 5: SUBMISSION\n# =============================================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"📍 STAGE 5: SUBMISSION\")\nprint(\"=\" * 70)\n\nrows = []\nfor img_id in sorted(VAL_IDS):\n    cats = results.get(img_id, [])\n    valid_cats = sorted([int(c) for c in cats if c in COCO_ID_SET])\n    \n    rows.append({\n        'image_id': int(img_id),\n        'categories': json.dumps(valid_cats)\n    })\n\ndf = pd.DataFrame(rows)\ndf = df.sort_values('image_id').reset_index(drop=True)\n\nerrors = []\nif not df['image_id'].is_unique:\n    errors.append(\"Duplicate image_ids!\")\nif len(df) != len(VAL_IDS):\n    errors.append(f\"Count mismatch: {len(df)} vs {len(VAL_IDS)}\")\nif set(df['image_id']) != set(VAL_IDS):\n    errors.append(\"ID mismatch!\")\n\nfor _, row in df.iterrows():\n    cats = json.loads(row['categories'])\n    if cats != sorted(cats):\n        errors.append(f\"Not sorted: {row['image_id']}\")\n        break\n    for c in cats:\n        if c not in COCO_ID_SET:\n            errors.append(f\"Invalid category: {c}\")\n            break\n\nif errors:\n    print(f\"❌ Errors: {errors}\")\nelse:\n    print(\"✅ Validation passed!\")\n\nsubmission_path = f\"{cfg.WORK_DIR}/submission.csv\"\ndf.to_csv(submission_path, index=False)\nprint(f\"✅ Saved: {submission_path}\")\n\ntotal_objects = sum(len(json.loads(c)) for c in df['categories'])\nempty = sum(1 for c in df['categories'] if c == '[]')\ncount_dist_pred = Counter(len(json.loads(c)) for c in df['categories'])\n\nprint(f\"\"\"\n{'─' * 60}\n📊 SUBMISSION STATISTICS\n{'─' * 60}\n   Images: {len(df)}\n   Total objects: {total_objects}\n   Avg/image: {total_objects / len(df):.2f}\n   Empty predictions: {empty}\n\"\"\")\n\nelapsed = time.time() - PIPELINE_START\nprint(f\"\"\"\n{'─' * 60}\n⏱️  Time: {elapsed/60:.1f}min ({elapsed/3600:.2f}h)\n{'─' * 60}\n\"\"\")\n\n# 🔥 FIX 1: NEVER DELETE MODELS - Preserve for reuse\nprint(\"🛡️ Preserving models and checkpoint for reuse\")\n# shutil.rmtree(cfg.YOLO_DIR, ignore_errors=True)  # DISABLED\n# shutil.rmtree(cfg.SYNTHETIC_DIR, ignore_errors=True)  # DISABLED\n# shutil.rmtree(f\"{cfg.WORK_DIR}/yolo\", ignore_errors=True)  # DISABLED\n# ckpt.clear()  # DISABLED\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"🏆 GRANDMASTER PIPELINE COMPLETE!\")\nprint(\"=\" * 70)\nprint(\"\"\"\n✅ GRANDMASTER FIXES APPLIED:\n\n   🔥 Trust COUNT MODEL more (trained on real data)\n   🔥 ZERO HALLUCINATION - Never invent categories\n   🔥 Layout-aware synthetic (shelf rows)\n   🔥 Class-aware arbitration (rare vs common)\n   🔥 Reduced epochs: 6+4+8 = 18 (~6-7 hours)\n   \n   🛡️ CRASH RESILIENT:\n   ✅ FIX 1: Never delete models - preserved for reuse\n   ✅ FIX 2: Skip count training if model exists\n   ✅ FIX 3: Skip YOLO training if best.pt exists\n   ✅ FIX 4: Atomic checkpoint saves (no corruption)\n   ✅ FIX 5: Verify synthetic files exist on reload\n\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-29T17:37:40.210775Z","iopub.execute_input":"2026-01-29T17:37:40.211810Z","iopub.status.idle":"2026-01-29T19:02:49.303720Z","shell.execute_reply.started":"2026-01-29T17:37:40.211761Z","shell.execute_reply":"2026-01-29T19:02:49.303040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}