{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91196,"databundleVersionId":12020672,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install rasterio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:10.627779Z","iopub.execute_input":"2025-06-01T12:23:10.627991Z","iopub.status.idle":"2025-06-01T12:23:15.98054Z","shell.execute_reply.started":"2025-06-01T12:23:10.62796Z","shell.execute_reply":"2025-06-01T12:23:15.979917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport tqdm\nimport rasterio\nimport numpy as np\nimport timm\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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:15.982712Z","iopub.execute_input":"2025-06-01T12:23:15.982954Z","iopub.status.idle":"2025-06-01T12:23:29.344979Z","shell.execute_reply.started":"2025-06-01T12:23:15.982929Z","shell.execute_reply":"2025-06-01T12:23:29.344426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport os\n\nimport matplotlib.pyplot as plt\nimport imageio.v3 as  imageio\nimport matplotlib.image as mpimg\nfrom PIL import Image\n\nimport imageio.v3 as  imageio\nimport albumentations as A\n\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch import nn\nfrom tqdm.notebook import tqdm\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.model_selection import train_test_split\n\nfrom sklearn.model_selection import KFold\nimport torch.nn.functional as F\nfrom torchmetrics import F1Score\nfrom sklearn.preprocessing import MinMaxScaler\n\nimport torch\nimport timm\nimport glob\nimport torchmetrics\nimport time\nimport psutil\nimport os\nimport math\nimport gc\nimport tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:29.345625Z","iopub.execute_input":"2025-06-01T12:23:29.346013Z","iopub.status.idle":"2025-06-01T12:23:33.355226Z","shell.execute_reply.started":"2025-06-01T12:23:29.345995Z","shell.execute_reply":"2025-06-01T12:23:33.354442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_metadata = pd.read_csv(\"/kaggle/input/geolifeclef-2025/GLC25_PA_metadata_train.csv\")\n\ntest_metadata = pd.read_csv(\"/kaggle/input/geolifeclef-2025/GLC25_PA_metadata_test.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:33.356227Z","iopub.execute_input":"2025-06-01T12:23:33.357195Z","iopub.status.idle":"2025-06-01T12:23:35.635447Z","shell.execute_reply.started":"2025-06-01T12:23:33.357154Z","shell.execute_reply":"2025-06-01T12:23:35.634845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unique, counts = np.unique(train_metadata.speciesId.values, return_counts=True)\nprint(len(unique))\n\nnew_unique, new_counts = [], []\nfor u, c in zip(unique, counts):\n    if c > 5:\n        new_unique.append(u)\n        new_counts.append(c)\nunique = np.array(new_unique)\ncounts = np.array(new_counts)\nprint(len(unique))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:35.636289Z","iopub.execute_input":"2025-06-01T12:23:35.636567Z","iopub.status.idle":"2025-06-01T12:23:35.701002Z","shell.execute_reply.started":"2025-06-01T12:23:35.636545Z","shell.execute_reply":"2025-06-01T12:23:35.700389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = len(unique)\nnum_surveys = len(np.unique(train_metadata.surveyId.values))\n\nspecies_dict = {}\nfor i in range(num_classes):\n    species_dict[unique[i]] = i","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:35.701731Z","iopub.execute_input":"2025-06-01T12:23:35.701917Z","iopub.status.idle":"2025-06-01T12:23:35.733208Z","shell.execute_reply.started":"2025-06-01T12:23:35.701903Z","shell.execute_reply":"2025-06-01T12:23:35.732536Z"}},"outputs":[],"execution_count":null},{"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, bioclim_data_dir, landsat_data_dir, sentinel_data_dir, metadata, transform=None):\n        self.transform = transform\n        self.sentinel_transform = A.Compose([\n            A.Rotate(limit=(-10, 10)),\n            A.RandomBrightnessContrast(brightness_limit=(-0.05, 0.05), contrast_limit=(-0.05, 0.05), p=0.3),\n            A.Normalize(mean=(0.5, 0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5, 0.5), max_pixel_value=1),\n            ToTensorV2(),\n        ])\n      \n        self.bioclim_data_dir = bioclim_data_dir\n        self.landsat_data_dir = landsat_data_dir\n        self.sentinel_data_dir = sentinel_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        \n        survey_id = self.metadata.surveyId[idx]\n        \n        landsat_sample = torch.nan_to_num(torch.load(os.path.join(self.landsat_data_dir, f\"GLC25-PA-train-landsat-time-series_{survey_id}_cube.pt\")))\n        bioclim_sample = torch.nan_to_num(torch.load(os.path.join(self.bioclim_data_dir, f\"GLC25-PA-train-bioclimatic_monthly_{survey_id}_cube.pt\")))\n        \n        \n        tiff_path = construct_patch_path(self.sentinel_data_dir, survey_id)\n        with rasterio.open(tiff_path) as dataset:\n            sentinel_sample = dataset.read(out_dtype=np.float32)  # Read all bands\n            sentinel_sample = np.array([quantile_normalize(band) for band in sentinel_sample])  # Apply quantile normalization\n        sentinel_sample = np.transpose(sentinel_sample, (1, 2, 0)) \n\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            if label_id in species_dict.keys():\n                label[species_dict[label_id]] = 1 # Set the corresponding class index to 1 for each species\n        \n        if isinstance(landsat_sample, torch.Tensor):\n            landsat_sample = landsat_sample.permute(1, 2, 0)  # Change tensor shape from (C, H, W) to (H, W, C)\n            landsat_sample = landsat_sample.numpy()  # Convert tensor to numpy array\n            \n        if isinstance(bioclim_sample, torch.Tensor):\n            bioclim_sample = bioclim_sample.permute(1, 2, 0)  # Change tensor shape from (C, H, W) to (H, W, C)\n            bioclim_sample = bioclim_sample.numpy()  # Convert tensor to numpy array   \n        \n        if self.transform:\n            landsat_sample = self.transform(landsat_sample)\n            bioclim_sample = self.transform(bioclim_sample)\n            sentinel_sample = self.sentinel_transform(image=sentinel_sample)['image']\n        \n        \n        return landsat_sample, bioclim_sample, sentinel_sample, label, survey_id\n    \nclass TestDataset(TrainDataset):\n    def __init__(self, bioclim_data_dir, landsat_data_dir, sentinel_data_dir, metadata, transform=None):\n        self.transform = transform\n        self.sentinel_transform = A.Compose([\n            A.Normalize(mean=(0.5, 0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5, 0.5), max_pixel_value=1),\n            ToTensorV2(),\n        ])\n      \n        self.bioclim_data_dir = bioclim_data_dir\n        self.landsat_data_dir = landsat_data_dir\n        self.sentinel_data_dir = sentinel_data_dir\n        self.metadata = metadata\n        \n    def __getitem__(self, idx):\n        \n        survey_id = self.metadata.surveyId[idx]\n        landsat_sample = torch.nan_to_num(torch.load(os.path.join(self.landsat_data_dir, f\"GLC25-PA-test-landsat_time_series_{survey_id}_cube.pt\")))\n        bioclim_sample = torch.nan_to_num(torch.load(os.path.join(self.bioclim_data_dir, f\"GLC25-PA-test-bioclimatic_monthly_{survey_id}_cube.pt\")))\n        \n        \n        tiff_path = construct_patch_path(self.sentinel_data_dir, survey_id)\n        with rasterio.open(tiff_path) as dataset:\n            sentinel_sample = dataset.read(out_dtype=np.float32)  # Read all bands\n            sentinel_sample = np.array([quantile_normalize(band) for band in sentinel_sample])  # Apply quantile normalization\n        sentinel_sample = np.transpose(sentinel_sample, (1, 2, 0)) \n            \n        if isinstance(landsat_sample, torch.Tensor):\n            landsat_sample = landsat_sample.permute(1, 2, 0)  # Change tensor shape from (C, H, W) to (H, W, C)\n            landsat_sample = landsat_sample.numpy()  # Convert tensor to numpy array\n        if isinstance(bioclim_sample, torch.Tensor):\n            bioclim_sample = bioclim_sample.permute(1, 2, 0)  # Change tensor shape from (C, H, W) to (H, W, C)\n            bioclim_sample = bioclim_sample.numpy()  # Convert tensor to numpy array   \n        \n        if self.transform:\n            landsat_sample = self.transform(landsat_sample)\n            bioclim_sample = self.transform(bioclim_sample)\n            sentinel_sample = self.sentinel_transform(image=sentinel_sample)['image']\n        \n        return landsat_sample, bioclim_sample, sentinel_sample, survey_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:23:35.735358Z","iopub.execute_input":"2025-06-01T12:23:35.735577Z","iopub.status.idle":"2025-06-01T12:23:35.752697Z","shell.execute_reply.started":"2025-06-01T12:23:35.73556Z","shell.execute_reply":"2025-06-01T12:23:35.751962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dataset and DataLoader\nbatch_size = 300\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_landsat_data_path = \"/kaggle/input/geolifeclef-2025/SateliteTimeSeries-Landsat/cubes/PA-train/\"\ntrain_bioclim_data_path = \"/kaggle/input/geolifeclef-2025/BioclimTimeSeries/cubes/PA-train/\"\ntrain_sentinel_data_path=\"/kaggle/input/geolifeclef-2025/SatelitePatches/PA-train/\"\ntrain_metadata_path = \"/kaggle/input/geolifeclef-2025/GLC25_PA_metadata_train.csv\"\ntrain_metadata = pd.read_csv(train_metadata_path)\ndataset_alpine = TrainDataset(train_bioclim_data_path, train_landsat_data_path, train_sentinel_data_path, train_metadata, transform=transform)\ntrain_loader = DataLoader(dataset_alpine, batch_size=batch_size, shuffle=True, num_workers=4)\n\n# Load Test metadata\ntest_landsat_data_path = \"/kaggle/input/geolifeclef-2025/SateliteTimeSeries-Landsat/cubes/PA-test/\"\ntest_bioclim_data_path = \"/kaggle/input/geolifeclef-2025/BioclimTimeSeries/cubes/PA-test/\"\ntest_sentinel_data_path = \"/kaggle/input/geolifeclef-2025/SatelitePatches/PA-test/\"\ntest_metadata_path = \"/kaggle/input/geolifeclef-2025/GLC25_PA_metadata_test.csv\"\ntest_metadata = pd.read_csv(test_metadata_path)\ntest_dataset = TestDataset(test_bioclim_data_path, test_landsat_data_path, test_sentinel_data_path, test_metadata, transform=transform)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:40.118395Z","iopub.execute_input":"2025-06-01T12:26:40.119043Z","iopub.status.idle":"2025-06-01T12:26:42.673145Z","shell.execute_reply.started":"2025-06-01T12:26:40.119019Z","shell.execute_reply":"2025-06-01T12:26:42.672606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass MultimodalEnsemble(nn.Module):\n    def __init__(self, num_classes):\n        super(MultimodalEnsemble, self).__init__()\n        \n        self.landsat_norm = nn.LayerNorm([6,4,21])\n        self.landsat_model = models.swin_t(weights=None)\n        # Modify the first convolutional layer to accept 6 channels instead of 3\n        self.landsat_model.features[0][0] = nn.Conv2d(6, 96, kernel_size=(4, 4), stride=(4, 4))\n        self.landsat_model.head = nn.Identity()\n        \n        self.bioclim_norm = nn.LayerNorm([4,19,12])\n        self.bioclim_model = models.swin_t(weights=None)  \n        # Modify the first convolutional layer to accept 4 channels instead of 3\n        self.bioclim_model.features[0][0] = nn.Conv2d(4, 96, kernel_size=(4, 4), stride=(4, 4))\n        self.bioclim_model.head = nn.Identity()\n        \n        self.sentinel_model = models.swin_t(weights=\"IMAGENET1K_V1\")\n        # Modify the first layer to accept 4 channels instead of 3\n        self.sentinel_model.features[0][0] = nn.Conv2d(4, 96, kernel_size=(4, 4), stride=(4, 4))\n        self.sentinel_model.head = nn.Identity()\n        \n        \n        self.proj1 = nn.Sequential(\n            nn.Linear(768, 1000),\n            nn.BatchNorm1d(1000),\n            nn.GELU(),\n            nn.Dropout(0.2)\n        )\n        self.proj2 = nn.Sequential(\n            nn.Linear(768, 1000),\n            nn.BatchNorm1d(1000),\n            nn.GELU(),\n            nn.Dropout(0.2)\n        )\n        self.proj3 = nn.Sequential(\n            nn.Linear(768, 1000),\n            nn.BatchNorm1d(1000),\n            nn.GELU(),\n            nn.Dropout(0.2)\n        )\n        self.label = nn.Sequential(\n            nn.Linear(3000, 4096),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            nn.Linear(4096, num_classes),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            nn.Linear(num_classes, num_classes),\n        )\n        \n    def forward(self, x, y, z):\n        \n        x = self.landsat_norm(x)\n        x = self.landsat_model(x)\n        x = self.proj1(x)\n        \n        y = self.bioclim_norm(y)\n        y = self.bioclim_model(y)\n        y = self.proj2(y)\n        \n        z = self.proj3(self.sentinel_model(z))\n        \n        \n        xyz = torch.cat((x, y, z), dim=1)\n        out = self.label(xyz)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:42.674493Z","iopub.execute_input":"2025-06-01T12:26:42.674789Z","iopub.status.idle":"2025-06-01T12:26:42.686374Z","shell.execute_reply.started":"2025-06-01T12:26:42.674764Z","shell.execute_reply":"2025-06-01T12:26:42.685808Z"}},"outputs":[],"execution_count":null},{"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(69)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:42.687009Z","iopub.execute_input":"2025-06-01T12:26:42.687216Z","iopub.status.idle":"2025-06-01T12:26:42.708195Z","shell.execute_reply.started":"2025-06-01T12:26:42.687173Z","shell.execute_reply":"2025-06-01T12:26:42.707638Z"}},"outputs":[],"execution_count":null},{"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\n# num_classes = 11255 # Number of all unique classes within the PO and PA data.\nmodel = MultimodalEnsemble(num_classes).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:42.709563Z","iopub.execute_input":"2025-06-01T12:26:42.710009Z","iopub.status.idle":"2025-06-01T12:26:44.724667Z","shell.execute_reply.started":"2025-06-01T12:26:42.709994Z","shell.execute_reply":"2025-06-01T12:26:44.724085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Hyperparameters\nlearning_rate = 8e-5\nnum_epochs = 12\npositive_weigh_factor = 1.0\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate)\nscheduler = CosineAnnealingLR(optimizer, T_max=25, verbose=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:44.725431Z","iopub.execute_input":"2025-06-01T12:26:44.725649Z","iopub.status.idle":"2025-06-01T12:26:44.733233Z","shell.execute_reply.started":"2025-06-01T12:26:44.725632Z","shell.execute_reply":"2025-06-01T12:26:44.732548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Training for {num_epochs} epochs started.\")\n\nfor epoch in range(num_epochs):\n    model.train()\n    \n    for batch_idx, (data1, data2, data3, targets, _) in enumerate(train_loader):\n\n        data1 = data1.to(device)\n        data2 = data2.to(device)\n        data3 = data3.to(device)\n        targets = targets.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(data1, data2, data3)\n\n        pos_weight = targets*positive_weigh_factor  \n        criterion = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n        loss = criterion(outputs, targets)\n\n        loss.backward()\n        optimizer.step()\n\n        if batch_idx % 128 == 0:\n            print(f\"Epoch {epoch+1}/{num_epochs}, Batch {batch_idx}/{len(train_loader)}, Loss: {loss.item()}\")\n\n\n    # Save the trained model\n    # print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {loss.item()}\")\n    model.eval()\n    torch.save(model.state_dict(), f\"{epoch}-multimodal-model.pth\")\n    \n    scheduler.step()\n    print(\"Scheduler:\",scheduler.state_dict())\n\n# Save the trained model\nmodel.eval()\ntorch.save(model.state_dict(), \"multimodal-model.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model.load_state_dict(torch.load(\"/kaggle/input/train-old-architecture/11-multimodal-model.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:44.734145Z","iopub.execute_input":"2025-06-01T12:26:44.734408Z","iopub.status.idle":"2025-06-01T12:26:45.17302Z","shell.execute_reply.started":"2025-06-01T12:26:44.734392Z","shell.execute_reply":"2025-06-01T12:26:45.172451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model.eval()\n# print(\"Okay!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:46.294437Z","iopub.execute_input":"2025-06-01T12:26:46.294716Z","iopub.status.idle":"2025-06-01T12:26:46.300829Z","shell.execute_reply.started":"2025-06-01T12:26:46.294694Z","shell.execute_reply":"2025-06-01T12:26:46.300005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nprint(\"Done\")\n\nwith torch.no_grad():\n    surveys = []\n    predictions_list = []\n    top_k_indices = None\n    for batch_idx, (data1, data2, data3, surveyID) in enumerate(test_loader):\n\n        data1 = data1.to(device)\n        data2 = data2.to(device)\n        data3 = data3.to(device)\n\n        outputs = model(data1, data2, data3)\n        predictions = torch.sigmoid(outputs).cpu().numpy()\n        \n        batch_top_predictions = []\n        for el in predictions:\n            answ = np.array(list(np.nonzero(el > 0.18)[0]))\n            \n            if len(answ) < 14:\n                answ = np.array(np.argsort(-el)[:14])\n            batch_top_predictions.append(answ)\n        \n        batch_top_unique = [unique[el] for el in batch_top_predictions]\n        \n        predictions_list.extend(batch_top_unique)\n        \n        surveys.extend(surveyID.cpu().numpy())\n\n    top_k_indices = predictions_list","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:26:46.601305Z","iopub.execute_input":"2025-06-01T12:26:46.601547Z","iopub.status.idle":"2025-06-01T12:42:11.734517Z","shell.execute_reply.started":"2025-06-01T12:26:46.60153Z","shell.execute_reply":"2025-06-01T12:42:11.733437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_concatenated = [' '.join(map(lambda x: str(int(x)), row)) for row in top_k_indices]\n\n\npd.DataFrame(\n    {'surveyId': surveys,\n     'predictions': data_concatenated,\n    }).to_csv(\"submission.csv\", index = False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-01T12:42:11.736127Z","iopub.execute_input":"2025-06-01T12:42:11.737304Z","iopub.status.idle":"2025-06-01T12:42:11.996094Z","shell.execute_reply.started":"2025-06-01T12:42:11.737278Z","shell.execute_reply":"2025-06-01T12:42:11.995413Z"}},"outputs":[],"execution_count":null}]}