{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":114201,"databundleVersionId":13622514,"sourceType":"competition"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":35265.370426,"end_time":"2025-04-12T04:18:24.491981","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-04-11T18:30:39.121555","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"e76c598c","cell_type":"markdown","source":"## Simple baseline with Sentinel-2 data — ResNet-6 + Binary Cross Entropy\n\nThe occurrence of different types of organisms, whether plants or animals, is generally associated with the characteristics of the environment or ecosystem in which they live. This relationship between the presence of species and their habitat is often interdependent and can be affected by various factors, such as climate, which is another modality we provide.\n\nTo demonstrate the performance while using just the _image data_, i.e., Sentinel Image Patches, we provide a straightforward baseline that is based on a custom ResNet6-like architecture and Binary Cross-Entropy. \nAs described above, the satellite patches provide an image-like modality that captures habitats and other aspects of the locality.\n\nConsidering the significant extent of enhancing the performance of this baseline, we encourage you to experiment with various techniques, architectures, losses, etc.\n\n#### **Have Fun!**","metadata":{"papermill":{"duration":0.004282,"end_time":"2025-04-11T18:30:41.945879","exception":false,"start_time":"2025-04-11T18:30:41.941597","status":"completed"},"tags":[]}},{"id":"3ddbd166-2be6-4dfe-82c3-9fb1b4a11d4b","cell_type":"code","source":"!pip install rasterio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-05T05:48:21.677694Z","iopub.execute_input":"2025-09-05T05:48:21.678176Z","iopub.status.idle":"2025-09-05T05:48:28.038525Z","shell.execute_reply.started":"2025-09-05T05:48:21.678153Z","shell.execute_reply":"2025-09-05T05:48:28.037869Z"}},"outputs":[],"execution_count":null},{"id":"55c27752","cell_type":"code","source":"import os\nimport torch\nimport tqdm\nimport rasterio\nimport numpy as np\nimport pandas as pd\nimport albumentations as A\nimport torchvision.models as models\nimport torchvision.transforms as transforms\nimport torch.nn as nn\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom sklearn.metrics import precision_recall_fscore_support","metadata":{"ExecuteTime":{"end_time":"2024-04-30T21:25:07.29831Z","start_time":"2024-04-30T21:25:05.354584Z"},"execution":{"iopub.status.busy":"2025-09-05T05:48:28.040058Z","iopub.execute_input":"2025-09-05T05:48:28.040345Z","iopub.status.idle":"2025-09-05T05:48:41.640197Z","shell.execute_reply.started":"2025-09-05T05:48:28.040311Z","shell.execute_reply":"2025-09-05T05:48:41.639593Z"},"papermill":{"duration":9.684855,"end_time":"2025-04-11T18:30:51.634701","exception":false,"start_time":"2025-04-11T18:30:41.949846","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e91674b2","cell_type":"markdown","source":"## Data description\n\nThe Sentinel-2 data was acquired through the Sentinel2 satellite program and pre-processed by [Ecodatacube](https://stac.ecodatacube.eu/) to produce raster files scaled to the entire European continent and projected into a unique CRS. \nEach TIFF file corresponds to a unique observation location (via \"surveyId\"). To load the patches for a selected observation, take the \"surveyId\" from any occurrence CSV and load it following this rule --> '…/CD/AB/XXXXABCD.jpeg'. For example, the image location for the surveyId 3018575 is \"./75/85/3018575.tiff\". For all \"surveyId\" with less than four digits, you can use a similar rule. For a \"surveyId\" 1 is \"./1/1.tiff\".\nThe data can simply be loaded using the following method:\n\n```python\ndef construct_patch_path(output_path, survey_id):\n    \"\"\"Construct the patch file path based on survey_id as './CD/AB/XXXXABCD.tiff'\"\"\"\n    path = output_path\n    for d in (str(survey_id)[-2:], str(survey_id)[-4:-2]):\n        path = os.path.join(path, d)\n\n    path = os.path.join(path, f\"{survey_id}.tiff\")\n\n    return path\n```\n\n**For more information about data processing, normalization, and visualization, please refer to the following notebook**: [Kaggle Notebook](https://www.kaggle.com/code/picekl/sentinel-2-data-processing-and-normalization).\n\n**References:**\n- *Traceability (lineage): The dataset was produced entirely by mosaicking and seasonally aggregating imagery from the Sentinel-2 Level-2A product (https://sentinels.copernicus.eu/web/sentinel/user-guides/sentinel-2-msi/product-types/level-2a)*\n- *Ecodatacube.eu: Analysis-ready open environmental data cube for Europe (https://doi.org/10.21203/rs.3.rs-2277090/v3)*","metadata":{"execution":{"iopub.execute_input":"2024-05-01T13:30:07.054038Z","iopub.status.busy":"2024-05-01T13:30:07.053659Z","iopub.status.idle":"2024-05-01T13:30:07.058148Z","shell.execute_reply":"2024-05-01T13:30:07.057269Z","shell.execute_reply.started":"2024-05-01T13:30:07.054008Z"},"papermill":{"duration":0.003839,"end_time":"2025-04-11T18:30:51.642722","exception":false,"start_time":"2025-04-11T18:30:51.638883","status":"completed"},"tags":[]}},{"id":"5213a77a","cell_type":"markdown","source":"## Prepare custom dataset loader\n\nWe have to slightly update the Dataset to provide the relevant data in the appropriate format.","metadata":{"papermill":{"duration":0.00379,"end_time":"2025-04-11T18:30:51.650496","exception":false,"start_time":"2025-04-11T18:30:51.646706","status":"completed"},"tags":[]}},{"id":"122e2fee","cell_type":"code","source":"def construct_patch_path(data_path, survey_id):\n    \"\"\"Construct the patch file path based on plot_id as './CD/AB/XXXXABCD.jpeg'\"\"\"\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\n    path = os.path.join(path, f\"{survey_id}.tiff\")\n\n    return path\n\ndef quantile_normalize(band, low=2, high=98):\n    sorted_band = np.sort(band.flatten())\n    quantiles = np.percentile(sorted_band, np.linspace(low, high, len(sorted_band)))\n    normalized_band = np.interp(band.flatten(), sorted_band, quantiles).reshape(band.shape)\n    \n    min_val, max_val = np.min(normalized_band), np.max(normalized_band)\n    \n    # Prevent division by zero if min_val == max_val\n    if max_val == min_val:\n        return np.zeros_like(normalized_band, dtype=np.float32)  # Return an array of zeros\n\n    # Perform normalization (min-max scaling)\n    return ((normalized_band - min_val) / (max_val - min_val)).astype(np.float32)\n\nclass TrainDataset(Dataset):\n    def __init__(self, data_dir, metadata, transform=None):\n        self.transform = transform\n        self.data_dir = data_dir\n        self.metadata = metadata\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        \n        self.metadata = self.metadata.drop_duplicates(subset=\"surveyId\").reset_index(drop=True)\n\n    def __len__(self):\n        return len(self.metadata)\n\n    def __getitem__(self, idx):\n        \n        survey_id = self.metadata.surveyId[idx]\n        species_ids = self.label_dict.get(survey_id, [])  # Get list of species IDs for the survey ID\n        label = torch.zeros(num_classes)  # Initialize label tensor\n        for species_id in species_ids:\n            label_id = species_id\n            label[label_id] = 1  # Set the corresponding class index to 1 for each species\n        \n        # Read TIFF files (multispectral bands)\n        tiff_path = construct_patch_path(self.data_dir, survey_id)\n        with rasterio.open(tiff_path) as dataset:\n            image = dataset.read(out_dtype=np.float32)  # Read all bands\n            image = np.array([quantile_normalize(band) for band in image])  # Apply quantile normalization\n\n        image = np.transpose(image, (1, 2, 0))  # Convert to HWC format\n        image = self.transform(image)\n\n        return image, label, survey_id\n    \nclass TestDataset(TrainDataset):\n    def __init__(self, data_dir, metadata, transform=None):\n        self.transform = transform\n        self.data_dir = data_dir\n        self.metadata = metadata\n        \n    def __getitem__(self, idx):\n        \n        survey_id = self.metadata.surveyId[idx]\n        \n        # Read TIFF files (multispectral bands)\n        tiff_path = construct_patch_path(self.data_dir, survey_id)\n        with rasterio.open(tiff_path) as dataset:\n            image = dataset.read(out_dtype=np.float32)  # Read all bands\n            image = np.array([quantile_normalize(band) for band in image])  # Apply quantile normalization\n\n        image = np.transpose(image, (1, 2, 0))  # Convert to HWC format\n        \n        image = self.transform(image)\n        return image, survey_id","metadata":{"ExecuteTime":{"end_time":"2024-04-30T21:25:32.627928Z","start_time":"2024-04-30T21:25:32.612131Z"},"collapsed":false,"execution":{"iopub.status.busy":"2025-09-05T05:48:41.640955Z","iopub.execute_input":"2025-09-05T05:48:41.641373Z","iopub.status.idle":"2025-09-05T05:48:41.653866Z","shell.execute_reply.started":"2025-09-05T05:48:41.641353Z","shell.execute_reply":"2025-09-05T05:48:41.653173Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.022281,"end_time":"2025-04-11T18:30:51.676595","exception":false,"start_time":"2025-04-11T18:30:51.654314","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"201a02b1","cell_type":"markdown","source":"### Load metadata and prepare data loaders","metadata":{"papermill":{"duration":0.003709,"end_time":"2025-04-11T18:30:51.684240","exception":false,"start_time":"2025-04-11T18:30:51.680531","status":"completed"},"tags":[]}},{"id":"5b5fc9db","cell_type":"code","source":"# Dataset and DataLoader\nbatch_size = 128\n\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.5, 0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5, 0.5)),\n])\n\n# Load Training metadata\ntrain_data_path = \"/kaggle/input/geoplant-at-paiss/SatelitePatches/PA-train\"\ntrain_metadata_path = \"/kaggle/input/geoplant-at-paiss/GLC25_PA_metadata_train.csv\"\ntrain_metadata = pd.read_csv(train_metadata_path)\ntrain_dataset = TrainDataset(train_data_path, train_metadata, transform=transform)\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\n\n# Load Test metadata\ntest_data_path = \"/kaggle/input/geoplant-at-paiss/SatelitePatches/PA-test/\"\ntest_metadata_path = \"/kaggle/input/geoplant-at-paiss/GLC25_PA_metadata_test.csv\"\ntest_metadata = pd.read_csv(test_metadata_path)\ntest_dataset = TestDataset(test_data_path, test_metadata, transform=transform)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=4)","metadata":{"ExecuteTime":{"end_time":"2024-04-30T21:25:34.532017Z","start_time":"2024-04-30T21:25:32.615562Z"},"collapsed":false,"execution":{"iopub.status.busy":"2025-09-05T05:48:41.655626Z","iopub.execute_input":"2025-09-05T05:48:41.656253Z","iopub.status.idle":"2025-09-05T05:48:45.614291Z","shell.execute_reply.started":"2025-09-05T05:48:41.656235Z","shell.execute_reply":"2025-09-05T05:48:45.613376Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":5.745402,"end_time":"2025-04-11T18:30:57.433545","exception":false,"start_time":"2025-04-11T18:30:51.688143","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"757e8b59","cell_type":"markdown","source":"## Define custom ResNet6-like architecture for Sentinel-2 images\n\nSentinel-2 patches provide **four spectral channels**: **Red, Green, Blue, and Near-Infrared (NIR)**. Since standard ResNet architectures are designed for **3-channel RGB images**, the input layer must be adapted.\n\n### Key adjustments\n- **Input adaptation**: Replace the first convolutional layer so it accepts **4 input channels** instead of 3, allowing the model to fully exploit the NIR band.\n- **Preserve architecture**: The rest of the ResNet backbone (ResNet-6 or ResNet-18) remains unchanged and can still leverage pretrained weights.\n- **Benefit of NIR**: The NIR band is highly informative for vegetation and biodiversity prediction, making it crucial to include rather than discard.\n\n### Note\nThis is a minimal but essential modification. Without it, the network would ignore the NIR band, losing valuable spectral information for species distribution modeling.","metadata":{"papermill":{"duration":0.004179,"end_time":"2025-04-11T18:30:57.441900","exception":false,"start_time":"2025-04-11T18:30:57.437721","status":"completed"},"tags":[]}},{"id":"7e3ad6be-2177-4ce5-ad64-9a8bc32f1047","cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass BasicBlock(nn.Module):\n    \"\"\"2×(3x3 Conv + BN + ReLU) with optional downsampling on the skip path.\"\"\"\n    def __init__(self, in_c: int, out_c: int, stride: int = 1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_c, out_c, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn1   = nn.BatchNorm2d(out_c)\n        self.conv2 = nn.Conv2d(out_c, out_c, kernel_size=3, stride=1, padding=1, bias=False)\n        self.bn2   = nn.BatchNorm2d(out_c)\n\n        self.downsample = None\n        if stride != 1 or in_c != out_c:\n            self.downsample = nn.Sequential(\n                nn.Conv2d(in_c, out_c, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_c)\n            )\n\n    def forward(self, x):\n        identity = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = F.relu(out, inplace=True)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        out += identity\n        out = F.relu(out, inplace=True)\n        return out\n\nclass ResNet6(nn.Module):\n    \"\"\"\n    ResNet-6 for 4×64×64 Sentinel-2 patches.\n      - Input: [B, 4, 64, 64]\n      - Stem: 3×3 conv with stride=2 -> 32×32\n      - Blocks:\n          * Block1: 64 → 64, stride=1 (32×32)\n          * Block2: 64 → 128, stride=2 (16×16)\n          * Block3: 128 → 128, stride=1 (16×16)\n      - GAP + MLP head -> logits [B, num_classes]\n    \"\"\"\n    def __init__(self, num_classes: int, stem_channels: int = 64, mlp_hidden: int = 512, p_drop: float = 0.1):\n        super().__init__()\n\n        # Normalize the 4×64×64 tensor per sample (robust for small batches)\n        self.norm_input = nn.LayerNorm([4, 64, 64])\n\n        # Stem: light downsampling to 32×32\n        self.stem = nn.Sequential(\n            nn.Conv2d(4, stem_channels, kernel_size=3, stride=2, padding=1, bias=False),  # 64->32\n            nn.BatchNorm2d(stem_channels),\n            nn.ReLU(inplace=True)\n        )\n\n        # 3 residual blocks (6 conv layers total)\n        self.block1 = BasicBlock(stem_channels, stem_channels, stride=1)   # 32×32\n        self.block2 = BasicBlock(stem_channels, stem_channels*2, stride=2) # 32->16\n        self.block3 = BasicBlock(stem_channels*2, stem_channels*2, stride=1) # 16×16\n\n        self.gap = nn.AdaptiveAvgPool2d(1)\n\n        self.head = nn.Sequential(\n            nn.Flatten(),                                # [B, C, 1, 1] -> [B, C]\n            nn.LayerNorm(stem_channels*2),\n            nn.Linear(stem_channels*2, mlp_hidden),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p_drop),\n            nn.Linear(mlp_hidden, num_classes)          # logits, use BCEWithLogitsLoss\n        )\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.ones_(m.weight); nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.Linear):\n                nn.init.trunc_normal_(m.weight, std=0.02); nn.init.zeros_(m.bias)\n\n    def forward(self, x):\n        # x: [B, 4, 64, 64]\n        x = self.norm_input(x)\n        x = self.stem(x)        # [B, 64, 32, 32]\n\n        x = self.block1(x)      # [B, 64, 32, 32]\n        x = self.block2(x)      # [B, 128, 16, 16]\n        x = self.block3(x)      # [B, 128, 16, 16]\n\n        x = self.gap(x)         # [B, 128, 1, 1]\n        x = self.head(x)        # [B, num_classes] (logits)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-05T05:48:45.615332Z","iopub.execute_input":"2025-09-05T05:48:45.615633Z","iopub.status.idle":"2025-09-05T05:48:45.631326Z","shell.execute_reply.started":"2025-09-05T05:48:45.615611Z","shell.execute_reply":"2025-09-05T05:48:45.630751Z"}},"outputs":[],"execution_count":null},{"id":"85a7f756","cell_type":"code","source":"# Check if cuda is available\ndevice = torch.device(\"cpu\")\n\nif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n    print(\"DEVICE = CUDA\")\n\nnum_classes = 11255 # Number of all unique classes within the PO and PA data.\nmodel = ResNet6(num_classes).to(device)","metadata":{"execution":{"iopub.status.busy":"2025-09-05T05:48:45.632043Z","iopub.execute_input":"2025-09-05T05:48:45.632226Z","iopub.status.idle":"2025-09-05T05:48:46.064515Z","shell.execute_reply.started":"2025-09-05T05:48:45.632213Z","shell.execute_reply":"2025-09-05T05:48:46.063754Z"},"papermill":{"duration":0.933587,"end_time":"2025-04-11T18:30:58.468395","exception":false,"start_time":"2025-04-11T18:30:57.534808","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"9a7710e4-d790-4eb4-b728-74ec15160626","cell_type":"code","source":"# Hyperparameters\nlearning_rate = 0.0002\nnum_epochs = 4\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)\nscheduler = CosineAnnealingLR(optimizer, T_max=num_epochs, verbose=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-05T05:48:46.065292Z","iopub.execute_input":"2025-09-05T05:48:46.065562Z","iopub.status.idle":"2025-09-05T05:48:46.072428Z","shell.execute_reply.started":"2025-09-05T05:48:46.065543Z","shell.execute_reply":"2025-09-05T05:48:46.071769Z"}},"outputs":[],"execution_count":null},{"id":"b7946956","cell_type":"code","source":"def set_seed(seed):\n    # Set seed for Python's built-in random number generator\n    torch.manual_seed(seed)\n    # Set seed for numpy\n    np.random.seed(seed)\n    # Set seed for CUDA if available\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        # Set cuDNN's random number generator seed for deterministic behavior\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\nset_seed(77)","metadata":{"execution":{"iopub.status.busy":"2025-09-05T05:48:46.073257Z","iopub.execute_input":"2025-09-05T05:48:46.073530Z","iopub.status.idle":"2025-09-05T05:48:46.099544Z","shell.execute_reply.started":"2025-09-05T05:48:46.073515Z","shell.execute_reply":"2025-09-05T05:48:46.098909Z"},"papermill":{"duration":0.011755,"end_time":"2025-04-11T18:30:58.484569","exception":false,"start_time":"2025-04-11T18:30:58.472814","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"1dd18b4b","cell_type":"markdown","source":"## Training Loop\n\nNothing special, just a standard Pytorch training loop.","metadata":{"papermill":{"duration":0.003826,"end_time":"2025-04-11T18:30:58.492394","exception":false,"start_time":"2025-04-11T18:30:58.488568","status":"completed"},"tags":[]}},{"id":"146d3632","cell_type":"code","source":"import time\nfrom tqdm import tqdm\nimport torch\n\nprint(f\"Training for {num_epochs} epochs started.\")\nstart_time = time.time()\n\nfor epoch in range(num_epochs):\n    epoch_start = time.time()\n    model.train()\n\n    running_loss = 0.0\n    pbar = tqdm(enumerate(train_loader), total=len(train_loader), desc=f\"Epoch {epoch+1}/{num_epochs}\")\n\n    for batch_idx, (data, targets, _) in pbar:\n        data = data.to(device)\n        targets = targets.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(data)\n\n        criterion = torch.nn.BCEWithLogitsLoss()\n        loss = criterion(outputs, targets)\n\n        loss.backward()\n        optimizer.step()\n\n        # Update running loss\n        running_loss += loss.item()\n        avg_loss = running_loss / (batch_idx + 1)\n\n        # Show loss in the tqdm bar\n        pbar.set_postfix({\n            \"batch_loss\": f\"{loss.item():.4f}\",\n            \"avg_loss\": f\"{avg_loss:.4f}\"\n        })\n\n    scheduler.step()\n    epoch_time = time.time() - epoch_start\n    print(f\"\\nEpoch {epoch+1} finished in {epoch_time:.2f} seconds\")\n    print(\"Scheduler:\", scheduler.state_dict())\n\n# Save the trained model\nmodel.eval()\ntorch.save(model.state_dict(), \"resnet6-with-sentinel2-cubes.pth\")\n\ntotal_time = time.time() - start_time\nprint(f\"Training completed in {total_time/60:.2f} minutes\")","metadata":{"ExecuteTime":{"start_time":"2024-04-30T21:25:34.536634Z"},"collapsed":false,"execution":{"iopub.status.busy":"2025-09-05T05:48:46.100252Z","iopub.execute_input":"2025-09-05T05:48:46.100422Z","iopub.status.idle":"2025-09-05T07:10:04.312889Z","shell.execute_reply.started":"2025-09-05T05:48:46.100409Z","shell.execute_reply":"2025-09-05T07:10:04.312013Z"},"is_executing":true,"jupyter":{"outputs_hidden":false},"papermill":{"duration":34991.420477,"end_time":"2025-04-12T04:14:09.917095","exception":false,"start_time":"2025-04-11T18:30:58.496618","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"994c12d7","cell_type":"markdown","source":"## Test Loop\n\nAgain, nothing special, just a standard inference.","metadata":{"papermill":{"duration":0.013546,"end_time":"2025-04-12T04:14:09.944423","exception":false,"start_time":"2025-04-12T04:14:09.930877","status":"completed"},"tags":[]}},{"id":"65a38793","cell_type":"code","source":"with torch.no_grad():\n    all_predictions = []\n    surveys = []\n    top_k_indices = None\n    for data, surveyID in tqdm(test_loader, total=len(test_loader)):\n\n        data = data.to(device)\n        \n        outputs = model(data)\n        predictions = torch.sigmoid(outputs).cpu().numpy()\n\n        # Sellect top-25 values as predictions\n        top_25 = np.argsort(-predictions, axis=1)[:, :25] \n        if top_k_indices is None:\n            top_k_indices = top_25\n        else:\n            top_k_indices = np.concatenate((top_k_indices, top_25), axis=0)\n\n        surveys.extend(surveyID.cpu().numpy())","metadata":{"collapsed":false,"execution":{"iopub.status.busy":"2025-09-05T07:10:04.315365Z","iopub.execute_input":"2025-09-05T07:10:04.315781Z","iopub.status.idle":"2025-09-05T07:13:50.336341Z","shell.execute_reply.started":"2025-09-05T07:10:04.315757Z","shell.execute_reply":"2025-09-05T07:13:50.335559Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":251.427356,"end_time":"2025-04-12T04:18:21.384873","exception":false,"start_time":"2025-04-12T04:14:09.957517","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"67773ed5","cell_type":"markdown","source":"## Save prediction file! 🎉🥳🙌🤗","metadata":{"papermill":{"duration":0.026488,"end_time":"2025-04-12T04:18:21.438850","exception":false,"start_time":"2025-04-12T04:18:21.412362","status":"completed"},"tags":[]}},{"id":"50ff556b","cell_type":"code","source":"data_concatenated = [' '.join(map(str, row)) for row in top_k_indices]\n\npd.DataFrame(\n    {'surveyId': surveys,\n     'predictions': data_concatenated,\n    }).to_csv(\"submission.csv\", index = False)","metadata":{"execution":{"iopub.status.busy":"2025-09-05T07:13:50.337324Z","iopub.execute_input":"2025-09-05T07:13:50.337528Z","iopub.status.idle":"2025-09-05T07:13:50.572984Z","shell.execute_reply.started":"2025-09-05T07:13:50.337506Z","shell.execute_reply":"2025-09-05T07:13:50.572262Z"},"papermill":{"duration":0.209123,"end_time":"2025-04-12T04:18:21.675144","exception":false,"start_time":"2025-04-12T04:18:21.466021","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4001785f-5f98-4334-9ab6-1c23f59973e2","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}