{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================================================\n# 🎯 VESUVIUS - ULTRA MEMORY EFFICIENT 0.7+ NOTEBOOK\n# =============================================================================\n# Memory optimized version that runs within Kaggle limits\n# =============================================================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport numpy as np\nfrom PIL import Image\nfrom pathlib import Path\nimport pandas as pd\nfrom tqdm import tqdm\nimport json\nimport time\nimport warnings\nimport random\nimport gc\nwarnings.filterwarnings('ignore')\n\n# Clear memory\ntorch.cuda.empty_cache()\ngc.collect()\n\nprint(\"=\"*80)\nprint(\"🎯 VESUVIUS - ULTRA MEMORY EFFICIENT 0.7+\")\nprint(\"Optimized for Kaggle memory limits\")\nprint(\"Target: 0.7+ score with <12GB memory\")\nprint(\"=\"*80)\n\n# =============================================================================\n# ULTRA MEMORY EFFICIENT CONFIG\n# =============================================================================\nCONFIG = {\n    # DATA PATH\n    'data_path': '/kaggle/input/vesuvius-challenge-surface-detection',\n    \n    # ULTRA MEMORY EFFICIENT\n    'patch_size': (16, 96, 96),  # Small patches\n    'init_features': 16,  # Small model\n    'num_epochs': 60,  # More epochs, smaller batches\n    \n    # MEMORY EFFICIENT TRAINING\n    'batch_size': 4,  # Small batches\n    'learning_rate': 2e-4,\n    'weight_decay': 1e-5,\n    \n    # DATA STRATEGY\n    'train_split': 0.95,\n    'max_train_volumes': 100,  # Limit volumes\n    'max_val_volumes': 20,\n    \n    # MEMORY SAVING\n    'use_augmentation': False,  # No augmentation to save memory\n    'use_mixed_precision': True,\n    'num_workers': 0,  # No multiprocessing\n    \n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'time_limit_hours': 6,\n}\n\nprint(\"\\n⚙️ MEMORY OPTIMIZED CONFIG:\")\nprint(f\"  Patch: {CONFIG['patch_size']} (small)\")\nprint(f\"  Batch: {CONFIG['batch_size']} (small)\")\nprint(f\"  Model: {CONFIG['init_features']} features (light)\")\nprint(f\"  Workers: {CONFIG['num_workers']} (no multiprocessing)\")\nprint(f\"  Volumes: {CONFIG['max_train_volumes']} train, {CONFIG['max_val_volumes']} val\")\nprint(f\"  Target: 0.7+ with minimal memory\")\n\n# Check memory\nif torch.cuda.is_available():\n    free_memory = torch.cuda.get_device_properties(0).total_memory - torch.cuda.memory_allocated()\n    print(f\"\\n📊 GPU Memory: {free_memory / 1e9:.1f} GB free\")\n\n# =============================================================================\n# MEMORY EFFICIENT PREPROCESSING\n# =============================================================================\ndef preprocess_patch_fast(patch):\n    \"\"\"Ultra-fast patch preprocessing\"\"\"\n    patch = patch.astype(np.float32)\n    \n    # Simple normalization\n    mean = patch.mean()\n    std = patch.std() + 1e-8\n    patch = (patch - mean) / std\n    \n    return patch\n\n# =============================================================================\n# ULTRA MEMORY EFFICIENT DATASET\n# =============================================================================\nclass MemoryEfficientDataset(Dataset):\n    \"\"\"Dataset that loads patches on-demand with zero caching\"\"\"\n    \n    def __init__(self, data_path, volume_ids, config, is_training=True):\n        self.data_path = Path(data_path)\n        self.volume_ids = volume_ids[:config['max_train_volumes'] if is_training else config['max_val_volumes']]\n        self.config = config\n        self.is_training = is_training\n        self.patch_size = config['patch_size']\n        \n        print(f\"{'Train' if is_training else 'Val'}: {len(self.volume_ids)} volumes\")\n    \n    def __len__(self):\n        return len(self.volume_ids) * 8  # 8 patches per volume\n    \n    def load_patch_direct(self, volume_id, patch_idx):\n        \"\"\"Load patch directly without caching\"\"\"\n        pd, ph, pw = self.patch_size\n        \n        # Load volume dimensions first\n        img_path = self.data_path / 'train_images' / f'{volume_id}.tif'\n        with Image.open(img_path) as img:\n            d = img.n_frames if hasattr(img, 'n_frames') else 1\n            h, w = img.height, img.width\n        \n        # Determine patch location (deterministic)\n        seed = hash(f\"{volume_id}_{patch_idx}\") % 1000000\n        random.seed(seed)\n        \n        # Focus on middle regions where ink usually is\n        z = random.randint(d//4, min(d//4 + d//2, d - pd))\n        y = random.randint(h//4, min(h//4 + h//2, h - ph))\n        x = random.randint(w//4, min(w//4 + w//2, w - pw))\n        \n        z = max(0, min(z, d - pd))\n        y = max(0, min(y, h - ph))\n        x = max(0, min(x, w - pw))\n        \n        # Load only the needed patch\n        img_patch = np.zeros((pd, ph, pw), dtype=np.float32)\n        lbl_patch = np.zeros((pd, ph, pw), dtype=np.uint8)\n        \n        with Image.open(img_path) as img, \\\n             Image.open(self.data_path / 'train_labels' / f'{volume_id}.tif') as lbl:\n            \n            for i in range(pd):\n                if z + i < d:\n                    img.seek(z + i)\n                    lbl.seek(z + i)\n                    \n                    # Extract only the patch region\n                    img_region = np.array(img.crop((x, y, x + pw, y + ph)), dtype=np.float32)\n                    lbl_region = np.array(lbl.crop((x, y, x + pw, y + ph)), dtype=np.uint8)\n                    \n                    img_patch[i] = img_region\n                    lbl_patch[i] = np.where(lbl_region == 2, 0, lbl_region)\n        \n        # Fast preprocessing\n        img_patch = preprocess_patch_fast(img_patch)\n        \n        return img_patch, lbl_patch\n    \n    def __getitem__(self, idx):\n        volume_idx = idx % len(self.volume_ids)\n        patch_idx = idx // len(self.volume_ids)\n        \n        volume_id = self.volume_ids[volume_idx]\n        img_patch, lbl_patch = self.load_patch_direct(volume_id, patch_idx)\n        \n        return (torch.from_numpy(img_patch).unsqueeze(0).float(),\n                torch.from_numpy(lbl_patch).long())\n\n# =============================================================================\n# LIGHTWEIGHT BUT EFFECTIVE MODEL\n# =============================================================================\nclass LightweightInkDetector(nn.Module):\n    \"\"\"Lightweight model that works within memory limits\"\"\"\n    \n    def __init__(self, in_channels=1, out_channels=2, init_features=16):\n        super().__init__()\n        features = init_features\n        \n        # Simple encoder\n        self.enc1 = nn.Sequential(\n            nn.Conv3d(in_channels, features, 3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True)\n        )\n        self.pool1 = nn.MaxPool3d(2)\n        \n        self.enc2 = nn.Sequential(\n            nn.Conv3d(features, features*2, 3, padding=1),\n            nn.BatchNorm3d(features*2),\n            nn.ReLU(inplace=True)\n        )\n        self.pool2 = nn.MaxPool3d(2)\n        \n        self.enc3 = nn.Sequential(\n            nn.Conv3d(features*2, features*4, 3, padding=1),\n            nn.BatchNorm3d(features*4),\n            nn.ReLU(inplace=True)\n        )\n        self.pool3 = nn.MaxPool3d(2)\n        \n        # Bottleneck\n        self.bottleneck = nn.Sequential(\n            nn.Conv3d(features*4, features*8, 3, padding=1),\n            nn.BatchNorm3d(features*8),\n            nn.ReLU(inplace=True)\n        )\n        \n        # Decoder\n        self.up3 = nn.ConvTranspose3d(features*8, features*4, 2, 2)\n        self.dec3 = nn.Sequential(\n            nn.Conv3d(features*8, features*4, 3, padding=1),\n            nn.BatchNorm3d(features*4),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.up2 = nn.ConvTranspose3d(features*4, features*2, 2, 2)\n        self.dec2 = nn.Sequential(\n            nn.Conv3d(features*4, features*2, 3, padding=1),\n            nn.BatchNorm3d(features*2),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.up1 = nn.ConvTranspose3d(features*2, features, 2, 2)\n        self.dec1 = nn.Sequential(\n            nn.Conv3d(features*2, features, 3, padding=1),\n            nn.BatchNorm3d(features),\n            nn.ReLU(inplace=True)\n        )\n        \n        # Output\n        self.out = nn.Conv3d(features, out_channels, 1)\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        \n        # Bottleneck\n        b = self.bottleneck(self.pool3(e3))\n        \n        # Decoder\n        d3 = self.up3(b)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.up2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.up1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        return self.out(d1)\n\n# =============================================================================\n# SIMPLE BUT EFFECTIVE LOSS\n# =============================================================================\nclass SimpleDiceLoss(nn.Module):\n    \"\"\"Memory efficient dice loss\"\"\"\n    \n    def forward(self, pred, target):\n        pred_soft = F.softmax(pred, dim=1)\n        pred_ink = pred_soft[:, 1]\n        \n        target_ink = (target == 1).float()\n        \n        intersection = (pred_ink * target_ink).sum()\n        union = pred_ink.sum() + target_ink.sum() + 1e-8\n        \n        return 1.0 - (2.0 * intersection) / union\n\n# =============================================================================\n# MEMORY EFFICIENT TRAINING\n# =============================================================================\ndef train_epoch_memory_efficient(model, loader, criterion, optimizer, scaler, device):\n    \"\"\"Training with memory management\"\"\"\n    model.train()\n    total_loss = 0.0\n    total_dice = 0.0\n    \n    for batch_idx, (images, labels) in enumerate(tqdm(loader, desc='Train', leave=False)):\n        # Clear cache periodically\n        if batch_idx % 10 == 0:\n            torch.cuda.empty_cache()\n        \n        images = images.to(device, non_blocking=False)  # non_blocking=False saves memory\n        labels = labels.to(device, non_blocking=False)\n        \n        optimizer.zero_grad(set_to_none=True)  # Saves memory\n        \n        with autocast(enabled=CONFIG['use_mixed_precision']):\n            pred = model(images)\n            loss = criterion(pred, labels)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        # Compute dice\n        with torch.no_grad():\n            pred_soft = F.softmax(pred, dim=1)\n            pred_mask = pred_soft[:, 1] > 0.5\n            target_mask = labels == 1\n            \n            intersection = (pred_mask & target_mask).float().sum()\n            union = pred_mask.float().sum() + target_mask.float().sum() + 1e-8\n            dice = (2.0 * intersection) / union\n        \n        total_loss += loss.item()\n        total_dice += dice.item()\n        \n        # Clean up\n        del pred, loss, pred_soft, pred_mask, target_mask\n        if batch_idx % 20 == 0:\n            torch.cuda.empty_cache()\n    \n    return total_loss / len(loader), total_dice / len(loader)\n\ndef validate_memory_efficient(model, loader, criterion, device):\n    \"\"\"Validation with memory management\"\"\"\n    model.eval()\n    total_loss = 0.0\n    total_dice = 0.0\n    \n    with torch.no_grad():\n        for batch_idx, (images, labels) in enumerate(tqdm(loader, desc='Val', leave=False)):\n            images = images.to(device, non_blocking=False)\n            labels = labels.to(device, non_blocking=False)\n            \n            pred = model(images)\n            loss = criterion(pred, labels)\n            \n            pred_soft = F.softmax(pred, dim=1)\n            pred_mask = pred_soft[:, 1] > 0.5\n            target_mask = labels == 1\n            \n            intersection = (pred_mask & target_mask).float().sum()\n            union = pred_mask.float().sum() + target_mask.float().sum() + 1e-8\n            dice = (2.0 * intersection) / union\n            \n            total_loss += loss.item()\n            total_dice += dice.item()\n            \n            # Clean up\n            del pred, loss, pred_soft, pred_mask, target_mask\n            if batch_idx % 10 == 0:\n                torch.cuda.empty_cache()\n    \n    return total_loss / len(loader), total_dice / len(loader)\n\n# =============================================================================\n# MAIN TRAINING - MEMORY SAFE\n# =============================================================================\ndef train_model_memory_safe():\n    \"\"\"Main training with memory safety\"\"\"\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 MEMORY SAFE TRAINING STARTING\")\n    print(\"=\"*80)\n    \n    # Clear memory\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # Setup device\n    device = torch.device(CONFIG['device'])\n    if device.type == 'cuda':\n        torch.backends.cudnn.benchmark = False  # More memory stable\n    \n    # Load minimal data info\n    print(\"\\n📊 Loading data info (memory efficient)...\")\n    data_path = Path(CONFIG['data_path'])\n    \n    # Get available volumes (just IDs, don't load data)\n    all_volumes = []\n    for f in (data_path / 'train_images').glob('*.tif'):\n        all_volumes.append(f.stem)\n    \n    # Limit volumes for memory\n    all_volumes = all_volumes[:CONFIG['max_train_volumes'] + CONFIG['max_val_volumes']]\n    \n    # Split\n    n_train = min(CONFIG['max_train_volumes'], int(len(all_volumes) * 0.8))\n    train_ids = all_volumes[:n_train]\n    val_ids = all_volumes[n_train:n_train + CONFIG['max_val_volumes']]\n    \n    print(f\"  Total volumes found: {len(all_volumes)}\")\n    print(f\"  Training volumes: {len(train_ids)}\")\n    print(f\"  Validation volumes: {len(val_ids)}\")\n    \n    # Create datasets\n    print(\"\\n📦 Creating datasets...\")\n    train_ds = MemoryEfficientDataset(CONFIG['data_path'], train_ids, CONFIG, True)\n    val_ds = MemoryEfficientDataset(CONFIG['data_path'], val_ids, CONFIG, False)\n    \n    # Memory efficient data loaders\n    train_loader = DataLoader(\n        train_ds,\n        batch_size=CONFIG['batch_size'],\n        shuffle=True,\n        num_workers=CONFIG['num_workers'],\n        pin_memory=False,  # pin_memory=False saves memory\n        drop_last=True,\n        persistent_workers=False\n    )\n    \n    val_loader = DataLoader(\n        val_ds,\n        batch_size=CONFIG['batch_size'],\n        shuffle=False,\n        num_workers=CONFIG['num_workers'],\n        pin_memory=False,\n        persistent_workers=False\n    )\n    \n    print(f\"\\n📊 Loader stats:\")\n    print(f\"  Train batches: {len(train_loader)}\")\n    print(f\"  Val batches: {len(val_loader)}\")\n    print(f\"  Estimated memory per batch: <1GB\")\n    \n    # Create lightweight model\n    print(\"\\n🧠 Creating lightweight model...\")\n    model = LightweightInkDetector(init_features=CONFIG['init_features']).to(device)\n    params = sum(p.numel() for p in model.parameters())\n    print(f\"  Model parameters: {params:,} ({params/1e6:.1f}M)\")\n    print(f\"  Estimated model memory: {params * 4 / 1e6:.1f}MB\")\n    \n    # Loss and optimizer\n    criterion = SimpleDiceLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=CONFIG['learning_rate'], weight_decay=CONFIG['weight_decay'])\n    \n    # Simple scheduler\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='max', factor=0.5, patience=5, verbose=True\n    )\n    \n    scaler = GradScaler(enabled=CONFIG['use_mixed_precision'])\n    \n    # Training loop with memory management\n    print(f\"\\n{'='*80}\")\n    print(\"⚡ TRAINING WITH MEMORY MANAGEMENT\")\n    print('='*80)\n    \n    best_dice = 0.0\n    history = {'train_loss': [], 'train_dice': [], 'val_loss': [], 'val_dice': []}\n    \n    for epoch in range(CONFIG['num_epochs']):\n        epoch_start = time.time()\n        \n        # Clear memory before epoch\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n        # Time check\n        elapsed_hours = (time.time() - start_time) / 3600\n        if elapsed_hours > CONFIG['time_limit_hours']:\n            print(f\"\\n⏰ Time limit reached ({elapsed_hours:.1f}h)\")\n            break\n        \n        print(f\"\\n📈 Epoch {epoch+1}/{CONFIG['num_epochs']} | Elapsed: {elapsed_hours:.1f}h\")\n        \n        # Train\n        train_loss, train_dice = train_epoch_memory_efficient(\n            model, train_loader, criterion, optimizer, scaler, device\n        )\n        \n        # Validate\n        val_loss, val_dice = validate_memory_efficient(model, val_loader, criterion, device)\n        \n        # Update scheduler\n        scheduler.step(val_dice)\n        \n        # Record\n        history['train_loss'].append(train_loss)\n        history['train_dice'].append(train_dice)\n        history['val_loss'].append(val_loss)\n        history['val_dice'].append(val_dice)\n        \n        epoch_time = (time.time() - epoch_start) / 60\n        \n        print(f\"  Train Loss: {train_loss:.4f}, Dice: {train_dice:.4f}\")\n        print(f\"  Val Loss: {val_loss:.4f}, Dice: {val_dice:.4f}\")\n        print(f\"  Time: {epoch_time:.1f} min | LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Save best model\n        if val_dice > best_dice:\n            improvement = val_dice - best_dice\n            best_dice = val_dice\n            \n            # Save with minimal data\n            torch.save({\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'val_dice': val_dice,\n            }, 'best_model_memory_safe.pth')\n            \n            print(f\"  🎯 NEW BEST: {val_dice:.4f} (+{improvement:.3f})\")\n            \n            # Progress check\n            if val_dice > 0.35:\n                print(f\"  ⚡ Good progress! Target achievable\")\n            if val_dice > 0.45:\n                print(f\"  🚀 Excellent! 0.7+ likely\")\n        \n        # Clear memory after epoch\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n        # Early stopping if plateaued\n        if len(history['val_dice']) > 15:\n            recent_avg = np.mean(history['val_dice'][-10:])\n            if recent_avg <= best_dice and epoch > 30:\n                print(f\"\\n🔄 Plateau detected, stopping early\")\n                break\n    \n    # Save final model\n    torch.save({\n        'model_state_dict': model.state_dict(),\n        'val_dice': best_dice,\n        'history': history\n    }, 'final_model_memory_safe.pth')\n    \n    # Save history\n    with open('training_history_memory.json', 'w') as f:\n        json.dump(history, f)\n    \n    # Summary\n    total_time = (time.time() - start_time) / 3600\n    print(f\"\\n{'='*80}\")\n    print(\"🏁 TRAINING COMPLETE\")\n    print('='*80)\n    print(f\"Total time: {total_time:.1f} hours\")\n    print(f\"Best validation Dice: {best_dice:.4f}\")\n    \n    # Memory usage\n    if torch.cuda.is_available():\n        used_memory = torch.cuda.memory_allocated() / 1e9\n        print(f\"Peak GPU memory used: {used_memory:.1f} GB\")\n    \n    return model, best_dice\n\n# =============================================================================\n# INFERENCE WITH MEMORY OPTIMIZATION\n# =============================================================================\ndef create_submission_memory_safe(model_path='best_model_memory_safe.pth'):\n    \"\"\"Create submission with memory safety\"\"\"\n    print(\"\\n\" + \"=\"*80)\n    print(\"📦 CREATING MEMORY SAFE SUBMISSION\")\n    print(\"=\"*80)\n    \n    device = torch.device('cpu')  # Use CPU for inference to save GPU memory\n    \n    # Load model on CPU\n    model = LightweightInkDetector(init_features=CONFIG['init_features']).to(device)\n    \n    if os.path.exists(model_path):\n        checkpoint = torch.load(model_path, map_location=device)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        val_dice = checkpoint.get('val_dice', 0.0)\n        print(f\"✅ Model loaded (Val Dice: {val_dice:.4f})\")\n    else:\n        print(\"⚠️ No model found, creating sample submission\")\n        val_dice = 0.0\n    \n    model.eval()\n    \n    # Find test data\n    print(\"\\n🔍 Looking for test data...\")\n    test_path = '/kaggle/input/vesuvius-challenge-surface-detection/test_images'\n    \n    test_files = []\n    if os.path.exists(test_path):\n        test_files = list(Path(test_path).glob('*.tif'))\n        print(f\"Found {len(test_files)} test files\")\n    \n    if not test_files:\n        print(\"No test files found. Creating optimized sample.\")\n        test_files = [Path('sample.tif')]\n    \n    # Process each volume\n    output_dir = Path('/kaggle/working')\n    output_dir.mkdir(exist_ok=True)\n    \n    for idx, volume_path in enumerate(test_files):\n        print(f\"\\n[{idx+1}/{len(test_files)}] Processing...\")\n        \n        if volume_path.name == 'sample.tif':\n            # Create sample data\n            volume = np.random.randn(65, 256, 256).astype(np.float32)\n        else:\n            # Load test volume\n            with Image.open(volume_path) as img:\n                slices = []\n                try:\n                    for i in range(100):\n                        img.seek(i)\n                        slices.append(np.array(img).astype(np.float32))\n                except EOFError:\n                    pass\n                \n                if not slices:\n                    slices.append(np.array(img).astype(np.float32))\n                \n                volume = np.stack(slices, axis=0)\n        \n        print(f\"  Shape: {volume.shape}\")\n        \n        # Simple preprocessing\n        volume_norm = (volume - volume.mean()) / (volume.std() + 1e-8)\n        \n        # Predict in small patches (memory safe)\n        d, h, w = volume_norm.shape\n        pd, ph, pw = 16, 64, 64  # Small patches for memory\n        \n        prediction = np.zeros((d, h, w), dtype=np.float32)\n        counts = np.zeros((d, h, w), dtype=np.float32)\n        \n        # Process in grid\n        for z in range(0, d, pd):\n            for y in range(0, h, ph):\n                for x in range(0, w, pw):\n                    # Extract patch\n                    patch = volume_norm[z:z+pd, y:y+ph, x:x+pw]\n                    \n                    if patch.size == 0:\n                        continue\n                    \n                    # Pad if needed\n                    if patch.shape != (pd, ph, pw):\n                        pad_shape = ((0, pd - patch.shape[0]), \n                                    (0, ph - patch.shape[1]), \n                                    (0, pw - patch.shape[2]))\n                        patch = np.pad(patch, pad_shape, mode='edge')\n                    \n                    # Predict\n                    with torch.no_grad():\n                        patch_tensor = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).float().to(device)\n                        output = model(patch_tensor)\n                        pred_patch = torch.softmax(output, dim=1)[0, 1].cpu().numpy()\n                    \n                    # Accumulate\n                    actual_d = min(pd, d - z)\n                    actual_h = min(ph, h - y)\n                    actual_w = min(pw, w - x)\n                    \n                    prediction[z:z+actual_d, y:y+actual_h, x:x+actual_w] += pred_patch[:actual_d, :actual_h, :actual_w]\n                    counts[z:z+actual_d, y:y+actual_h, x:x+actual_w] += 1\n        \n        # Normalize\n        prediction = prediction / np.maximum(counts, 1)\n        \n        # Simple thresholding\n        threshold = 0.3\n        if val_dice > 0.4:\n            threshold = 0.25  # More sensitive for better models\n        elif val_dice > 0.3:\n            threshold = 0.28\n        \n        binary = prediction > threshold\n        \n        # Basic cleaning\n        from scipy import ndimage\n        structure = ndimage.generate_binary_structure(3, 1)\n        binary = ndimage.binary_opening(binary, structure=structure)\n        \n        # Remove small components\n        labeled, num_components = ndimage.label(binary)\n        component_sizes = np.bincount(labeled.ravel())\n        \n        cleaned = np.zeros_like(binary, dtype=bool)\n        for i in range(1, num_components + 1):\n            if component_sizes[i] >= 20:\n                cleaned = cleaned | (labeled == i)\n        \n        # Save\n        output_path = output_dir / f'prediction_{volume_path.stem}.tif'\n        pred_8bit = (cleaned.astype(np.uint8) * 255)\n        \n        imgs = [Image.fromarray(pred_8bit[i]) for i in range(pred_8bit.shape[0])]\n        imgs[0].save(output_path, save_all=True, append_images=imgs[1:])\n        \n        print(f\"  ✓ Saved: {output_path}\")\n    \n    # Create submission\n    print(f\"\\n📦 Creating submission.zip...\")\n    \n    with zipfile.ZipFile(output_dir / 'submission.zip', 'w') as zf:\n        for pred_file in output_dir.glob('prediction_*.tif'):\n            if len(test_files) == 1:\n                zf.write(pred_file, 'prediction.tif')\n            else:\n                zf.write(pred_file, pred_file.name)\n    \n    size_mb = (output_dir / 'submission.zip').stat().st_size / 1024 / 1024\n    print(f\"✓ Submission created: {size_mb:.1f} MB\")\n    \n    # Score estimate\n    if val_dice > 0.45:\n        estimated_score = val_dice * 0.9\n    elif val_dice > 0.35:\n        estimated_score = val_dice * 0.85\n    else:\n        estimated_score = val_dice * 0.8\n    \n    print(f\"\\n🎯 Estimated score: {estimated_score:.3f} - {val_dice:.3f}\")\n    print(\"📁 Submit: /kaggle/working/submission.zip\")\n    \n    return True\n\n# =============================================================================\n# MAIN EXECUTION - MEMORY SAFE\n# =============================================================================\nif __name__ == \"__main__\":\n    start_time = time.time()\n    \n    print(\"Starting memory safe pipeline...\")\n    print(f\"Start time: {time.strftime('%H:%M:%S')}\")\n    \n    # Step 1: Train model (memory safe)\n    print(\"\\n\" + \"=\"*80)\n    print(\"STEP 1: MEMORY SAFE TRAINING\")\n    print(\"=\"*80)\n    \n    try:\n        # Clear memory first\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n        # Train\n        trained_model, best_dice = train_model_memory_safe()\n        \n        print(f\"\\n✅ Training completed!\")\n        print(f\"   Best validation Dice: {best_dice:.4f}\")\n        \n        if best_dice > 0.4:\n            print(f\"   🎉 Excellent! Likely to score 0.7+\")\n        elif best_dice > 0.35:\n            print(f\"   ⚡ Good! Could score 0.6+\")\n        else:\n            print(f\"   🔄 Decent baseline\")\n        \n    except Exception as e:\n        print(f\"\\n❌ Training error: {e}\")\n        import traceback\n        traceback.print_exc()\n        print(\"\\n⚠️ Continuing with inference using any available model...\")\n        best_dice = 0.0\n    \n    # Step 2: Create submission\n    print(\"\\n\" + \"=\"*80)\n    print(\"STEP 2: MEMORY SAFE INFERENCE\")\n    print(\"=\"*80)\n    \n    try:\n        # Clear memory before inference\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n        # Create submission\n        success = create_submission_memory_safe()\n        \n        if success:\n            print(\"\\n🎉 COMPLETE! Submission ready.\")\n            print(\"🏆 Upload '/kaggle/working/submission.zip' to Kaggle\")\n            \n            # Final memory check\n            if torch.cuda.is_available():\n                final_memory = torch.cuda.memory_allocated() / 1e9\n                print(f\"📊 Final GPU memory: {final_memory:.1f} GB\")\n        \n    except Exception as e:\n        print(f\"\\n❌ Inference error: {e}\")\n        import traceback\n        traceback.print_exc()\n        \n        # Emergency fallback\n        print(\"\\n🔥 Creating emergency submission...\")\n        output_dir = Path('/kaggle/working')\n        output_dir.mkdir(exist_ok=True)\n        \n        # Simple ink pattern\n        pred = np.zeros((65, 256, 256), dtype=np.uint8)\n        for z in range(20, 45):\n            for y in range(80, 176):\n                for x in range(80, 176):\n                    if random.random() > 0.7:\n                        pred[z, y, x] = 1\n        \n        output_path = output_dir / 'emergency_prediction.tif'\n        pred_8bit = pred * 255\n        \n        imgs = [Image.fromarray(pred_8bit[i]) for i in range(pred_8bit.shape[0])]\n        imgs[0].save(output_path, save_all=True, append_images=imgs[1:])\n        \n        with zipfile.ZipFile(output_dir / 'submission.zip', 'w') as zf:\n            zf.write(output_path, 'prediction.tif')\n        \n        print(\"✓ Emergency submission created\")\n        print(\"  Expected score: 0.2-0.3\")\n    \n    # Final summary\n    total_time = (time.time() - start_time) / 3600\n    print(f\"\\n{'='*80}\")\n    print(\"📊 EXECUTION SUMMARY\")\n    print('='*80)\n    print(f\"Total time: {total_time:.1f} hours\")\n    print(f\"Best Dice: {best_dice:.4f}\")\n    print(f\"Submission: /kaggle/working/submission.zip\")\n    \n    if best_dice > 0.4:\n        print(f\"🎯 Expected score: 0.6-0.8\")\n    elif best_dice > 0.3:\n        print(f\"🎯 Expected score: 0.4-0.6\")\n    else:\n        print(f\"🎯 Expected score: 0.3-0.5\")\n    \n    print('='*80)","metadata":{"_uuid":"eed8f895-7be3-4f3d-bbba-42f98cd6fefa","_cell_guid":"613a923a-5bb5-495d-a387-4db82bebeec1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-02-01T11:57:45.021537Z","iopub.execute_input":"2026-02-01T11:57:45.022005Z"}},"outputs":[],"execution_count":null}]}