{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":114201,"databundleVersionId":13622514,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Multi-Modal Fusion Baseline for GeoPlant Challenge\n\n## Overview\nThis notebook implements a multi-modal deep learning approach for the GeoPlant@PAISS challenge. The goal is to predict plant species that should grow at a location given GPS coordinates and various predictors, e.g., satellite images, climatic time series, land cover, human footprint, etc.\n\n## Key Components of this notebook:\n1. **Multi-Modal Data Sources**:\n   - **Sentinel-2 patches**: 64x64 satellite imagery with 4 channels (RGB + NIR)\n   - **Landsat time series**: Historical satellite data with 6 bands × 4 quarters × 21 years\n   - **Bioclimatic cubes**: Environmental variables with 4 variables × 12 months × 19 years\n\n2. **Model Architecture**:\n   - Separate encoders for each data modality\n   - Cross-modal attention mechanism to learn relationships between modalities\n   - SE (Squeeze-and-Excitation) blocks for channel attention\n   - Multi-label classification for 11,255 plant species\n\n3. **Training Techniques**:\n   - Asymmetric Loss to handle class imbalance in multi-label setting\n   - MixUp augmentation for better generalization\n   - Gradient clipping for training stability\n   - OneCycleLR scheduling for optimal learning rate\n\n4. **Robustness Features**:\n   - Checkpoint saving/loading for resuming interrupted training","metadata":{}},{"cell_type":"markdown","source":"## Dependencies\n- rasterio : a Python library to access geospatial raster data\n- torchview : for model architecture visualization\n- PyTorch, numpy, tqdm and pandas","metadata":{}},{"cell_type":"code","source":"!pip install rasterio torchview -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:14:55.395776Z","iopub.execute_input":"2025-09-16T16:14:55.396273Z","iopub.status.idle":"2025-09-16T16:15:02.108088Z","shell.execute_reply.started":"2025-09-16T16:14:55.39624Z","shell.execute_reply":"2025-09-16T16:15:02.106901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport rasterio\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:15:02.109691Z","iopub.execute_input":"2025-09-16T16:15:02.110065Z","iopub.status.idle":"2025-09-16T16:15:08.592485Z","shell.execute_reply.started":"2025-09-16T16:15:02.110027Z","shell.execute_reply":"2025-09-16T16:15:08.591476Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss Functions","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    \"\"\"Focal Loss for addressing class imbalance\"\"\"\n    def __init__(self, alpha=0.25, gamma=2):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        \n    def forward(self, inputs, targets):\n        bce_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-bce_loss)\n        focal_loss = self.alpha * (1-pt)**self.gamma * bce_loss\n        return focal_loss.mean()\n\nclass AsymmetricLoss(nn.Module):\n    \"\"\"Asymmetric Loss for multi-label classification\"\"\"\n    def __init__(self, gamma_neg=4, gamma_pos=1, clip=0.05, eps=1e-8):\n        super().__init__()\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.eps = eps\n\n    def forward(self, x, y):\n        x_sigmoid = torch.sigmoid(x)\n        xs_pos = x_sigmoid\n        xs_neg = 1 - x_sigmoid\n\n        if self.clip is not None and self.clip > 0:\n            xs_neg = (xs_neg + self.clip).clamp(max=1)\n\n        los_pos = y * torch.log(xs_pos.clamp(min=self.eps))\n        los_neg = (1 - y) * torch.log(xs_neg.clamp(min=self.eps))\n        loss = los_pos + los_neg\n\n        pt0 = xs_pos * y\n        pt1 = xs_neg * (1 - y)\n        pt = pt0 + pt1\n        one_sided_gamma = self.gamma_pos * y + self.gamma_neg * (1 - y)\n        one_sided_w = torch.pow(1 - pt, one_sided_gamma)\n        loss *= one_sided_w\n\n        return -loss.sum(dim=1).mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:15:08.593472Z","iopub.execute_input":"2025-09-16T16:15:08.593879Z","iopub.status.idle":"2025-09-16T16:15:08.605359Z","shell.execute_reply.started":"2025-09-16T16:15:08.593855Z","shell.execute_reply":"2025-09-16T16:15:08.603963Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Multi-Modal Dataset","metadata":{}},{"cell_type":"code","source":"def construct_patch_path(data_path, survey_id):\n    \"\"\"Construct Sentinel-2 patch path\"\"\"\n    path = data_path\n    for d in (str(survey_id)[-2:], str(survey_id)[-4:-2]):\n        path = os.path.join(path, d)\n    return os.path.join(path, f\"{survey_id}.tiff\")\n\nclass MultiModalDataset(Dataset):\n    def __init__(self, sentinel_dir, landsat_dir, bioclim_dir, metadata, subset, is_train=True, mixup_alpha=1.0):\n        self.sentinel_dir = sentinel_dir\n        self.landsat_dir = landsat_dir\n        self.bioclim_dir = bioclim_dir\n        self.metadata = metadata\n        self.subset = subset\n        self.is_train = is_train\n        self.mixup_alpha = mixup_alpha if is_train else 0\n        \n        if is_train:\n            self.metadata = self.metadata.dropna(subset=\"speciesId\").reset_index(drop=True)\n            self.metadata['speciesId'] = self.metadata['speciesId'].astype(int)\n            self.label_dict = self.metadata.groupby('surveyId')['speciesId'].apply(list).to_dict()\n            self.metadata = self.metadata.drop_duplicates(subset=\"surveyId\").reset_index(drop=True)\n            self.num_classes = 11255\n\n    def __len__(self):\n        return len(self.metadata)\n    \n    def normalize_data(self, data):\n        \"\"\"Simple normalization\"\"\"\n        return (data - data.mean()) / (data.std() + 1e-8)\n\n    def __getitem__(self, idx):\n        row = self.metadata.iloc[idx]\n        survey_id = row['surveyId']\n        \n        # Load Sentinel-2 patch\n        sentinel_path = construct_patch_path(self.sentinel_dir, survey_id)\n        with rasterio.open(sentinel_path) as dataset:\n            sentinel_data = dataset.read(out_dtype=np.float32)\n            sentinel_data = torch.from_numpy(sentinel_data)\n            sentinel_data = self.normalize_data(sentinel_data)\n        \n        # Load Landsat cube\n        landsat_path = os.path.join(self.landsat_dir, f\"GLC25-PA-{self.subset}-landsat-time-series_{survey_id}_cube.pt\")\n        if self.subset == \"test\":\n            landsat_path = os.path.join(self.landsat_dir, f\"GLC25-PA-{self.subset}-landsat_time_series_{survey_id}_cube.pt\")\n        landsat_data = torch.nan_to_num(torch.load(landsat_path, weights_only=True))\n        landsat_data = self.normalize_data(landsat_data)\n        \n        # Load Bioclim cube\n        bioclim_path = os.path.join(self.bioclim_dir, f\"GLC25-PA-{self.subset}-bioclimatic_monthly_{survey_id}_cube.pt\")\n        bioclim_data = torch.load(bioclim_path, weights_only=True)\n        bioclim_data = self.normalize_data(bioclim_data)\n        \n        if self.is_train:\n            # Create multi-label target\n            species_ids = self.label_dict.get(survey_id, [])\n            label = torch.zeros(self.num_classes)\n            for species_id in species_ids:\n                if species_id < self.num_classes:\n                    label[species_id] = 1.0\n            return sentinel_data, landsat_data, bioclim_data, label, survey_id\n        else:\n            return sentinel_data, landsat_data, bioclim_data, survey_id\n\nclass MixUpCollate:\n    \"\"\"Collate function with MixUp augmentation\"\"\"\n    def __init__(self, alpha=1.0, p=0.5):\n        self.alpha = alpha\n        self.p = p\n\n    def __call__(self, batch):\n        if len(batch[0]) == 5:  # Training data\n            sentinel = torch.stack([item[0] for item in batch])\n            landsat = torch.stack([item[1] for item in batch])\n            bioclim = torch.stack([item[2] for item in batch])\n            labels = torch.stack([item[3] for item in batch])\n            survey_ids = [item[4] for item in batch]\n            \n            # Apply MixUp\n            if np.random.random() < self.p:\n                batch_size = sentinel.size(0)\n                lam = np.random.beta(self.alpha, self.alpha)\n                indices = torch.randperm(batch_size)\n                \n                sentinel = lam * sentinel + (1 - lam) * sentinel[indices]\n                landsat = lam * landsat + (1 - lam) * landsat[indices]\n                bioclim = lam * bioclim + (1 - lam) * bioclim[indices]\n                labels = lam * labels + (1 - lam) * labels[indices]\n            \n            return sentinel, landsat, bioclim, labels, survey_ids\n        else:  # Test data\n            sentinel = torch.stack([item[0] for item in batch])\n            landsat = torch.stack([item[1] for item in batch])\n            bioclim = torch.stack([item[2] for item in batch])\n            survey_ids = [item[3] for item in batch]\n            return sentinel, landsat, bioclim, survey_ids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:15:08.607473Z","iopub.execute_input":"2025-09-16T16:15:08.607754Z","iopub.status.idle":"2025-09-16T16:15:08.63318Z","shell.execute_reply.started":"2025-09-16T16:15:08.607726Z","shell.execute_reply":"2025-09-16T16:15:08.632046Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architecture with SE Blocks and Cross-Modal Attention","metadata":{}},{"cell_type":"code","source":"class SEBlock(nn.Module):\n    \"\"\"Squeeze-and-Excitation block\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.squeeze = nn.AdaptiveAvgPool2d(1)\n        self.excitation = nn.Sequential(\n            nn.Linear(channels, channels // reduction, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(channels // reduction, channels, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, _, _ = x.size()\n        y = self.squeeze(x).view(b, c)\n        y = self.excitation(y).view(b, c, 1, 1)\n        return x * y.expand_as(x)\n\nclass BasicBlockSE(nn.Module):\n    \"\"\"Basic residual block with SE\"\"\"\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, 3, stride, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, 3, 1, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        self.se = SEBlock(out_channels)\n        \n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, 1, stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        \n    def forward(self, x):\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out = self.se(out)\n        out += self.shortcut(x)\n        return F.relu(out)\n\nclass CrossModalAttention(nn.Module):\n    \"\"\"Cross-modal attention mechanism\"\"\"\n    def __init__(self, dim, num_heads=8):\n        super().__init__()\n        self.num_heads = num_heads\n        self.dim = dim\n        self.head_dim = dim // num_heads\n        \n        self.qkv = nn.Linear(dim, dim * 3)\n        self.proj = nn.Linear(dim, dim)\n        self.norm1 = nn.LayerNorm(dim)\n        self.norm2 = nn.LayerNorm(dim)\n        self.ffn = nn.Sequential(\n            nn.Linear(dim, dim * 4),\n            nn.GELU(),\n            nn.Linear(dim * 4, dim)\n        )\n        \n    def forward(self, x):\n        # x shape: (batch, num_modalities, dim)\n        B, N, C = x.shape\n        \n        # Self-attention\n        x_norm = self.norm1(x)\n        qkv = self.qkv(x_norm).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n        \n        attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5)\n        attn = attn.softmax(dim=-1)\n        \n        out = (attn @ v).transpose(1, 2).reshape(B, N, C)\n        x = x + self.proj(out)\n        \n        # FFN\n        x = x + self.ffn(self.norm2(x))\n        \n        return x\n\nclass MultiModalFusionModel(nn.Module):\n    \"\"\"Multi-modal fusion model with cross-modal attention\"\"\"\n    def __init__(self, num_classes=11255):\n        super().__init__()\n        \n        # Sentinel-2 encoder (4 channels, 64x64)\n        self.sentinel_encoder = nn.Sequential(\n            nn.Conv2d(4, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            BasicBlockSE(64, 128, 2),  # 32x32\n            BasicBlockSE(128, 256, 2),  # 16x16\n            BasicBlockSE(256, 512, 2),  # 8x8\n            nn.AdaptiveAvgPool2d(1)\n        )\n        \n        # Landsat encoder (6 channels, 4x21)\n        self.landsat_encoder = nn.Sequential(\n            nn.Conv2d(6, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            BasicBlockSE(64, 128),\n            BasicBlockSE(128, 256),\n            BasicBlockSE(256, 512),\n            nn.AdaptiveAvgPool2d(1)\n        )\n        \n        # Bioclim encoder (4 channels, 19x12)\n        self.bioclim_encoder = nn.Sequential(\n            nn.Conv2d(4, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            BasicBlockSE(64, 128),\n            BasicBlockSE(128, 256),\n            BasicBlockSE(256, 512),\n            nn.AdaptiveAvgPool2d(1)\n        )\n        \n        # Cross-modal attention\n        self.cross_modal_attention = CrossModalAttention(512, num_heads=8)\n        \n        # Final classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(512 * 3, 1024),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(1024, num_classes)\n        )\n        \n    def forward(self, sentinel, landsat, bioclim):\n        # Encode each modality\n        sentinel_feat = self.sentinel_encoder(sentinel).squeeze(-1).squeeze(-1)  # (B, 512)\n        landsat_feat = self.landsat_encoder(landsat).squeeze(-1).squeeze(-1)    # (B, 512)\n        bioclim_feat = self.bioclim_encoder(bioclim).squeeze(-1).squeeze(-1)    # (B, 512)\n        \n        # Stack features for cross-modal attention\n        features = torch.stack([sentinel_feat, landsat_feat, bioclim_feat], dim=1)  # (B, 3, 512)\n        \n        # Apply cross-modal attention\n        features = self.cross_modal_attention(features)\n        \n        # Concatenate all features\n        combined_features = features.reshape(features.size(0), -1)  # (B, 3*512)\n        \n        # Classification\n        logits = self.classifier(combined_features)\n        \n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:15:08.634174Z","iopub.execute_input":"2025-09-16T16:15:08.634428Z","iopub.status.idle":"2025-09-16T16:15:08.65938Z","shell.execute_reply.started":"2025-09-16T16:15:08.634407Z","shell.execute_reply":"2025-09-16T16:15:08.658451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model architecture visualization\nfrom torchview import draw_graph\n\nmodel = MultiModalFusionModel(num_classes=11255)\n# Shapes RÉELLES de votre modèle\nmodel_graph = draw_graph(\n    model, \n    input_size=[\n        (32, 4, 64, 64),   # Sentinel: 4 bands, 64x64 pixels\n        (32, 6, 4, 21),    # Landsat: 6 bands, 4 quarters × 21 years\n        (32, 4, 19, 12)    # Bioclim: 4 variables, 19 years × 12 months\n    ]\n)\nmodel_graph.visual_graph","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T17:00:11.110934Z","iopub.execute_input":"2025-09-16T17:00:11.111244Z","iopub.status.idle":"2025-09-16T17:00:12.624738Z","shell.execute_reply.started":"2025-09-16T17:00:11.111224Z","shell.execute_reply":"2025-09-16T17:00:12.623959Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Setup","metadata":{}},{"cell_type":"code","source":"# Configuration\nbatch_size = 32\nnum_workers = 4\nnum_classes = 11255\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Data paths\ntrain_sentinel_path = \"/kaggle/input/geoplant-at-paiss/SatelitePatches/PA-train\"\ntrain_landsat_path = \"/kaggle/input/geoplant-at-paiss/SateliteTimeSeries-Landsat/cubes/PA-train/\"\ntrain_bioclim_path = \"/kaggle/input/geoplant-at-paiss/BioclimTimeSeries/cubes/PA-train\"\ntrain_metadata = pd.read_csv(\"/kaggle/input/geoplant-at-paiss/GLC25_PA_metadata_train.csv\")\n\ntest_sentinel_path = \"/kaggle/input/geoplant-at-paiss/SatelitePatches/PA-test/\"\ntest_landsat_path = \"/kaggle/input/geoplant-at-paiss/SateliteTimeSeries-Landsat/cubes/PA-test/\"\ntest_bioclim_path = \"/kaggle/input/geoplant-at-paiss/BioclimTimeSeries/cubes/PA-test\"\ntest_metadata = pd.read_csv(\"/kaggle/input/geoplant-at-paiss/GLC25_PA_metadata_test.csv\")\n\n# Create datasets\ntrain_dataset = MultiModalDataset(\n    train_sentinel_path, train_landsat_path, train_bioclim_path,\n    train_metadata, subset=\"train\", is_train=True\n)\n\ntest_dataset = MultiModalDataset(\n    test_sentinel_path, test_landsat_path, test_bioclim_path,\n    test_metadata, subset=\"test\", is_train=False\n)\n\n# Data loaders with MixUp\ntrain_loader = DataLoader(\n    train_dataset, batch_size=batch_size, shuffle=True, \n    num_workers=num_workers, collate_fn=MixUpCollate(alpha=1.0, p=0.5)\n)\n\ntest_loader = DataLoader(\n    test_dataset, batch_size=batch_size, shuffle=False, \n    num_workers=num_workers, collate_fn=MixUpCollate(alpha=0, p=0)\n)\n\nprint(f\"Training samples: {len(train_dataset)}\")\nprint(f\"Test samples: {len(test_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-16T16:20:45.757553Z","iopub.execute_input":"2025-09-16T16:20:45.757869Z","iopub.status.idle":"2025-09-16T16:20:49.714711Z","shell.execute_reply.started":"2025-09-16T16:20:45.757847Z","shell.execute_reply":"2025-09-16T16:20:49.713816Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"# Initialize model\nmodel = MultiModalFusionModel(num_classes=num_classes).to(device)\n\n# Loss and optimizer\ncriterion = AsymmetricLoss(gamma_neg=4, gamma_pos=1, clip=0.05)\nn_epochs = 15\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer, max_lr=1e-4, steps_per_epoch=len(train_loader), epochs=n_epochs\n)\n\n# Training\nnum_epochs = n_epochs\nbest_loss = float('inf')\n\nfor epoch in range(num_epochs):\n    model.train()\n    train_loss = 0.0\n    \n    pbar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs}')\n    for batch_idx, (sentinel, landsat, bioclim, labels, _) in enumerate(pbar):\n        sentinel = sentinel.to(device)\n        landsat = landsat.to(device)\n        bioclim = bioclim.to(device)\n        labels = labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(sentinel, landsat, bioclim)\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        scheduler.step()\n        \n        train_loss += loss.item()\n        pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n    \n    avg_loss = train_loss / len(train_loader)\n    print(f'Epoch {epoch+1}/{num_epochs}, Average Loss: {avg_loss:.4f}')\n    \n    # Save best model\n    if avg_loss < best_loss:\n        best_loss = avg_loss\n        torch.save(model.state_dict(), 'best_multimodal_model.pth')\n        print(f'Saved best model with loss: {best_loss:.4f}')\n\nprint(\"Training completed!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Generate Predictions","metadata":{}},{"cell_type":"code","source":"# Load best model\nmodel.load_state_dict(torch.load('best_multimodal_model.pth'))\nmodel.eval()\n\n# Generate predictions\npredictions = []\nsurvey_ids = []\n\nwith torch.no_grad():\n    for sentinel, landsat, bioclim, surveyId in tqdm(test_loader, desc=\"Generating predictions\"):\n        sentinel = sentinel.to(device)\n        landsat = landsat.to(device)\n        bioclim = bioclim.to(device)\n        \n        outputs = model(sentinel, landsat, bioclim)\n        probs = torch.sigmoid(outputs)\n        \n        predictions.append(probs.cpu().numpy())\n        survey_ids.extend(surveyId)\n\n# Concatenate predictions\npredictions = np.vstack(predictions)\n\n# Test different top-k values - try smaller values first\ntop_k_values = [20, 25, 30]\n\nfor top_k in top_k_values:\n    print(f\"\\\\n=== Testing top-k = {top_k} ===\")\n    \n    # Get top-k predictions\n    top_k_indices = np.argsort(-predictions, axis=1)[:, :top_k]\n    \n    # Format submission\n    data_concatenated = []\n    for row in top_k_indices:\n        sorted_species = sorted(row.tolist())\n        data_concatenated.append(' '.join(map(str, sorted_species)))\n    \n    # Create submission\n    submission_df = pd.DataFrame({\n        'surveyId': survey_ids,\n        'predictions': data_concatenated\n    })\n    \n    # Save with different filename for each k\n    filename = f\"submission_top{top_k}.csv\"\n    submission_df.to_csv(filename, index=False)\n    print(f\"Submission saved as {filename}\")\n    \n    # Show sample prediction length\n    sample_pred = data_concatenated[0].split()\n    print(f\"Sample prediction length: {len(sample_pred)} species\")\n    print(f\"First few species: {' '.join(sample_pred[:10])}...\")\n\n# Create final submission\nfinal_top_k = 30\ntop_k_indices = np.argsort(-predictions, axis=1)[:, :final_top_k]\ndata_concatenated = [' '.join(map(str, sorted(row.tolist()))) for row in top_k_indices]\n\nsubmission_df = pd.DataFrame({\n    'surveyId': survey_ids,\n    'predictions': data_concatenated\n})\n\nsubmission_df.to_csv(\"submission.csv\", index=False)\nprint(f\"\\\\n=== FINAL SUBMISSION ===\")\nprint(f\"Using top-k={final_top_k} for final submission.csv\")\nprint(f\"Total predictions: {len(submission_df)}\")\nprint(\"\\\\nFirst few predictions:\")\nprint(submission_df.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}