{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":125981,"databundleVersionId":14910697,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports, Paths and Params","metadata":{}},{"cell_type":"code","source":"# Imports\nimport os\nfrom pathlib import Path\nimport torch\nimport json\nimport math\nimport heapq\nfrom collections import deque\nfrom typing import List, Tuple, Dict\nimport random\nimport numpy as np\nfrom multiprocessing import Pool, cpu_count\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torchvision import transforms\nimport torchvision.transforms as T\nfrom torchvision.utils import save_image\nfrom scipy.ndimage import gaussian_filter, map_coordinates, affine_transform\nfrom collections import Counter\nfrom PIL import Image, ImageEnhance, ImageDraw, ImageFilter, ImageOps\nfrom tqdm import tqdm\nimport pandas as pd","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:47:39.308139Z","iopub.execute_input":"2025-12-27T04:47:39.308390Z","iopub.status.idle":"2025-12-27T04:47:48.118098Z","shell.execute_reply.started":"2025-12-27T04:47:39.308358Z","shell.execute_reply":"2025-12-27T04:47:48.117544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Paths\npaths = {\n    \"train_images\": Path(\"/kaggle/input/the-blind-flight-synapse-drive-ps-1/SynapseDrive_Dataset/train/images\"),\n    \"train_labels\": Path(\"/kaggle/input/the-blind-flight-synapse-drive-ps-1/SynapseDrive_Dataset/train/labels\"),\n    \"test_images\": Path(\"/kaggle/input/the-blind-flight-synapse-drive-ps-1/SynapseDrive_Dataset/test/images\"),\n    \"test_velocities\": Path(\"/kaggle/input/the-blind-flight-synapse-drive-ps-1/SynapseDrive_Dataset/test/velocities\"),\n    \"mod\": Path(\"/kaggle/working/augmented\"),\n    \"saved_images\": Path(\"/kaggle/working/generated_sample/images\"),\n    \"saved_labels\": Path(\"/generated_sample/labels\"),\n    \"mod_images\": Path(\"/kaggle/working/augmented/images\"),\n    \"mod_labels\": Path(\"/kaggle/working/augmented/labels\"),\n    \"submission_path\":Path(\"/kaggle/working/submission_baseline.csv\"),\n    \"assets\": Path(\"/kaggle/input/the-blind-flight-synapse-drive-ps-1/SynapseDrive_Dataset/assets\"),\n    'mod_train_images': Path('/kaggle/working/test_data/images'),\n    'mod_train_labels': Path('/kaggle/working/test_data/labels')\n}\n\n# Make necessary dirs\npaths['saved_images'].mkdir(parents=True, exist_ok=True)\npaths['saved_labels'].mkdir(parents=True, exist_ok=True)\npaths['mod_images'].mkdir(parents = True, exist_ok = True)\npaths['mod_labels'].mkdir(parents = True, exist_ok = True)\npaths['mod_train_images'].mkdir(parents = True, exist_ok = True)\npaths['mod_train_labels'].mkdir(parents = True, exist_ok = True)\n\n# device\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-27T04:47:48.119729Z","iopub.execute_input":"2025-12-27T04:47:48.120396Z","iopub.status.idle":"2025-12-27T04:47:48.195281Z","shell.execute_reply.started":"2025-12-27T04:47:48.120370Z","shell.execute_reply":"2025-12-27T04:47:48.194540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Parameters\nNUM_TRAIN_SAMPLES = 20\nGRID_SIZE = 20\nNUM_CLASSES = 5  # 0..4\nBATCH_SIZE = 3\nEPOCHS = 30\nLR = 1e-4\nto_tensor = T.ToTensor()\nTARGET_TILE_SIZE = 36        # Final tile size in generated image         \nNUM_SAMPLES = 2000\nBORDER_WIDTH = 1             # 1px black borders (set to 0 for seamless like Image 1)\nPNG_COMPRESSION = 4        # Balance between size and speed\nUSE_MULTIPROCESSING = True\nNUM_WORKERS = max(1, cpu_count() - 1)\n\n# classes\nCLASS_WALK = 0\nCLASS_WALL = 1\nCLASS_HAZARD = 2\nCLASS_START = 3\nCLASS_GOAL = 4\n\nCLASS_WEIGHTS = torch.tensor(\n    [1.0, 2.0, 1.8, 3.0, 3.0],\n    dtype=torch.float32,\n    device=DEVICE\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:20.680925Z","iopub.execute_input":"2025-12-26T13:39:20.681333Z","iopub.status.idle":"2025-12-26T13:39:20.869254Z","shell.execute_reply.started":"2025-12-26T13:39:20.681312Z","shell.execute_reply":"2025-12-26T13:39:20.868668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"counter = Counter()\n\nfor json_file in paths[\"train_labels\"].glob(\"*.json\"):\n    with open(json_file, \"r\") as f:\n        data = json.load(f)\n\n        grid = np.array(data[\"grid\"])   # change key if needed\n        counter.update(grid.flatten().tolist())\n\nprint(\"Class counts and probabilites:\")\nfor cls in sorted(counter):\n    print(f\"class {cls}: {counter[cls]/NUM_TRAIN_SAMPLES} | probability: {counter[cls]/(NUM_TRAIN_SAMPLES * GRID_SIZE * GRID_SIZE)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:22.007001Z","iopub.execute_input":"2025-12-26T13:39:22.007260Z","iopub.status.idle":"2025-12-26T13:39:22.089480Z","shell.execute_reply.started":"2025-12-26T13:39:22.007240Z","shell.execute_reply":"2025-12-26T13:39:22.088763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"asset_map = {\n    \"desert\": {\n        \"0\": \"t1_sand\",\n        \"1\": {\"a\": \"t1_rocks\", \"b\": \"t1_cacti\"},\n        \"2\": \"t1_quicksand\",\n        \"3\": \"t1_rover\",\n        \"4\": \"t1_goal\",\n    },\n    \"forest\": {\n        \"0\": \"t0_dirt\",\n        \"1\": \"t0_tree\",\n        \"2\": \"t0_puddle\",\n        \"3\": \"t0_startship\",\n        \"4\": \"t0_goal\",\n    },\n    \"lab\": {\n        \"0\": \"t2_floor\",\n        \"1\": {\"a\": \"t2_wall\", \"b\": \"t2_plasma\"},\n        \"2\": \"t2_glue\",\n        \"3\": \"t2_drone\",\n        \"4\": \"t2_goal\",\n    },\n}\n\nTERRAIN_COSTS = {\n    'lab': {\n        CLASS_WALK: 1.0,\n        CLASS_WALL: float('inf'),\n        CLASS_HAZARD: 3.0,\n        CLASS_START: 1.0,\n        CLASS_GOAL: 2.0\n    },\n    'forest': {\n        CLASS_WALK: 1.5,\n        CLASS_WALL: float('inf'),\n        CLASS_HAZARD: 2.8,\n        CLASS_START: 1.5,\n        CLASS_GOAL: 2.5\n    },\n    'desert': {\n        CLASS_WALK: 1.2,\n        CLASS_WALL: float('inf'),\n        CLASS_HAZARD: 3.7,\n        CLASS_START: 1.2,\n        CLASS_GOAL: 2.2\n    }\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:23.158531Z","iopub.execute_input":"2025-12-26T13:39:23.159190Z","iopub.status.idle":"2025-12-26T13:39:23.164255Z","shell.execute_reply.started":"2025-12-26T13:39:23.159165Z","shell.execute_reply":"2025-12-26T13:39:23.163578Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Synthetic Map generation\n\n* Grid generation using dijkstra algorithm \n* load asset\n* rendering grid as image\n* Saving image and grid\n","metadata":{}},{"cell_type":"code","source":"def generate_grid(size=20, class_probs={0:0.5,1:0.28,2:0.21}, hazard_bias=0.3):\n    grid = np.full((size, size), 1, dtype=np.int8)\n\n    start = (random.randrange(size), random.randrange(size))\n    goal = (random.randrange(size), random.randrange(size))\n    while goal == start:\n        goal = (random.randrange(size), random.randrange(size))\n\n    cur = start\n    grid[cur] = 0\n\n    while cur != goal:\n        r, c = cur\n        moves = [(nr, nc) for nr, nc in\n                 [(r-1,c),(r+1,c),(r,c-1),(r,c+1)]\n                 if 0 <= nr < size and 0 <= nc < size]\n        nxt = random.choice(moves)\n        grid[nxt] = 2 if random.random() < hazard_bias else 0\n        cur = nxt\n\n    for r in range(size):\n        for c in range(size):\n            if (r, c) in (start, goal):\n                continue\n            if grid[r, c] == 1:\n                grid[r, c] = random.choices(\n                    [0, 1, 2],\n                    [class_probs[0], class_probs[1], class_probs[2]]\n                )[0]\n\n    grid[start] = 3\n    grid[goal] = 4\n    return grid, start, goal\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:25.571644Z","iopub.execute_input":"2025-12-26T13:39:25.572024Z","iopub.status.idle":"2025-12-26T13:39:25.579508Z","shell.execute_reply.started":"2025-12-26T13:39:25.571998Z","shell.execute_reply":"2025-12-26T13:39:25.578514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_assets(asset_root, asset_map, tile_size):\n    \"\"\"\n    Load assets with high-quality downsampling to preserve detail\n    \"\"\"\n    cache = {}\n    \n    for terrain, classes in asset_map.items():\n        cache[terrain] = {}\n        for val in classes.values():\n            names = val.values() if isinstance(val, dict) else [val]\n            for name in names:\n                img_path = asset_root / terrain / f\"{name}.png\"\n                \n                # Load high-res image\n                img = Image.open(img_path).convert(\"RGB\")\n                \n                # Use LANCZOS for high-quality downsampling (better than NEAREST for large->small)\n                # This preserves detail when going from 512x512 or 1024x1024 down to 32x32\n                img_resized = img.resize((tile_size, tile_size), Image.LANCZOS)\n                \n                cache[terrain][name] = to_tensor(img_resized)\n    \n    return cache\n\n\ndef choose_asset(asset_map, terrain, class_id):\n    entry = asset_map[terrain][str(class_id)]\n    if isinstance(entry, dict):\n        return random.choice(list(entry.values()))\n    return entry","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:27.031138Z","iopub.execute_input":"2025-12-26T13:39:27.031834Z","iopub.status.idle":"2025-12-26T13:39:27.037604Z","shell.execute_reply.started":"2025-12-26T13:39:27.031780Z","shell.execute_reply":"2025-12-26T13:39:27.036880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def render_grid(grid, terrain, asset_map, asset_cache, border_width=1):\n    \"\"\"\n    Render grid with optional borders\n    border_width=0 for seamless (like Image 1)\n    border_width=1 for bordered (like Images 2,3,4)\n    \"\"\"\n    grid = torch.as_tensor(grid)\n    tile = next(iter(asset_cache[terrain].values()))\n    C, H, W = tile.shape\n    \n    if border_width > 0:\n        # With borders\n        tile_h = H + border_width\n        tile_w = W + border_width\n        \n        canvas = torch.zeros((C, \n                             grid.shape[0] * tile_h + border_width,\n                             grid.shape[1] * tile_w + border_width))\n        \n        for r in range(grid.shape[0]):\n            for c in range(grid.shape[1]):\n                cls = int(grid[r, c])\n                name = choose_asset(asset_map, terrain, cls)\n                \n                y_start = r * tile_h + border_width\n                x_start = c * tile_w + border_width\n                canvas[:, y_start:y_start+H, x_start:x_start+W] = asset_cache[terrain][name]\n    else:\n        # Seamless (no borders)\n        canvas = torch.zeros((C, grid.shape[0] * H, grid.shape[1] * W))\n        \n        for r in range(grid.shape[0]):\n            for c in range(grid.shape[1]):\n                cls = int(grid[r, c])\n                name = choose_asset(asset_map, terrain, cls)\n                canvas[:, r*H:(r+1)*H, c*W:(c+1)*W] = asset_cache[terrain][name]\n    \n    return canvas","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:28.245632Z","iopub.execute_input":"2025-12-26T13:39:28.246390Z","iopub.status.idle":"2025-12-26T13:39:28.252920Z","shell.execute_reply.started":"2025-12-26T13:39:28.246358Z","shell.execute_reply":"2025-12-26T13:39:28.252274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_png_optimized(tensor, path, compress_level=6):\n    \"\"\"\n    Save PNG with optimal compression for 600-700 KB target\n    \"\"\"\n    img = (tensor.clamp(0, 1) * 255).byte()\n    img = img.permute(1, 2, 0).numpy()\n    pil = Image.fromarray(img)\n    pil.save(path, format=\"PNG\", compress_level=compress_level, optimize=True)\n\ndef generate_sample(args):\n    \"\"\"Generate a single sample (for multiprocessing)\"\"\"\n    i, asset_cache, asset_map, save_images_dir, save_labels_dir, border_width = args\n    \n    # Set seeds for reproducibility\n    random.seed(i)\n    np.random.seed(i)\n    torch.manual_seed(i)\n\n    terrain = random.choice(list(asset_map.keys()))\n    grid, start, goal = generate_grid(GRID_SIZE)\n\n    image = render_grid(grid, terrain, asset_map, asset_cache, border_width)\n\n    # Save image\n    save_png_optimized(\n        image,\n        save_images_dir / f\"sample_{i:04d}.png\",\n        compress_level=PNG_COMPRESSION\n    )\n\n    # Save label\n    label = {\n        \"id\": f\"{i:04d}\",\n        \"t\": terrain,\n        \"n\": GRID_SIZE,\n        \"s\": list(start),\n        \"g\": list(goal),\n        \"grid\": grid.tolist()\n    }\n    \n    with open(save_labels_dir / f\"sample_{i:04d}.json\", \"w\") as f:\n        json.dump(label, f, separators=(\",\", \":\"))\n    \n    return i","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:29.586509Z","iopub.execute_input":"2025-12-26T13:39:29.586858Z","iopub.status.idle":"2025-12-26T13:39:29.594520Z","shell.execute_reply.started":"2025-12-26T13:39:29.586794Z","shell.execute_reply":"2025-12-26T13:39:29.593827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Loading high-res assets and downsampling to {TARGET_TILE_SIZE}x{TARGET_TILE_SIZE}...\n    asset_cache = load_assets(paths['assets'], asset_map, TARGET_TILE_SIZE)\n    # Prepare arguments for all samples\n    args_list = [\n        (i, asset_cache, asset_map, paths['saved_images'], paths['saved_labels'], BORDER_WIDTH)\n        for i in range(NUM_SAMPLES)\n    ]\n    \n    if USE_MULTIPROCESSING and NUM_SAMPLES > 100:\n        with Pool(NUM_WORKERS) as pool:\n            for idx, result in enumerate(pool.imap_unordered(generate_sample, args_list)):\n                if idx % 100 == 0 and idx > 0:\n                    print(f\"Generated {idx}/{NUM_SAMPLES}\")\n        print(f\"Generated {NUM_SAMPLES}/{NUM_SAMPLES}\")\n    else:\n        print(f\"Generating {NUM_SAMPLES} samples (single-threaded)...\")\n        for i, args in enumerate(args_list):\n            generate_sample(args)\n            if i % 100 == 0:\n                print(f\"Generated {i}/{NUM_SAMPLES}\")\n    \n    print(\"✅ Dataset generation complete\")\n\n\nmain()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:31.814097Z","iopub.execute_input":"2025-12-26T13:39:31.814384Z","iopub.status.idle":"2025-12-26T13:39:34.855308Z","shell.execute_reply.started":"2025-12-26T13:39:31.814359Z","shell.execute_reply":"2025-12-26T13:39:34.854287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# helper function\ndef load_label_grid(json_path):\n    \"\"\"Load grid from JSON label file\"\"\"\n    with open(json_path, 'r') as f:\n        data = json.load(f)\n    return np.array(data['grid'], dtype=np.int8)\n\n\ndef save_label_grid(grid, json_path):\n    \"\"\"Save grid to JSON label file\"\"\"\n    # Load existing JSON to preserve metadata\n    base_name = json_path.stem.split('_')[0]  # Get original ID\n    original_json = json_path.parent.parent / \"labels\" / f\"{base_name}.json\"\n    \n    if original_json.exists():\n        with open(original_json, 'r') as f:\n            data = json.load(f)\n    else:\n        data = {\"id\": json_path.stem, \"n\": grid.shape[0]}\n    \n    data['grid'] = grid.tolist()\n    \n    with open(json_path, 'w') as f:\n        json.dump(data, f, separators=(\",\", \":\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:37.672451Z","iopub.execute_input":"2025-12-26T13:39:37.673225Z","iopub.status.idle":"2025-12-26T13:39:37.678919Z","shell.execute_reply.started":"2025-12-26T13:39:37.673188Z","shell.execute_reply":"2025-12-26T13:39:37.678133Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Defining Augmentation techniques","metadata":{}},{"cell_type":"code","source":"def add_color_blobs(img, num_blobs=None, intensity_range=(40, 120)):\n    \"\"\"\n    Add semi-transparent color blobs like in test images\n    \"\"\"\n    if num_blobs is None:\n        num_blobs = random.randint(2, 5)\n    \n    img = img.convert(\"RGBA\")\n    overlay = Image.new(\"RGBA\", img.size, (0, 0, 0, 0))\n    draw = ImageDraw.Draw(overlay)\n    \n    w, h = img.size\n    \n    for _ in range(num_blobs):\n        # Random position\n        cx = random.randint(-w//4, w + w//4)\n        cy = random.randint(-h//4, h + h//4)\n        \n        # Random size\n        rx = random.randint(w // 6, w // 3)\n        ry = random.randint(h // 6, h // 3)\n        \n        # Random color (vibrant tints like purple, pink, cyan, green)\n        color_choices = [\n            (200, 100, 255),  # Purple/Magenta\n            (255, 150, 200),  # Pink\n            (100, 255, 255),  # Cyan\n            (150, 255, 150),  # Light green\n            (255, 200, 100),  # Orange\n            (150, 150, 255),  # Light blue\n        ]\n        base_color = random.choice(color_choices)\n        alpha = random.randint(intensity_range[0], intensity_range[1])\n        color = (*base_color, alpha)\n        \n        # Draw ellipse\n        draw.ellipse(\n            [cx - rx, cy - ry, cx + rx, cy + ry],\n            fill=color\n        )\n    \n    # Blur the overlay for softer blobs\n    overlay = overlay.filter(ImageFilter.GaussianBlur(radius=random.randint(5, 15)))\n    \n    return Image.alpha_composite(img, overlay).convert(\"RGB\")\n\ndef elastic_deformation(img, alpha=None, sigma=None):\n    \"\"\"\n    Apply elastic deformation to simulate warping/distortion\n    \"\"\"\n    if alpha is None:\n        alpha = random.uniform(1.5, 4.0)\n    if sigma is None:\n        sigma = random.uniform(4.0, 8.0)\n    \n    img_np = np.array(img)\n    shape = img_np.shape[:2]\n    \n    # Generate random displacement fields\n    dx = gaussian_filter((np.random.rand(*shape) * 2 - 1), sigma) * alpha\n    dy = gaussian_filter((np.random.rand(*shape) * 2 - 1), sigma) * alpha\n    \n    x, y = np.meshgrid(np.arange(shape[1]), np.arange(shape[0]))\n    indices = (y + dy).reshape(-1), (x + dx).reshape(-1)\n    \n    # Apply to each channel\n    warped = np.zeros_like(img_np)\n    for c in range(3):\n        warped[..., c] = map_coordinates(\n            img_np[..., c],\n            indices,\n            order=1,\n            mode=\"reflect\"\n        ).reshape(shape)\n    \n    return Image.fromarray(warped)\n    \ndef random_crop(img, grid, crop_size=None):\n    \"\"\"\n    Randomly crop a portion of the image\n    Returns cropped image and corresponding grid portion\n    \"\"\"\n    w, h = img.size\n    grid_h, grid_w = grid.shape\n    \n    if crop_size is None:\n        # Random crop size between 50-90% of original\n        crop_ratio = random.uniform(0.5, 0.9)\n        crop_h = int(h * crop_ratio)\n        crop_w = int(w * crop_ratio)\n    else:\n        crop_h, crop_w = crop_size\n    \n    # Ensure crop size doesn't exceed image\n    crop_h = min(crop_h, h)\n    crop_w = min(crop_w, w)\n    \n    # Random crop position\n    top = random.randint(0, h - crop_h)\n    left = random.randint(0, w - crop_w)\n    \n    # Crop image\n    img_cropped = img.crop((left, top, left + crop_w, top + crop_h))\n    \n    # Calculate grid crop (assuming uniform tile size)\n    tile_h = h / grid_h\n    tile_w = w / grid_w\n    \n    grid_top = int(top / tile_h)\n    grid_left = int(left / tile_w)\n    grid_bottom = int((top + crop_h) / tile_h)\n    grid_right = int((left + crop_w) / tile_w)\n    \n    # Ensure valid grid bounds\n    grid_top = max(0, min(grid_top, grid_h - 1))\n    grid_left = max(0, min(grid_left, grid_w - 1))\n    grid_bottom = max(grid_top + 1, min(grid_bottom, grid_h))\n    grid_right = max(grid_left + 1, min(grid_right, grid_w))\n    \n    grid_cropped = grid[grid_top:grid_bottom, grid_left:grid_right]\n    \n    return img_cropped, grid_cropped\n\ndef photometric_augmentation(img):\n    \"\"\"\n    Apply random brightness, contrast, saturation, hue adjustments\n    \"\"\"\n    # Brightness\n    img = ImageEnhance.Brightness(img).enhance(random.uniform(0.85, 1.15))\n    \n    # Contrast\n    img = ImageEnhance.Contrast(img).enhance(random.uniform(0.85, 1.15))\n    \n    # Saturation\n    img = ImageEnhance.Color(img).enhance(random.uniform(0.8, 1.2))\n    \n    return img\n\ndef add_gaussian_noise(img, sigma=None):\n    \"\"\"Add Gaussian noise to image\"\"\"\n    if sigma is None:\n        sigma = random.randint(3, 12)\n    \n    arr = np.array(img).astype(np.float32)\n    noise = np.random.normal(0, sigma, arr.shape)\n    arr = np.clip(arr + noise, 0, 255).astype(np.uint8)\n    return Image.fromarray(arr)\n\n\ndef add_blur(img):\n    \"\"\"Add slight blur\"\"\"\n    blur_radius = random.uniform(0.3, 1.2)\n    return img.filter(ImageFilter.GaussianBlur(radius=blur_radius))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:39.565503Z","iopub.execute_input":"2025-12-26T13:39:39.565990Z","iopub.status.idle":"2025-12-26T13:39:39.581915Z","shell.execute_reply.started":"2025-12-26T13:39:39.565964Z","shell.execute_reply":"2025-12-26T13:39:39.581190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AdvancedAugmentation:\n    \"\"\"\n    Advanced augmentation system matching test set characteristics\n    \"\"\"\n    \n    def __init__(\n        self,\n        src_images_dir: Path,\n        src_labels_dir: Path,\n        aug_images_dir: Path,\n        aug_labels_dir: Path\n    ):\n        self.src_images_dir = src_images_dir\n        self.src_labels_dir = src_labels_dir\n        self.aug_images_dir = aug_images_dir\n        self.aug_labels_dir = aug_labels_dir\n        \n        # Load image IDs\n        self.image_ids = []\n        for p in sorted(self.src_labels_dir.glob(\"*.json\")):\n            image_id = p.stem\n            img_path = self.src_images_dir / f\"{image_id}.png\"\n            if img_path.is_file():\n                self.image_ids.append(image_id)\n        \n        if not self.image_ids:\n            raise RuntimeError(f\"No images found in {src_images_dir}\")\n        \n    def save_original(self):\n        \"\"\"Copy original data to augmented folder\"\"\"\n        for img_id in self.image_ids:\n            src_img = self.src_images_dir / f\"{img_id}.png\"\n            src_label = self.src_labels_dir / f\"{img_id}.json\"\n            \n            dst_img = self.aug_images_dir / f\"{img_id}.png\"\n            dst_label = self.aug_labels_dir / f\"{img_id}.json\"\n            \n            if not dst_img.exists():\n                Image.open(src_img).save(dst_img)\n                grid = load_label_grid(src_label)\n                save_label_grid(grid, dst_label)\n        \n        print(f\"✅ Original data saved ({len(self.image_ids)} images)\")\n    \n    \n    def color_blob_augmentation(self, num_variants=2):\n        \"\"\"\n        Apply color blob overlays (like test images)\n        \"\"\"\n        count = 0\n        \n        for img_id in self.image_ids:\n            img = Image.open(self.src_images_dir / f\"{img_id}.png\")\n            grid = load_label_grid(self.src_labels_dir / f\"{img_id}.json\")\n            \n            for v in range(num_variants):\n                aug_id = f\"{img_id}_cb{v}\"\n                aug_img_path = self.aug_images_dir / f\"{aug_id}.png\"\n                \n                if not aug_img_path.exists():\n                    # Apply color blobs with random intensity\n                    img_aug = add_color_blobs(img, num_blobs=random.randint(2, 4))\n                    \n                    # Optional: add slight photometric changes\n                    if random.random() < 0.5:\n                        img_aug = photometric_augmentation(img_aug)\n                    \n                    img_aug.save(aug_img_path)\n                    save_label_grid(grid, self.aug_labels_dir / f\"{aug_id}.json\")\n                    count += 1\n        \n        print(f\"✅ Color blob augmentation complete ({count} new images)\")\n    \n    \n    def elastic_augmentation(self, num_variants=2):\n        \"\"\"\n        Apply elastic deformation\n        \"\"\"\n        count = 0\n        \n        for img_id in self.image_ids:\n            img = Image.open(self.src_images_dir / f\"{img_id}.png\")\n            grid = load_label_grid(self.src_labels_dir / f\"{img_id}.json\")\n            \n            for v in range(num_variants):\n                aug_id = f\"{img_id}_ed{v}\"\n                aug_img_path = self.aug_images_dir / f\"{aug_id}.png\"\n                \n                if not aug_img_path.exists():\n                    img_aug = elastic_deformation(img)\n                    \n                    img_aug.save(aug_img_path)\n                    save_label_grid(grid, self.aug_labels_dir / f\"{aug_id}.json\")\n                    count += 1\n        \n        print(f\"✅ Elastic deformation complete ({count} new images)\")\n    \n    \n    def combined_augmentation(self, num_variants=3):\n        \"\"\"\n        Combine multiple augmentations (color blobs + deformation + photometric)\n        This creates images most similar to test set\n        \"\"\"\n        count = 0\n        \n        for img_id in self.image_ids:\n            img = Image.open(self.src_images_dir / f\"{img_id}.png\")\n            grid = load_label_grid(self.src_labels_dir / f\"{img_id}.json\")\n            \n            for v in range(num_variants):\n                aug_id = f\"{img_id}_comb{v}\"\n                aug_img_path = self.aug_images_dir / f\"{aug_id}.png\"\n                \n                if not aug_img_path.exists():\n                    img_aug = img.copy()\n                    \n                    # Random combination of augmentations\n                    augmentations = []\n                    \n                    # 70% chance of color blobs\n                    if random.random() < 0.7:\n                        augmentations.append('color_blob')\n                    \n                    # 50% chance of elastic deformation\n                    if random.random() < 0.5:\n                        augmentations.append('elastic')\n                    \n                    # 60% chance of photometric\n                    if random.random() < 0.6:\n                        augmentations.append('photometric')\n                    \n                    # 30% chance of noise\n                    if random.random() < 0.3:\n                        augmentations.append('noise')\n                    \n                    # 20% chance of blur\n                    if random.random() < 0.2:\n                        augmentations.append('blur')\n                    \n                    # Apply selected augmentations in random order\n                    random.shuffle(augmentations)\n                    \n                    for aug in augmentations:\n                        if aug == 'color_blob':\n                            img_aug = add_color_blobs(img_aug)\n                        elif aug == 'elastic':\n                            img_aug = elastic_deformation(img_aug)\n                        elif aug == 'photometric':\n                            img_aug = photometric_augmentation(img_aug)\n                        elif aug == 'noise':\n                            img_aug = add_gaussian_noise(img_aug)\n                        elif aug == 'blur':\n                            img_aug = add_blur(img_aug)\n                    \n                    img_aug.save(aug_img_path)\n                    save_label_grid(grid, self.aug_labels_dir / f\"{aug_id}.json\")\n                    count += 1\n        \n        print(f\"✅ Combined augmentation complete ({count} new images)\")\n    \n    \n    def crop_augmentation(self, num_variants=2, min_crop_ratio=0.6):\n        \"\"\"\n        Apply random crops\n        \"\"\"\n        count = 0\n        \n        for img_id in self.image_ids:\n            img = Image.open(self.src_images_dir / f\"{img_id}.png\")\n            grid = load_label_grid(self.src_labels_dir / f\"{img_id}.json\")\n            \n            for v in range(num_variants):\n                aug_id = f\"{img_id}_crop{v}\"\n                aug_img_path = self.aug_images_dir / f\"{aug_id}.png\"\n                \n                if not aug_img_path.exists():\n                    # Random crop\n                    crop_ratio = random.uniform(min_crop_ratio, 0.95)\n                    img_aug, grid_aug = random_crop(img, grid, crop_size=None)\n                    \n                    # Resize back to original size (optional)\n                    img_aug = img_aug.resize(img.size, Image.LANCZOS)\n                    \n                    img_aug.save(aug_img_path)\n                    save_label_grid(grid_aug, self.aug_labels_dir / f\"{aug_id}.json\")\n                    count += 1\n        \n        print(f\"✅ Crop augmentation complete ({count} new images)\")\n    \n    \n    def photometric_only(self, num_variants=2):\n        \"\"\"\n        Photometric augmentation only (brightness, contrast, saturation)\n        \"\"\"\n        count = 0\n        \n        for img_id in self.image_ids:\n            img = Image.open(self.src_images_dir / f\"{img_id}.png\")\n            grid = load_label_grid(self.src_labels_dir / f\"{img_id}.json\")\n            \n            for v in range(num_variants):\n                aug_id = f\"{img_id}_photo{v}\"\n                aug_img_path = self.aug_images_dir / f\"{aug_id}.png\"\n                \n                if not aug_img_path.exists():\n                    img_aug = photometric_augmentation(img)\n                    \n                    # Optional noise\n                    if random.random() < 0.4:\n                        img_aug = add_gaussian_noise(img_aug, sigma=random.randint(3, 8))\n                    \n                    img_aug.save(aug_img_path)\n                    save_label_grid(grid, self.aug_labels_dir / f\"{aug_id}.json\")\n                    count += 1\n        \n        print(f\"✅ Photometric augmentation complete ({count} new images)\")\n    \n    \n    def generate_all_augmentations(self):\n        \"\"\"\n        Generate all augmentation types\n        Recommended for maximum diversity\n        \"\"\"\n        self.save_original()\n        self.combined_augmentation(num_variants=1)  # Most important - matches test set\n        self.color_blob_augmentation(num_variants=1)\n        self.elastic_augmentation(num_variants=1)\n        self.photometric_only(num_variants=1)\n        self.crop_augmentation(num_variants=1)\n        \n        # Count total\n        total_images = len(list(self.aug_images_dir.glob(\"*.png\")))\n        print(f\"✅ AUGMENTATION COMPLETE\")\n        print(f\"📊 Total augmented images: {total_images}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:44.828024Z","iopub.execute_input":"2025-12-26T13:39:44.828611Z","iopub.status.idle":"2025-12-26T13:39:44.847363Z","shell.execute_reply.started":"2025-12-26T13:39:44.828585Z","shell.execute_reply":"2025-12-26T13:39:44.846696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"augmenter = AdvancedAugmentation(\n        src_images_dir=paths['saved_images'],\n        src_labels_dir=paths['saved_labels'],\n        aug_images_dir=paths['mod_images'],\n        aug_labels_dir=paths['mod_labels']\n    )\n\naugment_real_data = AdvancedAugmentation(\n        src_images_dir=paths['train_images'],\n        src_labels_dir=paths['train_labels'],\n        aug_images_dir=paths['mod_train_images'],\n        aug_labels_dir=paths['mod_train_labels']\n)\n    \naugmenter.generate_all_augmentations()   \naugment_real_data.generate_all_augmentations()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:48.346520Z","iopub.execute_input":"2025-12-26T13:39:48.346828Z","iopub.status.idle":"2025-12-26T13:39:50.636430Z","shell.execute_reply.started":"2025-12-26T13:39:48.346784Z","shell.execute_reply":"2025-12-26T13:39:50.635484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Extracting Data","metadata":{}},{"cell_type":"code","source":"# image trainsform (same dimension)\ntrain_transform = transforms.Compose([\n    transforms.Resize((GRID_SIZE, GRID_SIZE), interpolation=transforms.InterpolationMode.BILINEAR),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # ImageNet normalization\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((GRID_SIZE, GRID_SIZE), interpolation=transforms.InterpolationMode.BILINEAR),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\ndef custom_collate_fn(batch):\n    \"\"\"\n    Custom collate function to handle variable-sized grids\n    \"\"\"\n    images = []\n    grids = []\n    \n    for x, y in batch:\n        images.append(x)\n        grids.append(y)\n        \n    # stack images\n    images = torch.stack(images, dim=0)\n    \n    # Check if all grids are same size\n    grid_shapes = [g.shape for g in grids]\n    if len(set(grid_shapes)) == 1:\n        grids = torch.stack(grids, dim=0)\n    else:\n        target_size = 20\n        resized_grids = []\n        for g in grids:\n            if g.shape != (target_size, target_size):\n                g_tensor = g.unsqueeze(0).unsqueeze(0).float()\n                g_resized = F.interpolate(\n                    g_tensor,\n                    size=(target_size, target_size),\n                    mode='nearest'\n                )\n                g = g_resized.squeeze().long()\n            resized_grids.append(g)\n        grids = torch.stack(resized_grids, dim=0)\n    \n    return images, grids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:39:57.784128Z","iopub.execute_input":"2025-12-26T13:39:57.784826Z","iopub.status.idle":"2025-12-26T13:39:57.791849Z","shell.execute_reply.started":"2025-12-26T13:39:57.784780Z","shell.execute_reply":"2025-12-26T13:39:57.791186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GridDataset(Dataset):\n    def __init__(\n        self,\n        images_dir: Path,\n        labels_dir: Path,\n        grid_size: int = GRID_SIZE,\n        transform=None,\n    ):\n        self.images_dir = Path(images_dir)\n        self.labels_dir = Path(labels_dir)\n        self.grid_size = grid_size\n        self.transform = transform\n        \n        self.image_ids: List[str] = []\n        for p in sorted(labels_dir.glob(\"*.json\")):\n            image_id = p.stem\n            img_path = images_dir / f\"{image_id}.png\"\n            if img_path.is_file():\n                self.image_ids.append(image_id)\n        \n        if not self.image_ids:\n            raise RuntimeError(f\"No training labels/images found in {labels_dir}\")\n    \n    def __len__(self) -> int:\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx: int):\n        image_id = self.image_ids[idx]\n        img_path = self.images_dir / f\"{image_id}.png\"\n        label_path = self.labels_dir / f\"{image_id}.json\"\n        \n        img = Image.open(img_path).convert('RGB')\n        grid = load_label_grid(label_path)\n        \n        if self.transform:\n            x = self.transform(img)\n        else:\n            x = transforms.ToTensor()(img)\n        \n        y = torch.from_numpy(grid).long()\n        \n        return x, y\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:40:00.147980Z","iopub.execute_input":"2025-12-26T13:40:00.148488Z","iopub.status.idle":"2025-12-26T13:40:00.154857Z","shell.execute_reply.started":"2025-12-26T13:40:00.148461Z","shell.execute_reply":"2025-12-26T13:40:00.154080Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Architecture ","metadata":{}},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    \"\"\"Attention gate for U-Net\"\"\"\n    def __init__(self, F_g, F_l, F_int):\n        super().__init__()\n        self.W_g = nn.Sequential(\n            nn.Conv2d(F_g, F_int, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        \n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_l, F_int, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        \n        self.psi = nn.Sequential(\n            nn.Conv2d(F_int, 1, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n        \n        self.relu = nn.ReLU(inplace=True)\n    \n    def forward(self, g, x):\n        # g: decoder feature (gating signal)\n        # x: encoder feature (to be attention-weighted)\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        return x * psi\n\ndef conv_block(in_channels, out_channels):\n    \"\"\"Standard conv block with BatchNorm\"\"\"\n    return nn.Sequential(\n        nn.Conv2d(in_channels, out_channels, 3, padding=1),\n        nn.BatchNorm2d(out_channels),\n        nn.ReLU(inplace=True),\n        nn.Conv2d(out_channels, out_channels, 3, padding=1),\n        nn.BatchNorm2d(out_channels),\n        nn.ReLU(inplace=True),\n    )\n\n\nclass ImprovedUNet(nn.Module):\n    \"\"\"\n    Improved U-Net with optional attention gates\n    \"\"\"\n    def __init__(self, use_attention=True):\n        super().__init__()\n        self.use_attention = use_attention\n        \n        # Encoder\n        self.enc1 = conv_block(3, 64)\n        self.enc2 = conv_block(64, 128)\n        self.enc3 = conv_block(128, 256)\n        self.enc4 = conv_block(256, 512)\n        \n        self.pool = nn.MaxPool2d(2)\n        \n        # Bottleneck\n        self.bottleneck = conv_block(512, 1024)\n        \n        # Decoder\n        self.up4 = nn.ConvTranspose2d(1024, 512, 2, 2)\n        if use_attention:\n            self.att4 = AttentionBlock(F_g=512, F_l=512, F_int=256)\n        self.dec4 = conv_block(1024, 512)\n        \n        self.up3 = nn.ConvTranspose2d(512, 256, 2, 2)\n        if use_attention:\n            self.att3 = AttentionBlock(F_g=256, F_l=256, F_int=128)\n        self.dec3 = conv_block(512, 256)\n        \n        self.up2 = nn.ConvTranspose2d(256, 128, 2, 2)\n        if use_attention:\n            self.att2 = AttentionBlock(F_g=128, F_l=128, F_int=64)\n        self.dec2 = conv_block(256, 128)\n        \n        self.up1 = nn.ConvTranspose2d(128, 64, 2, 2)\n        if use_attention:\n            self.att1 = AttentionBlock(F_g=64, F_l=64, F_int=32)\n        self.dec1 = conv_block(128, 64)\n        \n        # Multi-task heads\n        self.seg_head = nn.Conv2d(64, 3, 1)  # walkable, wall, hazard\n        self.start_head = nn.Conv2d(64, 1, 1)\n        self.goal_head = nn.Conv2d(64, 1, 1)\n    \n    def forward(self, x):\n        # Encoder with size tracking\n        e1 = self.enc1(x)  # 64 channels, 20x20\n        e2 = self.enc2(self.pool(e1))  # 128 channels, 10x10\n        e3 = self.enc3(self.pool(e2))  # 256 channels, 5x5\n        e4 = self.enc4(self.pool(e3))  # 512 channels, 2x2\n        \n        # Bottleneck\n        b = self.bottleneck(self.pool(e4))  # 1024 channels, 1x1\n        \n        # Decoder with attention\n        d4 = self.up4(b)  # 512 channels\n        if d4.size() != e4.size():\n            d4 = F.interpolate(d4, size=e4.shape[2:], mode='bilinear', align_corners=False)\n        if self.use_attention:\n            e4 = self.att4(d4, e4)\n        d4 = self.dec4(torch.cat([d4, e4], dim=1))\n        \n        d3 = self.up3(d4)  # 256 channels\n        if d3.size() != e3.size():\n            d3 = F.interpolate(d3, size=e3.shape[2:], mode='bilinear', align_corners=False)\n        if self.use_attention:\n            e3 = self.att3(d3, e3)\n        d3 = self.dec3(torch.cat([d3, e3], dim=1))\n        \n        d2 = self.up2(d3)  # 128 channels\n        if d2.size() != e2.size():\n            d2 = F.interpolate(d2, size=e2.shape[2:], mode='bilinear', align_corners=False)\n        if self.use_attention:\n            e2 = self.att2(d2, e2)\n        d2 = self.dec2(torch.cat([d2, e2], dim=1))\n        \n        d1 = self.up1(d2)  # 64 channels\n        if d1.size() != e1.size():\n            d1 = F.interpolate(d1, size=e1.shape[2:], mode='bilinear', align_corners=False)\n        if self.use_attention:\n            e1 = self.att1(d1, e1)\n        d1 = self.dec1(torch.cat([d1, e1], dim=1))\n        \n        # Heads\n        seg_logits = self.seg_head(d1)\n        start_logits = self.start_head(d1)\n        goal_logits = self.goal_head(d1)\n        \n        return seg_logits, start_logits, goal_logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:40:02.875791Z","iopub.execute_input":"2025-12-26T13:40:02.876102Z","iopub.status.idle":"2025-12-26T13:40:02.892109Z","shell.execute_reply.started":"2025-12-26T13:40:02.876078Z","shell.execute_reply":"2025-12-26T13:40:02.891354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# loss\nclass DiceLoss(nn.Module):\n    \"\"\"Dice loss for better handling of class imbalance\"\"\"\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n    \n    def forward(self, logits, targets):\n        probs = F.softmax(logits, dim=1)\n        num_classes = logits.size(1)\n        \n        # One-hot encode targets\n        targets_one_hot = F.one_hot(targets, num_classes).permute(0, 3, 1, 2).float()\n        \n        # Calculate dice for each class\n        dice_scores = []\n        for c in range(num_classes):\n            pred_c = probs[:, c]\n            target_c = targets_one_hot[:, c]\n            \n            intersection = (pred_c * target_c).sum()\n            union = pred_c.sum() + target_c.sum()\n            \n            dice = (2.0 * intersection + self.smooth) / (union + self.smooth)\n            dice_scores.append(dice)\n        \n        return 1.0 - torch.stack(dice_scores).mean()\n\nclass CombinedLoss(nn.Module):\n    \"\"\"Combined CE + Dice loss\"\"\"\n    def __init__(self, weight=None, alpha=0.5):\n        super().__init__()\n        self.ce = nn.CrossEntropyLoss(weight=weight, label_smoothing=0.05)\n        self.dice = DiceLoss()\n        self.alpha = alpha\n    \n    def forward(self, logits, targets):\n        ce_loss = self.ce(logits, targets)\n        dice_loss = self.dice(logits, targets)\n        return self.alpha * ce_loss + (1 - self.alpha) * dice_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:40:05.708494Z","iopub.execute_input":"2025-12-26T13:40:05.709035Z","iopub.status.idle":"2025-12-26T13:40:05.716167Z","shell.execute_reply.started":"2025-12-26T13:40:05.709003Z","shell.execute_reply":"2025-12-26T13:40:05.715465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(\n    model,\n    train_loader,\n    val_loader,\n    epochs,\n    lr=1e-3,\n    grad_clip=5.0,\n    save_path=\"best_model.pth\",\n    eval_every=5\n):\n    model.to(DEVICE)\n    \n    # Improved loss functions with Dice\n    SEG_CLASS_WEIGHTS = torch.tensor([1.0, 3.0, 2.5], device=DEVICE)\n    seg_criterion = CombinedLoss(weight=SEG_CLASS_WEIGHTS, alpha=0.5)\n    \n    start_criterion = nn.BCEWithLogitsLoss(\n        pos_weight=torch.tensor([400.0], device=DEVICE)\n    )\n    goal_criterion = nn.BCEWithLogitsLoss(\n        pos_weight=torch.tensor([400.0], device=DEVICE)\n    )\n    \n    # AdamW optimizer (better for GPU)\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=lr,\n        weight_decay=1e-4,\n        betas=(0.9, 0.999)\n    )\n    \n    # Cosine annealing scheduler\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=10, T_mult=2, eta_min=1e-6\n    )\n    \n    # Mixed precision training for GPU\n    scaler = torch.cuda.amp.GradScaler() if DEVICE.type == 'cuda' else None\n    \n    best_val_loss = float('inf')\n    \n    print(f\"🚀 Starting training on {DEVICE}\")\n    print(f\"📊 Train batches: {len(train_loader)}, Val batches: {len(val_loader)}\")\n    print(\"=\"*60)\n    \n    for epoch in range(1, epochs + 1):\n        # Training phase\n        model.train()\n        total_loss = 0.0\n        seg_loss_sum = 0.0\n        start_loss_sum = 0.0\n        goal_loss_sum = 0.0\n        \n        # Progress bar for training\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch}/{epochs}\")\n        \n        for x, y in pbar:\n            x = x.to(DEVICE)\n            y = y.to(DEVICE)\n            \n            # Prepare targets\n            seg_target = y.clone()\n            seg_target[seg_target >= 3] = 0\n            \n            start_target = (y == CLASS_START).float().unsqueeze(1)\n            goal_target = (y == CLASS_GOAL).float().unsqueeze(1)\n            \n            optimizer.zero_grad()\n            \n            # Forward pass with mixed precision (GPU only)\n            if scaler is not None:\n                with torch.cuda.amp.autocast():\n                    seg_logits, start_logits, goal_logits = model(x)\n                    \n                    # Ensure size match\n                    if seg_logits.shape[2:] != seg_target.shape[1:]:\n                        seg_logits = F.interpolate(\n                            seg_logits, size=seg_target.shape[1:],\n                            mode='bilinear', align_corners=False\n                        )\n                        start_logits = F.interpolate(\n                            start_logits, size=start_target.shape[2:],\n                            mode='bilinear', align_corners=False\n                        )\n                        goal_logits = F.interpolate(\n                            goal_logits, size=goal_target.shape[2:],\n                            mode='bilinear', align_corners=False\n                        )\n                    \n                    seg_loss = seg_criterion(seg_logits, seg_target)\n                    start_loss = start_criterion(start_logits, start_target)\n                    goal_loss = goal_criterion(goal_logits, goal_target)\n                    \n                    loss = seg_loss + 0.5 * start_loss + 0.5 * goal_loss\n                \n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                \n                if grad_clip:\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n                \n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                # CPU fallback\n                seg_logits, start_logits, goal_logits = model(x)\n                \n                if seg_logits.shape[2:] != seg_target.shape[1:]:\n                    seg_logits = F.interpolate(\n                        seg_logits, size=seg_target.shape[1:],\n                        mode='bilinear', align_corners=False\n                    )\n                    start_logits = F.interpolate(\n                        start_logits, size=start_target.shape[2:],\n                        mode='bilinear', align_corners=False\n                    )\n                    goal_logits = F.interpolate(\n                        goal_logits, size=goal_target.shape[2:],\n                        mode='bilinear', align_corners=False\n                    )\n                \n                seg_loss = seg_criterion(seg_logits, seg_target)\n                start_loss = start_criterion(start_logits, start_target)\n                goal_loss = goal_criterion(goal_logits, goal_target)\n                \n                loss = seg_loss + 0.5 * start_loss + 0.5 * goal_loss\n                \n                loss.backward()\n                \n                if grad_clip:\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n                \n                optimizer.step()\n            \n            total_loss += loss.item()\n            seg_loss_sum += seg_loss.item()\n            start_loss_sum += start_loss.item()\n            goal_loss_sum += goal_loss.item()\n            \n            # Update progress bar\n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}',\n                'seg': f'{seg_loss.item():.4f}'\n            })\n        \n        avg_loss = total_loss / len(train_loader)\n        avg_seg = seg_loss_sum / len(train_loader)\n        avg_start = start_loss_sum / len(train_loader)\n        avg_goal = goal_loss_sum / len(train_loader)\n        \n        # Validation phase\n        val_loss = validate(model, val_loader, seg_criterion, start_criterion, goal_criterion)\n        \n        # Step scheduler\n        scheduler.step()\n        \n        print(\n            f\"\\n📈 Epoch {epoch:03d} | \"\n            f\"Train: {avg_loss:.4f} | \"\n            f\"Val: {val_loss:.4f} | \"\n            f\"LR: {optimizer.param_groups[0]['lr']:.6f}\"\n        )\n        print(f\"   Seg: {avg_seg:.4f} | Start: {avg_start:.4f} | Goal: {avg_goal:.4f}\")\n        \n        # Save best model\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_loss': val_loss,\n            }, save_path)\n            print(f\"✅ Saved best model (val_loss: {val_loss:.4f})\")\n        \n        # Evaluate periodically\n        if epoch % eval_every == 0:\n            evaluate_segmentation(model, val_loader)\n            evaluate_start_goal(model, val_loader)\n        \n        print(\"=\"*60)\n\n@torch.no_grad()\ndef validate(model, val_loader, seg_criterion, start_criterion, goal_criterion):\n    \"\"\"Validation loop\"\"\"\n    model.eval()\n    total_loss = 0.0\n    \n    for x, y in val_loader:\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        \n        seg_target = y.clone()\n        seg_target[seg_target >= 3] = 0\n        start_target = (y == CLASS_START).float().unsqueeze(1)\n        goal_target = (y == CLASS_GOAL).float().unsqueeze(1)\n        \n        seg_logits, start_logits, goal_logits = model(x)\n        \n        if seg_logits.shape[2:] != seg_target.shape[1:]:\n            seg_logits = F.interpolate(\n                seg_logits, size=seg_target.shape[1:],\n                mode='bilinear', align_corners=False\n            )\n            start_logits = F.interpolate(\n                start_logits, size=start_target.shape[2:],\n                mode='bilinear', align_corners=False\n            )\n            goal_logits = F.interpolate(\n                goal_logits, size=goal_target.shape[2:],\n                mode='bilinear', align_corners=False\n            )\n        \n        seg_loss = seg_criterion(seg_logits, seg_target)\n        start_loss = start_criterion(start_logits, start_target)\n        goal_loss = goal_criterion(goal_logits, goal_target)\n        \n        loss = seg_loss + 0.5 * start_loss + 0.5 * goal_loss\n        total_loss += loss.item()\n    \n    model.train()\n    return total_loss / len(val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:40:07.788515Z","iopub.execute_input":"2025-12-26T13:40:07.789235Z","iopub.status.idle":"2025-12-26T13:40:07.808093Z","shell.execute_reply.started":"2025-12-26T13:40:07.789204Z","shell.execute_reply":"2025-12-26T13:40:07.807355Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"def confusion_matrix(pred, target, num_classes=3):\n    pred = pred.view(-1)\n    target = target.view(-1)\n    cm = torch.bincount(\n        num_classes * target + pred,\n        minlength=num_classes**2\n    ).reshape(num_classes, num_classes)\n    return cm\n\ndef per_class_metrics(cm):\n    metrics = {}\n    num_classes = cm.shape[0]\n    for c in range(num_classes):\n        tp = cm[c, c]\n        fn = cm[c, :].sum() - tp\n        fp = cm[:, c].sum() - tp\n        \n        precision = tp / (tp + fp + 1e-6)\n        recall = tp / (tp + fn + 1e-6)\n        iou = tp / (tp + fp + fn + 1e-6)\n        f1 = 2 * precision * recall / (precision + recall + 1e-6)\n        \n        metrics[c] = {\n            \"precision\": precision.item(),\n            \"recall\": recall.item(),\n            \"iou\": iou.item(),\n            \"f1\": f1.item()\n        }\n    return metrics\n\n@torch.no_grad()\ndef evaluate_segmentation(model, val_loader):\n    model.eval()\n    total_cm = torch.zeros(3, 3, dtype=torch.int64)\n    \n    for x, y in tqdm(val_loader, desc=\"Evaluating segmentation\"):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        \n        seg_logits, _, _ = model(x)\n        seg_pred = torch.argmax(seg_logits, dim=1)\n        \n        seg_target = y.clone()\n        seg_target[seg_target >= 3] = 0\n        \n        cm = confusion_matrix(seg_pred.cpu(), seg_target.cpu(), num_classes=3)\n        total_cm += cm\n    \n    metrics = per_class_metrics(total_cm)\n    class_names = [\"walk\", \"wall\", \"hazard\"]\n    \n    print(\"\\n\" + \"=\"*50)\n    print(\"Segmentation Metrics:\")\n    print(\"=\"*50)\n    for c, name in enumerate(class_names):\n        m = metrics[c]\n        print(\n            f\"{name:7s} | \"\n            f\"P: {m['precision']:.3f} | \"\n            f\"R: {m['recall']:.3f} | \"\n            f\"IoU: {m['iou']:.3f} | \"\n            f\"F1: {m['f1']:.3f}\"\n        )\n    \n    mean_iou = sum(m[\"iou\"] for m in metrics.values()) / 3\n    mean_f1 = sum(m[\"f1\"] for m in metrics.values()) / 3\n    print(f\"\\n📊 Mean IoU: {mean_iou:.3f} | Mean F1: {mean_f1:.3f}\")\n    print(\"=\"*50 + \"\\n\")\n    \n    model.train()\n\n@torch.no_grad()\ndef evaluate_start_goal(model, val_loader):\n    model.eval()\n    start_correct = 0\n    goal_correct = 0\n    total = 0\n    \n    for x, y in tqdm(val_loader, desc=\"Evaluating start/goal\"):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        \n        seg_logits, start_logits, goal_logits = model(x)\n        seg_pred = torch.argmax(seg_logits, dim=1)\n        \n        wall_mask = (seg_pred == CLASS_WALL)\n        \n        # Mask walls\n        start_logits = start_logits.clone()\n        goal_logits = goal_logits.clone()\n        start_logits[wall_mask.unsqueeze(1)] = -1e9\n        goal_logits[wall_mask.unsqueeze(1)] = -1e9\n        \n        B = start_logits.size(0)\n        start_pred = torch.argmax(start_logits.view(B, -1), dim=1)\n        goal_pred = torch.argmax(goal_logits.view(B, -1), dim=1)\n        \n        start_gt = (y == CLASS_START).long().view(B, -1).argmax(dim=1)\n        goal_gt = (y == CLASS_GOAL).long().view(B, -1).argmax(dim=1)\n        \n        start_correct += (start_pred == start_gt).sum().item()\n        goal_correct += (goal_pred == goal_gt).sum().item()\n        total += B\n    \n    print(\"=\"*50)\n    print(\"Start/Goal Detection:\")\n    print(\"=\"*50)\n    print(f\"🎯 Start accuracy: {start_correct / total:.3f} ({start_correct}/{total})\")\n    print(f\"🏁 Goal  accuracy: {goal_correct / total:.3f} ({goal_correct}/{total})\")\n    print(\"=\"*50 + \"\\n\")\n    \n    model.train()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:40:20.170682Z","iopub.execute_input":"2025-12-26T13:40:20.171389Z","iopub.status.idle":"2025-12-26T13:40:20.184090Z","shell.execute_reply.started":"2025-12-26T13:40:20.171359Z","shell.execute_reply":"2025-12-26T13:40:20.183298Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# Create datasets\ntrain_dataset = GridDataset(\n        images_dir=paths['mod_images'],\n        labels_dir=paths['mod_labels'],\n        transform=train_transform\n    )\n    \nval_dataset = GridDataset(\n        images_dir=paths['mod_train_images'],\n        labels_dir=paths['mod_train_labels'],\n        transform=val_transform\n    )\n    \n    # GPU-optimized batch size and workers\ntrain_loader = DataLoader(\n        train_dataset, \n        batch_size=32,  # Larger batch for GPU\n        shuffle=True, \n        num_workers=4,  # Parallel data loading\n        pin_memory=True,  # Faster CPU->GPU transfer\n        collate_fn=custom_collate_fn  # Handle variable-sized grids\n    )\n    \nval_loader = DataLoader(\n        val_dataset, \n        batch_size=4, \n        shuffle=False, \n        num_workers=4,\n        pin_memory=True,\n        collate_fn=custom_collate_fn\n    )\n    \n    # Create improved model with attention\nmodel = ImprovedUNet(use_attention=True)\n    \n    # Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"📊 Model parameters: {total_params:,} (~{total_params/1e6:.2f}M)\")\nprint(\"=\"*60)\n    \n    # Train\ntrain_model(\n        model,\n        train_loader,\n        val_loader,\n        epochs=100,  # More epochs for GPU\n        lr=1e-3,\n        save_path=\"best_grid_model_gpu.pth\",\n        eval_every=5\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:41:34.217689Z","iopub.execute_input":"2025-12-26T13:41:34.218415Z","iopub.status.idle":"2025-12-26T13:41:34.228133Z","shell.execute_reply.started":"2025-12-26T13:41:34.218381Z","shell.execute_reply":"2025-12-26T13:41:34.227175Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Detecting Terrain","metadata":{}},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((GRID_SIZE, GRID_SIZE), interpolation=transforms.InterpolationMode.BILINEAR),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\ndef detect_terrain_type(img: Image.Image) -> str:\n    \"\"\"\n    Detect terrain type based on color statistics.\n    Lab: metallic/blue tones\n    Forest: green/brown tones\n    Desert: yellow/sandy tones\n    \"\"\"\n    img_array = np.array(img.resize((100, 100)))  # Downsample for speed\n    \n    # Average RGB values\n    avg_r = img_array[:, :, 0].mean()\n    avg_g = img_array[:, :, 1].mean()\n    avg_b = img_array[:, :, 2].mean()\n    \n    # Heuristic rules (tune these based on your data)\n    # Lab: bluish, metallic (high blue, balanced RGB)\n    if avg_b > avg_r + 10 and avg_b > avg_g + 10:\n        return 'lab'\n    \n    # Forest: greenish (high green)\n    if avg_g > avg_r + 15 and avg_g > avg_b + 15:\n        return 'forest'\n    \n    # Desert: sandy/yellow (high red+green, low blue)\n    if avg_r > avg_b + 20 and avg_g > avg_b + 20:\n        return 'desert'\n    \n    # Default fallback\n    return 'lab'\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:41:49.168691Z","iopub.execute_input":"2025-12-26T13:41:49.169295Z","iopub.status.idle":"2025-12-26T13:41:49.175118Z","shell.execute_reply.started":"2025-12-26T13:41:49.169265Z","shell.execute_reply":"2025-12-26T13:41:49.174277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# helper functions\ndef pick_start_goal_from_heatmaps(\n    start_logits: torch.Tensor,\n    goal_logits: torch.Tensor,\n    wall_mask: torch.Tensor\n) -> Tuple[Tuple[int, int], Tuple[int, int]]:\n    \"\"\"\n    Pick start and goal from heatmaps while avoiding walls.\n    \n    Args:\n        start_logits: (1, 1, G, G)\n        goal_logits: (1, 1, G, G)\n        wall_mask: (1, G, G) - True where walls are\n    \n    Returns:\n        start: (row, col)\n        goal: (row, col)\n    \"\"\"\n    _, _, G, _ = start_logits.shape\n    \n    # Apply wall masking\n    start_logits = start_logits.clone()\n    goal_logits = goal_logits.clone()\n    start_logits[0, 0][wall_mask[0]] = -1e9\n    goal_logits[0, 0][wall_mask[0]] = -1e9\n    \n    # Pick highest activation\n    start_flat = torch.argmax(start_logits.view(-1)).item()\n    goal_flat = torch.argmax(goal_logits.view(-1)).item()\n    \n    start = (start_flat // G, start_flat % G)\n    goal = (goal_flat // G, goal_flat % G)\n    \n    return start, goal\n\n\ndef load_velocity_boost(json_path: Path) -> np.ndarray:\n    \"\"\"Load velocity boost from JSON file.\"\"\"\n    try:\n        with json_path.open(\"r\", encoding=\"utf-8\") as f:\n            data = json.load(f)\n        \n        if \"boost\" not in data:\n            print(f\"⚠️  Warning: 'boost' not found in {json_path}. Using zero boost.\")\n            return np.zeros((GRID_SIZE, GRID_SIZE), dtype=np.float32)\n        \n        boost = np.array(data[\"boost\"], dtype=np.float32)\n        \n        if boost.shape != (GRID_SIZE, GRID_SIZE):\n            print(f\"⚠️  Warning: Boost shape {boost.shape} != ({GRID_SIZE}, {GRID_SIZE}). Resizing.\")\n            boost = np.zeros((GRID_SIZE, GRID_SIZE), dtype=np.float32)\n        \n        return boost\n    \n    except Exception as e:\n        print(f\"⚠️  Error loading boost from {json_path}: {e}. Using zero boost.\")\n        return np.zeros((GRID_SIZE, GRID_SIZE), dtype=np.float32)\n\n\ndef estimate_cost_grid(\n    pred_grid: np.ndarray,\n    boost: np.ndarray,\n    terrain: str = 'lab'\n) -> np.ndarray:\n    \"\"\"\n    Calculate step costs based on terrain, class, and boost.\n    \n    Formula: step_cost(i,j) = base_cost(i,j) - boost(i,j)\n    \"\"\"\n    G = pred_grid.shape[0]\n    cost_grid = np.zeros((G, G), dtype=np.float32)\n    \n    terrain_cost = TERRAIN_COSTS[terrain]\n    \n    for i in range(G):\n        for j in range(G):\n            cell_class = pred_grid[i, j]\n            base_cost = terrain_cost.get(cell_class, 1.0)\n            \n            # Apply boost\n            step_cost = base_cost - boost[i, j]\n            \n            # Ensure positive cost (safety check)\n            step_cost = max(step_cost, 0.01)\n            \n            cost_grid[i, j] = step_cost\n    \n    return cost_grid\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:41:50.181895Z","iopub.execute_input":"2025-12-26T13:41:50.182152Z","iopub.status.idle":"2025-12-26T13:41:50.191682Z","shell.execute_reply.started":"2025-12-26T13:41:50.182128Z","shell.execute_reply":"2025-12-26T13:41:50.190999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Paths detection with fallback","metadata":{}},{"cell_type":"code","source":"def dijkstra_path(\n    cost_grid: np.ndarray,\n    start: Tuple[int, int],\n    goal: Tuple[int, int]\n) -> List[Tuple[int, int]]:\n    \"\"\"\n    Dijkstra's algorithm to find optimal path.\n    \n    Returns:\n        List of (row, col) positions from start to goal.\n        Empty list if no path found.\n    \"\"\"\n    G = cost_grid.shape[0]\n    INF = float('inf')\n    \n    sr, sc = start\n    gr, gc = goal\n    \n    # Distance array\n    dist = np.full((G, G), INF, dtype=np.float32)\n    dist[sr, sc] = 0.0\n    \n    # Previous node for path reconstruction\n    prev = {}\n    \n    # Priority queue: (cost, (row, col))\n    heap = [(0.0, start)]\n    visited = set()\n    \n    while heap:\n        cur_cost, (r, c) = heapq.heappop(heap)\n        \n        if (r, c) in visited:\n            continue\n        visited.add((r, c))\n        \n        # Found goal\n        if (r, c) == goal:\n            break\n        \n        # Skip if we found a better path already\n        if cur_cost > dist[r, c]:\n            continue\n        \n        # Explore 4-connected neighbors\n        for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n            nr, nc = r + dr, c + dc\n            \n            # Boundary check\n            if not (0 <= nr < G and 0 <= nc < G):\n                continue\n            \n            # Wall check\n            if cost_grid[nr, nc] == INF:\n                continue\n            \n            # Calculate new cost\n            new_cost = cur_cost + cost_grid[nr, nc]\n            \n            # Update if better\n            if new_cost < dist[nr, nc]:\n                dist[nr, nc] = new_cost\n                prev[(nr, nc)] = (r, c)\n                heapq.heappush(heap, (new_cost, (nr, nc)))\n    \n    # Reconstruct path\n    if goal not in prev and goal != start:\n        return []\n    \n    path = []\n    cur = goal\n    while cur in prev:\n        path.append(cur)\n        cur = prev[cur]\n    path.append(start)\n    path.reverse()\n    \n    return path\n\ndef astar_fallback(\n    pred_grid: np.ndarray,\n    start: Tuple[int, int],\n    goal: Tuple[int, int]\n) -> List[Tuple[int, int]]:\n    \"\"\"\n    A* fallback (ignores boost, uses Manhattan distance heuristic).\n    \"\"\"\n    G = pred_grid.shape[0]\n    INF = float('inf')\n    \n    def heuristic(pos):\n        return abs(pos[0] - goal[0]) + abs(pos[1] - goal[1])\n    \n    sr, sc = start\n    gr, gc = goal\n    \n    dist = np.full((G, G), INF, dtype=np.float32)\n    dist[sr, sc] = 0.0\n    \n    prev = {}\n    heap = [(heuristic(start), 0.0, start)]\n    visited = set()\n    \n    while heap:\n        _, cur_cost, (r, c) = heapq.heappop(heap)\n        \n        if (r, c) in visited:\n            continue\n        visited.add((r, c))\n        \n        if (r, c) == goal:\n            break\n        \n        for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n            nr, nc = r + dr, c + dc\n            \n            if not (0 <= nr < G and 0 <= nc < G):\n                continue\n            \n            # Simple cost: 1 for walkable, INF for wall\n            if pred_grid[nr, nc] == CLASS_WALL:\n                continue\n            \n            new_cost = cur_cost + 1.0\n            \n            if new_cost < dist[nr, nc]:\n                dist[nr, nc] = new_cost\n                prev[(nr, nc)] = (r, c)\n                priority = new_cost + heuristic((nr, nc))\n                heapq.heappush(heap, (priority, new_cost, (nr, nc)))\n    \n    # Reconstruct\n    if goal not in prev and goal != start:\n        return []\n    \n    path = []\n    cur = goal\n    while cur in prev:\n        path.append(cur)\n        cur = prev[cur]\n    path.append(start)\n    path.reverse()\n    \n    return path\n\n\ndef greedy_fallback(\n    pred_grid: np.ndarray,\n    start: Tuple[int, int],\n    goal: Tuple[int, int],\n    max_steps: int = 500\n) -> List[Tuple[int, int]]:\n    \"\"\"\n    Greedy fallback: move toward goal, avoid walls.\n    Last resort if Dijkstra and A* fail.\n    \"\"\"\n    G = pred_grid.shape[0]\n    path = [start]\n    current = start\n    visited = {start}\n    \n    for _ in range(max_steps):\n        if current == goal:\n            break\n        \n        cr, cc = current\n        gr, gc = goal\n        \n        # Try to move closer to goal\n        candidates = []\n        for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n            nr, nc = cr + dr, cc + dc\n            \n            if not (0 <= nr < G and 0 <= nc < G):\n                continue\n            if pred_grid[nr, nc] == CLASS_WALL:\n                continue\n            if (nr, nc) in visited:\n                continue\n            \n            # Manhattan distance to goal\n            dist = abs(nr - gr) + abs(nc - gc)\n            candidates.append((dist, (nr, nc)))\n        \n        if not candidates:\n            # Stuck - try to backtrack\n            break\n        \n        # Pick closest to goal\n        candidates.sort()\n        _, next_pos = candidates[0]\n        \n        path.append(next_pos)\n        visited.add(next_pos)\n        current = next_pos\n    \n    return path\n\n\ndef desperate_fallback(\n    start: Tuple[int, int],\n    goal: Tuple[int, int]\n) -> List[Tuple[int, int]]:\n    \"\"\"\n    Absolute last resort: Generate ANY valid path from start to goal.\n    Ignores walls completely, just creates a Manhattan path.\n    \n    This guarantees a non-empty submission even if the map is impossible.\n    Strategy: Move horizontally first, then vertically.\n    \"\"\"\n    sr, sc = start\n    gr, gc = goal\n    \n    path = [start]\n    current_r, current_c = sr, sc\n    \n    # Move horizontally to goal column\n    while current_c != gc:\n        if current_c < gc:\n            current_c += 1\n        else:\n            current_c -= 1\n        path.append((current_r, current_c))\n    \n    # Move vertically to goal row\n    while current_r != gr:\n        if current_r < gr:\n            current_r += 1\n        else:\n            current_r -= 1\n        path.append((current_r, current_c))\n    \n    return path\n\n\ndef relaxed_greedy_fallback(\n    pred_grid: np.ndarray,\n    start: Tuple[int, int],\n    goal: Tuple[int, int],\n    max_steps: int = 500\n) -> List[Tuple[int, int]]:\n    \"\"\"\n    Even more relaxed greedy: allows moving through walls if necessary.\n    Used as second-to-last resort before desperate_fallback.\n    \"\"\"\n    G = pred_grid.shape[0]\n    path = [start]\n    current = start\n    visited = {start}\n    \n    for step in range(max_steps):\n        if current == goal:\n            break\n        \n        cr, cc = current\n        gr, gc = goal\n        \n        # Try to move closer to goal\n        candidates = []\n        for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n            nr, nc = cr + dr, cc + dc\n            \n            if not (0 <= nr < G and 0 <= nc < G):\n                continue\n            if (nr, nc) in visited:\n                continue\n            \n            # Manhattan distance to goal\n            dist = abs(nr - gr) + abs(nc - gc)\n            \n            # Penalty for walls, but still allow them\n            penalty = 100 if pred_grid[nr, nc] == CLASS_WALL else 0\n            candidates.append((dist + penalty, (nr, nc)))\n        \n        if not candidates:\n            # Allow revisiting if completely stuck\n            for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]:\n                nr, nc = cr + dr, cc + dc\n                \n                if not (0 <= nr < G and 0 <= nc < G):\n                    continue\n                \n                dist = abs(nr - gr) + abs(nc - gc)\n                penalty = 100 if pred_grid[nr, nc] == CLASS_WALL else 0\n                candidates.append((dist + penalty, (nr, nc)))\n        \n        if not candidates:\n            # Still stuck? Break and use desperate fallback\n            break\n        \n        # Pick best candidate\n        candidates.sort()\n        _, next_pos = candidates[0]\n        \n        path.append(next_pos)\n        visited.add(next_pos)\n        current = next_pos\n    \n    # If we didn't reach goal, continue with desperate fallback from current position\n    if current != goal:\n        remaining_path = desperate_fallback(current, goal)\n        if len(remaining_path) > 1:\n            path.extend(remaining_path[1:])  # Skip duplicate current position\n    \n    return path\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:41:54.468159Z","iopub.execute_input":"2025-12-26T13:41:54.468695Z","iopub.status.idle":"2025-12-26T13:41:54.492539Z","shell.execute_reply.started":"2025-12-26T13:41:54.468669Z","shell.execute_reply":"2025-12-26T13:41:54.491871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Path conversion\ndef path_to_lrud(path: List[Tuple[int, int]]) -> str:\n    \"\"\"\n    Convert list of (row, col) positions to lrud sequence.\n    \"\"\"\n    if len(path) <= 1:\n        return \"\"\n    \n    moves = []\n    for (r1, c1), (r2, c2) in zip(path[:-1], path[1:]):\n        dr, dc = r2 - r1, c2 - c1\n        \n        if dr == 1 and dc == 0:\n            moves.append(\"d\")  # down\n        elif dr == -1 and dc == 0:\n            moves.append(\"u\")  # up\n        elif dr == 0 and dc == 1:\n            moves.append(\"r\")  # right\n        elif dr == 0 and dc == -1:\n            moves.append(\"l\")  # left\n        else:\n            # Invalid move - should not happen\n            print(f\"⚠️  Warning: Invalid move from ({r1},{c1}) to ({r2},{c2})\")\n    \n    return \"\".join(moves)\n\ndef predict_outputs(\n    model: nn.Module,\n    img_path: Path\n) -> Tuple[np.ndarray, Tuple[int, int], Tuple[int, int], str]:\n    \"\"\"\n    Run model inference on a single image.\n    \n    Returns:\n        grid_pred: (G, G) predicted class grid\n        start: (row, col)\n        goal: (row, col)\n        terrain: detected terrain type\n    \"\"\"\n    # Load and transform image\n    img = Image.open(img_path).convert('RGB')\n    \n    # Detect terrain\n    terrain = detect_terrain_type(img)\n    \n    x = transform(img).unsqueeze(0).to(DEVICE)  # (1, 3, G, G)\n    \n    model.eval()\n    with torch.no_grad():\n        seg_logits, start_logits, goal_logits = model(x)\n        \n        # Segmentation prediction\n        seg_pred = torch.argmax(seg_logits, dim=1)  # (1, G, G)\n        \n        # Build wall mask\n        wall_mask = (seg_pred == CLASS_WALL)  # (1, G, G)\n        \n        # Pick start & goal with wall masking\n        start_rc, goal_rc = pick_start_goal_from_heatmaps(\n            start_logits.cpu(),\n            goal_logits.cpu(),\n            wall_mask.cpu()\n        )\n    \n    grid_pred = seg_pred.squeeze(0).cpu().numpy().astype(np.int64)\n    \n    return grid_pred, start_rc, goal_rc, terrain\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:41:59.327856Z","iopub.execute_input":"2025-12-26T13:41:59.328152Z","iopub.status.idle":"2025-12-26T13:41:59.336084Z","shell.execute_reply.started":"2025-12-26T13:41:59.328129Z","shell.execute_reply":"2025-12-26T13:41:59.335386Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"def run_inference_on_test(\n    model: nn.Module,\n    test_images_dir: Path,\n    test_velocities_dir: Path,\n    submission_path: Path\n):\n    \"\"\"\n    Complete inference pipeline with multi-level fallbacks.\n    Guarantees a path for every image.\n    \"\"\"\n\n    if not submission_path.exists() :\n        submission_path.parent.mkdir(parents=True, exist_ok=True)  # create folders if needed\n        submission_path.touch(exist_ok=True)\n        \n    model.eval()\n    \n    image_paths = sorted(test_images_dir.glob(\"*.png\"))\n    print(f\"🔍 Found {len(image_paths)} test images\")\n    \n    records = []\n    fallback_stats = {\n        'dijkstra': 0,\n        'astar': 0,\n        'greedy': 0,\n        'relaxed': 0,\n        'desperate': 0,\n        'error': 0\n    }\n    \n    for img_path in tqdm(image_paths, desc=\"Running inference\"):\n        image_id = img_path.stem\n        \n        try:\n            # 1. Predict grid, start, goal, terrain\n            grid_pred, start, goal, terrain = predict_outputs(model, img_path)\n            \n            # 2. Load velocity boost\n            boost = load_velocity_boost(test_velocities_dir / f\"{image_id}.json\")\n            \n            # 3. Calculate cost grid\n            cost_grid = estimate_cost_grid(grid_pred, boost, terrain)\n            \n            # 4. Try Dijkstra first (optimal with boost)\n            path = dijkstra_path(cost_grid, start, goal)\n            \n            if len(path) >= 2:\n                fallback_stats['dijkstra'] += 1\n            else:\n                # 5. Fallback to A* (ignores boost, uses heuristic)\n                print(f\"⚠️  {image_id}: Dijkstra failed, trying A*...\")\n                path = astar_fallback(grid_pred, start, goal)\n                \n                if len(path) >= 2:\n                    fallback_stats['astar'] += 1\n                else:\n                    # 6. Fallback to greedy (simple goal-seeking, avoids walls)\n                    print(f\"⚠️  {image_id}: A* failed, trying greedy...\")\n                    path = greedy_fallback(grid_pred, start, goal)\n                    \n                    if len(path) >= 2:\n                        fallback_stats['greedy'] += 1\n                    else:\n                        # 7. Fallback to relaxed greedy (allows walls with penalty)\n                        print(f\"⚠️  {image_id}: Greedy failed, trying relaxed greedy...\")\n                        path = relaxed_greedy_fallback(grid_pred, start, goal)\n                        \n                        if len(path) >= 2:\n                            fallback_stats['relaxed'] += 1\n                        else:\n                            # 8. ABSOLUTE LAST RESORT: Desperate fallback (ignores everything)\n                            print(f\"🚨 {image_id}: All intelligent pathfinding failed! Using desperate fallback...\")\n                            path = desperate_fallback(start, goal)\n                            fallback_stats['desperate'] += 1\n            \n            # 9. Convert to moves\n            moves = path_to_lrud(path)\n            \n            # 10. Final validation\n            if len(moves) == 0 and start != goal:\n                print(f\"❌ {image_id}: Path conversion failed! Using direct path...\")\n                path = desperate_fallback(start, goal)\n                moves = path_to_lrud(path)\n            \n            records.append({\"image_id\": image_id, \"path\": moves})\n            \n        except Exception as e:\n            print(f\"❌ Critical error processing {image_id}: {e}\")\n            # Even on error, try to provide a reasonable path\n            try:\n                # Try to at least create a Manhattan path\n                path = desperate_fallback((0, 0), (GRID_SIZE-1, GRID_SIZE-1))\n                moves = path_to_lrud(path)\n                records.append({\"image_id\": image_id, \"path\": moves})\n                fallback_stats['error'] += 1\n            except:\n                # Absolute worst case: empty path\n                records.append({\"image_id\": image_id, \"path\": \"\"})\n                fallback_stats['error'] += 1\n    \n    # Write submission CSV\n    df = pd.DataFrame(records)\n    df.to_csv(submission_path, index=False)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"✅ INFERENCE COMPLETE!\")\n    print(\"=\"*60)\n    print(f\"📁 Submission written to: {submission_path}\")\n    print(f\"📊 Total images processed: {len(records)}\")\n    print(\"\\n📈 Pathfinding Statistics:\")\n    print(f\"  🥇 Dijkstra (optimal):        {fallback_stats['dijkstra']:4d} ({fallback_stats['dijkstra']/len(records)*100:.1f}%)\")\n    print(f\"  🥈 A* (heuristic):            {fallback_stats['astar']:4d} ({fallback_stats['astar']/len(records)*100:.1f}%)\")\n    print(f\"  🥉 Greedy (goal-seeking):     {fallback_stats['greedy']:4d} ({fallback_stats['greedy']/len(records)*100:.1f}%)\")\n    print(f\"  ⚠️  Relaxed (wall-tolerant):  {fallback_stats['relaxed']:4d} ({fallback_stats['relaxed']/len(records)*100:.1f}%)\")\n    print(f\"  🚨 Desperate (Manhattan):     {fallback_stats['desperate']:4d} ({fallback_stats['desperate']/len(records)*100:.1f}%)\")\n    print(f\"  ❌ Errors:                    {fallback_stats['error']:4d} ({fallback_stats['error']/len(records)*100:.1f}%)\")\n    \n    empty_paths = sum(1 for r in records if r['path'] == '')\n    if empty_paths > 0:\n        print(f\"\\n⚠️  WARNING: {empty_paths} images have empty paths!\")\n    else:\n        print(f\"\\n✅ All {len(records)} images have valid paths!\")\n    \n    print(\"=\"*60)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:42:01.609149Z","iopub.execute_input":"2025-12-26T13:42:01.609781Z","iopub.status.idle":"2025-12-26T13:42:01.622174Z","shell.execute_reply.started":"2025-12-26T13:42:01.609756Z","shell.execute_reply":"2025-12-26T13:42:01.621526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model\nMODEL_PATH = Path(\"best_grid_model_gpu.pth\")\nmodel = ImprovedUNet(use_attention=True)  # Use your model class\n    \ncheckpoint = torch.load(MODEL_PATH, map_location=DEVICE)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.to(DEVICE)\nmodel.eval()\n\n# Run inference\nrun_inference_on_test(\n        model=model,\n        test_images_dir=paths['test_images'],\n        test_velocities_dir=paths['test_velocities'],\n        submission_path=paths['submission_path']\n    )\nprint(\"🎉 Inference complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-26T13:42:05.743945Z","iopub.execute_input":"2025-12-26T13:42:05.744200Z","iopub.status.idle":"2025-12-26T13:42:05.994167Z","shell.execute_reply.started":"2025-12-26T13:42:05.744181Z","shell.execute_reply":"2025-12-26T13:42:05.993345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}