{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install --no-index --find-links /kaggle/input/onnxruntime/ onnxruntime","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:23.622140Z","iopub.execute_input":"2023-05-25T10:50:23.622599Z","iopub.status.idle":"2023-05-25T10:50:37.471837Z","shell.execute_reply.started":"2023-05-25T10:50:23.622556Z","shell.execute_reply":"2023-05-25T10:50:37.470325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport gc\n\nimport os\nfrom pathlib import Path\n\nimport yaml\nfrom tqdm.notebook import tqdm\nfrom types import SimpleNamespace\n\nimport timm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchaudio\nimport librosa\nimport soundfile as sf\n\nimport onnxruntime\nimport concurrent\n\nimport time\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:37.477068Z","iopub.execute_input":"2023-05-25T10:50:37.477474Z","iopub.status.idle":"2023-05-25T10:50:42.035599Z","shell.execute_reply.started":"2023-05-25T10:50:37.477425Z","shell.execute_reply":"2023-05-25T10:50:42.034098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Utils","metadata":{}},{"cell_type":"code","source":"def load_config(filepath):\n    with open(filepath, 'rb') as file:\n        data = yaml.safe_load(file)\n    data = dictionary_to_namespace(data)\n    return data\n\n\ndef dictionary_to_namespace(data):\n    if type(data) is list:\n        return list(map(dictionary_to_namespace, data))\n    elif type(data) is dict:\n        sns = SimpleNamespace()\n        for key, value in data.items():\n            setattr(sns, key, dictionary_to_namespace(value))\n        return sns\n    else:\n        return data\n    \n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \n    \nseed_everything()","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:42.037698Z","iopub.execute_input":"2023-05-25T10:50:42.038189Z","iopub.status.idle":"2023-05-25T10:50:42.052575Z","shell.execute_reply.started":"2023-05-25T10:50:42.038131Z","shell.execute_reply":"2023-05-25T10:50:42.051605Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"class AudioTransform:\n    def __init__(self, always_apply=False, p=0.5):\n        self.always_apply = always_apply\n        self.p = p\n\n    def __call__(self, wav):\n        if self.always_apply:\n            return self.apply(wav)\n        else:\n            if np.random.rand() < self.p:\n                return self.apply(wav)\n            else:\n                return wav\n\n    def apply(self, wav):\n        raise NotImplementedError\n        \n\nclass Normalize(AudioTransform):\n    def __init__(self, always_apply=False, p=1):\n        super().__init__(always_apply, p)\n\n    def apply(self, wav):\n        max_vol = np.abs(wav).max()\n        y_vol = wav * 1 / max_vol\n        return np.asfortranarray(y_vol)\n    \n\nclass NormalizeMelSpec(nn.Module):\n    def __init__(self, eps=1e-6):\n        super().__init__()\n        self.eps = eps\n\n    def forward(self, X):\n        mean = X.mean((1, 2), keepdim=True)\n        std = X.std((1, 2), keepdim=True)\n        Xstd = (X - mean) / (std + self.eps)\n        norm_min, norm_max = Xstd.min(-1)[0].min(-1)[0], Xstd.max(-1)[0].max(-1)[0]\n        fix_ind = (norm_max - norm_min) > self.eps * torch.ones_like(\n            (norm_max - norm_min)\n        )\n        V = torch.zeros_like(Xstd)\n        if fix_ind.sum():\n            V_fix = Xstd[fix_ind]\n            norm_max_fix = norm_max[fix_ind, None, None]\n            norm_min_fix = norm_min[fix_ind, None, None]\n            V_fix = torch.max(\n                torch.min(V_fix, norm_max_fix),\n                norm_min_fix,\n            )\n            V_fix = (V_fix - norm_min_fix) / (norm_max_fix - norm_min_fix)\n            V[fix_ind] = V_fix\n        return V\n    \n    \nclass TestDataset(Dataset):\n    def __init__(self, audios, dataframe, config):\n        self.audios = audios\n        \n        dfs = []\n        for path, id in zip(dataframe.path.values, dataframe.filename.values):\n            for i in range(5, 605, 5):\n                start_seconds = i - 5\n                end_seconds = i\n                end_index = int(32000 * end_seconds)\n                start_index = int(32000 * start_seconds)\n                row_id = f'{id}_{end_seconds}'\n                dfs.append([path, start_index, end_index, row_id])\n        self.dataframe = pd.DataFrame(dfs, columns=['path', 'start', 'stop', 'row_id'])\n        \n        self.config = config\n        self.sr = self.config.dataset.sample_rate\n\n        self.mel_spec = torchaudio.transforms.MelSpectrogram(\n            32000,\n            n_mels=self.config.dataset.n_mels,\n            n_fft=self.config.dataset.nfft,\n            hop_length=self.config.dataset.hop_length,\n            f_max=self.config.dataset.fmax,\n            f_min=self.config.dataset.fmin,\n        )\n        self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB(top_db=config.dataset.top_db)\n        self.spec_norm = NormalizeMelSpec()\n        \n        self.wav_norm = Normalize(p=1)\n        \n        self.spec_transforms = nn.Sequential(\n                self.mel_spec,\n                self.amplitude_to_db,\n                self.spec_norm,\n            )\n        \n        self.train_period = config.dataset.train_duration\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, item):\n        row = self.dataframe.iloc[item]\n        fp = row.path\n        start = row.start\n        stop = row.stop\n        \n        audio = self.audios[str(fp)][start:stop]\n        audio = torch.tensor(self.wav_norm(audio))\n        image = self.spec_transforms(audio.view(1, audio.shape[0]))\n        return image","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:42.055761Z","iopub.execute_input":"2023-05-25T10:50:42.056892Z","iopub.status.idle":"2023-05-25T10:50:42.081703Z","shell.execute_reply.started":"2023-05-25T10:50:42.056844Z","shell.execute_reply":"2023-05-25T10:50:42.080494Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"def gem_freq(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), 1)).pow(1.0 / p)\n\n\nclass GeMFreq(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super().__init__()\n        self.p = torch.nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return gem_freq(x, p=self.p, eps=self.eps)\n\n\nclass AttHead(nn.Module):\n    def __init__(\n        self, in_chans, p=0.5, num_class=264, train_period=15.0, infer_period=5.0\n    ):\n        super().__init__()\n        self.train_period = train_period\n        self.infer_period = infer_period\n        self.pooling = GeMFreq()\n\n        self.dense_layers = nn.Sequential(\n            nn.Dropout(p / 2),\n            nn.Linear(in_chans, 512),\n            nn.ReLU(),\n            nn.Dropout(p),\n        )\n        self.attention = nn.Conv1d(\n            in_channels=512,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.fix_scale = nn.Conv1d(\n            in_channels=512,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n\n    def forward(self, feat):\n        feat = self.pooling(feat).squeeze(-2).permute(0, 2, 1)  # (bs, time, ch)\n\n        feat = self.dense_layers(feat).permute(0, 2, 1)  # (bs, 512, time)\n        # print(feat.shape)\n        \n        time_att = torch.tanh(self.attention(feat))\n        \n        assert self.train_period >= self.infer_period\n        \n        if self.training or self.train_period == self.infer_period: # or True\n            # print('train')\n            clipwise_pred = torch.sum(\n                torch.sigmoid(self.fix_scale(feat)) * torch.softmax(time_att, dim=-1),\n                dim=-1,\n            )  # sum((bs, 24, time), -1) -> (bs, 24)\n            logits = torch.sum(\n                self.fix_scale(feat) * torch.softmax(time_att, dim=-1),\n                dim=-1,\n            )\n        else:\n            # print('eval')\n            clipwise_pred_long = torch.sum(\n                torch.sigmoid(self.fix_scale(feat)) * torch.softmax(time_att, dim=-1),\n                dim=-1,\n            )  # sum((bs, 24, time), -1) -> (bs, 24)\n            \n            feat_time = feat.size(-1)\n            start = feat_time / 2 - feat_time * (self.infer_period / self.train_period) / 2\n            end = start + feat_time * (self.infer_period / self.train_period)\n            \n            start = int(start)\n            end = int(end)\n            \n            feat = feat[:, :, start:end]\n            att = torch.softmax(time_att[:, :, start:end], dim=-1)\n            \n            # print(feat.shape)\n            \n            clipwise_pred = torch.sum(\n                torch.sigmoid(self.fix_scale(feat)) * att,\n                dim=-1,\n            )\n            logits = torch.sum(\n                self.fix_scale(feat) * att,\n                dim=-1,\n            )\n            time_att = time_att[:, :, start:end]\n        return (\n            logits,\n            clipwise_pred,\n            self.fix_scale(feat).permute(0, 2, 1),\n            time_att.permute(0, 2, 1),\n        )\n\n\nclass TattakaModel(nn.Module):\n    def __init__(\n        self,\n        config,\n        pretrained,\n    ):\n        super().__init__()\n\n        # self.model = get_timm_backbone(config, pretrained)\n        self.model = timm.create_model(\n            config.model.backbone_type, features_only=True, pretrained=False, in_chans=1\n        )\n        encoder_channels = self.model.feature_info.channels()\n        dense_input = encoder_channels[-1]\n        self.head = AttHead(\n            dense_input,\n            p=config.model.dropout,\n            num_class=len(config.dataset.labels),\n            train_period=config.dataset.train_duration,\n            infer_period=config.dataset.valid_duration,\n        )\n        self.criterion = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, images):\n        spec = images\n        \n        feats = self.model(spec)\n        logits, output_clip, output_frame, output_attention = self.head(feats[-1])\n        return output_clip\n","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-05-25T10:50:42.083298Z","iopub.execute_input":"2023-05-25T10:50:42.084357Z","iopub.status.idle":"2023-05-25T10:50:42.110469Z","shell.execute_reply.started":"2023-05-25T10:50:42.084315Z","shell.execute_reply":"2023-05-25T10:50:42.109363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"exp_list = [\n    [\n        'exp290_enc', [[['exp290'], [-1]],]\n    ],\n    [\n        'exp302_enc', [[['exp302'], [-1]],]\n    ],\n    [\n        'exp295_enc', [\n            [['exp294'], [-1]],\n            \n            [['exp295'], [-1]],\n            [['exp295'], [0]],\n            [['exp295'], [1]],\n            [['exp295'], [2]],\n            [['exp295'], [3]],\n        ],\n    ],\n    [\n        'exp122_enc', [[['exp122',], [-1]],]\n    ],\n    [\n        'exp122_0_enc', [[['exp122',], [0]],]\n    ],\n    [\n        'exp122_1_enc', [[['exp122',], [1]],]\n    ],\n    [\n        'exp168_enc', [[['exp168',], [-1]],]\n    ],\n    [\n        'exp170_enc', [[['exp170'], [-1]]]\n    ],\n    [\n        'exp280_enc', [[['exp280'], [-1]]]\n    ],\n    [\n        'exp325_-1_enc', [[['exp325'], [-1]]]\n    ],\n]","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:42.111963Z","iopub.execute_input":"2023-05-25T10:50:42.113050Z","iopub.status.idle":"2023-05-25T10:50:42.129565Z","shell.execute_reply.started":"2023-05-25T10:50:42.113011Z","shell.execute_reply":"2023-05-25T10:50:42.128105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = ['abethr1', 'abhori1', 'abythr1', 'afbfly1', 'afdfly1', 'afecuc1',\n       'affeag1', 'afgfly1', 'afghor1', 'afmdov1', 'afpfly1', 'afpkin1',\n       'afpwag1', 'afrgos1', 'afrgrp1', 'afrjac1', 'afrthr1', 'amesun2',\n       'augbuz1', 'bagwea1', 'barswa', 'bawhor2', 'bawman1', 'bcbeat1',\n       'beasun2', 'bkctch1', 'bkfruw1', 'blacra1', 'blacuc1', 'blakit1',\n       'blaplo1', 'blbpuf2', 'blcapa2', 'blfbus1', 'blhgon1', 'blhher1',\n       'blksaw1', 'blnmou1', 'blnwea1', 'bltapa1', 'bltbar1', 'bltori1',\n       'blwlap1', 'brcale1', 'brcsta1', 'brctch1', 'brcwea1', 'brican1',\n       'brobab1', 'broman1', 'brosun1', 'brrwhe3', 'brtcha1', 'brubru1',\n       'brwwar1', 'bswdov1', 'btweye2', 'bubwar2', 'butapa1', 'cabgre1',\n       'carcha1', 'carwoo1', 'categr', 'ccbeat1', 'chespa1', 'chewea1',\n       'chibat1', 'chtapa3', 'chucis1', 'cibwar1', 'cohmar1', 'colsun2',\n       'combul2', 'combuz1', 'comsan', 'crefra2', 'crheag1', 'crohor1',\n       'darbar1', 'darter3', 'didcuc1', 'dotbar1', 'dutdov1', 'easmog1',\n       'eaywag1', 'edcsun3', 'egygoo', 'equaka1', 'eswdov1', 'eubeat1',\n       'fatrav1', 'fatwid1', 'fislov1', 'fotdro5', 'gabgos2', 'gargan',\n       'gbesta1', 'gnbcam2', 'gnhsun1', 'gobbun1', 'gobsta5', 'gobwea1',\n       'golher1', 'grbcam1', 'grccra1', 'grecor', 'greegr', 'grewoo2',\n       'grwpyt1', 'gryapa1', 'grywrw1', 'gybfis1', 'gycwar3', 'gyhbus1',\n       'gyhkin1', 'gyhneg1', 'gyhspa1', 'gytbar1', 'hadibi1', 'hamerk1',\n       'hartur1', 'helgui', 'hipbab1', 'hoopoe', 'huncis1', 'hunsun2',\n       'joygre1', 'kerspa2', 'klacuc1', 'kvbsun1', 'laudov1', 'lawgol',\n       'lesmaw1', 'lessts1', 'libeat1', 'litegr', 'litswi1', 'litwea1',\n       'loceag1', 'lotcor1', 'lotlap1', 'luebus1', 'mabeat1', 'macshr1',\n       'malkin1', 'marsto1', 'marsun2', 'mcptit1', 'meypar1', 'moccha1',\n       'mouwag1', 'ndcsun2', 'nobfly1', 'norbro1', 'norcro1', 'norfis1',\n       'norpuf1', 'nubwoo1', 'pabspa1', 'palfly2', 'palpri1', 'piecro1',\n       'piekin1', 'pitwhy', 'purgre2', 'pygbat1', 'quailf1', 'ratcis1',\n       'raybar1', 'rbsrob1', 'rebfir2', 'rebhor1', 'reboxp1', 'reccor',\n       'reccuc1', 'reedov1', 'refbar2', 'refcro1', 'reftin1', 'refwar2',\n       'rehblu1', 'rehwea1', 'reisee2', 'rerswa1', 'rewsta1', 'rindov',\n       'rocmar2', 'rostur1', 'ruegls1', 'rufcha2', 'sacibi2', 'sccsun2',\n       'scrcha1', 'scthon1', 'shesta1', 'sichor1', 'sincis1', 'slbgre1',\n       'slcbou1', 'sltnig1', 'sobfly1', 'somgre1', 'somtit4', 'soucit1',\n       'soufis1', 'spemou2', 'spepig1', 'spewea1', 'spfbar1', 'spfwea1',\n       'spmthr1', 'spwlap1', 'squher1', 'strher', 'strsee1', 'stusta1',\n       'subbus1', 'supsta1', 'tacsun1', 'tafpri1', 'tamdov1', 'thrnig1',\n       'trobou1', 'varsun2', 'vibsta2', 'vilwea1', 'vimwea1', 'walsta1',\n       'wbgbir1', 'wbrcha2', 'wbswea1', 'wfbeat1', 'whbcan1', 'whbcou1',\n       'whbcro2', 'whbtit5', 'whbwea1', 'whbwhe3', 'whcpri2', 'whctur2',\n       'wheslf1', 'whhsaw1', 'whihel1', 'whrshr1', 'witswa1', 'wlwwar',\n       'wookin1', 'woosan', 'wtbeat1', 'yebapa1', 'yebbar1', 'yebduc1',\n       'yebere1', 'yebgre1', 'yebsto1', 'yeccan1', 'yefcan', 'yelbis1',\n       'yenspu1', 'yertin1', 'yesbar1', 'yespet1', 'yetgre1', 'yewgre1']\n\nlabels = sorted(labels)","metadata":{"papermill":{"duration":0.035353,"end_time":"2022-04-22T06:01:17.283888","exception":false,"start_time":"2022-04-22T06:01:17.248535","status":"completed"},"tags":[],"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-05-25T10:50:42.131020Z","iopub.execute_input":"2023-05-25T10:50:42.131398Z","iopub.status.idle":"2023-05-25T10:50:42.148233Z","shell.execute_reply.started":"2023-05-25T10:50:42.131361Z","shell.execute_reply":"2023-05-25T10:50:42.147336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = Path('/kaggle/input/birdclef-2023/test_soundscapes')\n\ntest_df = pd.DataFrame(\n     [(path.stem, *path.stem.split(\"_\"), path) for path in path.glob(\"*.ogg\")],\n    columns = [\"filename\", \"name\" ,\"id\", \"path\"]\n)\n\naudios = {}\npaths = test_df.path.unique()\nidxs = [i for i in range(len(paths))]\n\ndef read_audios(idx):\n    fp = paths[idx]\n    wav, _ = sf.read(fp, dtype='float32')\n    audios[str(fp)] = wav\n    return True\n\nwith concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n    _ = executor.map(read_audios, idxs)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:42.149507Z","iopub.execute_input":"2023-05-25T10:50:42.150576Z","iopub.status.idle":"2023-05-25T10:50:42.849446Z","shell.execute_reply.started":"2023-05-25T10:50:42.150524Z","shell.execute_reply":"2023-05-25T10:50:42.848034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = load_config('/kaggle/input/birdclef-models/config.yaml')\n    \nconfig.dataset.train_duration = 5\nconfig.dataset.valid_duration = 5\nconfig.dataset.labels = labels\n\ntest_ds = TestDataset(audios, test_df, config)\ndataloader = torch.utils.data.DataLoader(\n    test_ds,\n    batch_size=8,\n    num_workers=4,\n    shuffle=False,\n    pin_memory=True,\n    drop_last=False,\n)\n\nbatches = [batch for batch in dataloader]\nidxs = [i for i in range(len(batches))]","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:42.850917Z","iopub.execute_input":"2023-05-25T10:50:42.851378Z","iopub.status.idle":"2023-05-25T10:50:43.661146Z","shell.execute_reply.started":"2023-05-25T10:50:42.851328Z","shell.execute_reply":"2023-05-25T10:50:43.659539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_encoder_output(idx):\n    images = batches[idx]\n    with torch.no_grad():\n        logits = encoder_session.run([], {'input': images.numpy()})[0]\n        encoder_out[idx] = logits\n    return True\n\n\ndef get_backbone_output(idx):\n    images = encoder_out[idx]\n    with torch.no_grad():\n        embeddings = backbone_session.run([], {'input': images})[0]\n        backbone_out[idx] = embeddings\n    return True\n\n\ndef get_head_output(idx):\n    embeddings = backbone_out[idx]\n    with torch.no_grad():\n        logits = head_session.run([], {'input': embeddings})[0]\n        predictions[idx] = logits\n    return True","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:43.666256Z","iopub.execute_input":"2023-05-25T10:50:43.666937Z","iopub.status.idle":"2023-05-25T10:50:43.675735Z","shell.execute_reply.started":"2023-05-25T10:50:43.666885Z","shell.execute_reply":"2023-05-25T10:50:43.674375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_predictions = None\ncntr = 0\n\nfor enc_name, same_encoder in tqdm(exp_list):\n    encoder_state = f'/kaggle/input/birdclef-models/{enc_name}.onnx'\n    encoder_session = onnxruntime.InferenceSession(encoder_state, providers=['CPUExecutionProvider'])\n    print('Encoder from: ', encoder_state)\n    \n    encoder_out = {}\n    with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n        _ = executor.map(get_encoder_output, idxs)\n    \n    for same_backbone in same_encoder:\n        backbone_state = f'/kaggle/input/birdclef-models/{same_backbone[0][0]}_{same_backbone[1][0]}.onnx'\n        backbone_session = onnxruntime.InferenceSession(backbone_state, providers=['CPUExecutionProvider'])\n        print('\\tBackbone from: ', backbone_state)\n    \n        backbone_out = {}\n        with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n            _ = executor.map(get_backbone_output, idxs)\n            \n        exp_l, fold = same_backbone\n        for exp in exp_l:\n            head_state = f'/kaggle/input/birdclef-models/{exp}_{fold[0]}_head.onnx'\n            head_session = onnxruntime.InferenceSession(head_state, providers=['CPUExecutionProvider'])\n            print('\\t\\tHead from: ', head_state)\n            \n            predictions = {}\n            with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n                _ = executor.map(get_head_output, idxs)\n                \n            preds = []\n            for i in range(len(idxs)):\n                preds.append(predictions[i])\n            preds = np.concatenate(preds)\n            \n            if all_predictions is None:\n                all_predictions = preds\n            else:\n                all_predictions += preds\n            \n            all_predictions = np.round(all_predictions, 5)\n            cntr += 1 \n            \n            del head_session, predictions, preds\n            gc.collect()\n        \n        del backbone_session, backbone_out\n        gc.collect()\n            \n    print('\\n\\n')\n    \n    del encoder_session, encoder_out\n    gc.collect()\n\n    \npredictions = all_predictions / cntr","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:50:43.677359Z","iopub.execute_input":"2023-05-25T10:50:43.677762Z","iopub.status.idle":"2023-05-25T10:51:36.223985Z","shell.execute_reply.started":"2023-05-25T10:50:43.677723Z","shell.execute_reply":"2023-05-25T10:51:36.222701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idxs = []\nfor fn in test_df.filename.values:\n    for i in range(5, 605, 5):\n        idxs.append(f'{fn}_{i}')\n        \nsub_df = pd.DataFrame(columns=['row_id']+labels)\nsub_df['row_id'] = idxs\nsub_df[labels] = predictions","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:51:36.225261Z","iopub.execute_input":"2023-05-25T10:51:36.225603Z","iopub.status.idle":"2023-05-25T10:51:36.312063Z","shell.execute_reply.started":"2023-05-25T10:51:36.225570Z","shell.execute_reply":"2023-05-25T10:51:36.310827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"debug = sub_df.shape[0]<=120\n\nif debug:\n    import matplotlib.pyplot as plt\n\n    fig, ax = plt.subplots(figsize=(20, 30))\n    heatmap = ax.pcolor(sub_df.iloc[:,1:].values.T, edgecolors='k', linewidths=0.1, vmin=0, vmax=1, cmap='Blues')\n    ax.set_xticks(np.arange(0, sub_df.shape[0]+0.5, 12))\n    ax.set_yticks(np.arange(sub_df.shape[1]-1))\n    ax.set_xticklabels(np.arange(0,605,60))\n    ax.set_yticklabels(labels)\n    plt.xlabel('sec')\n    plt.ylabel('species')\n    fig.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:51:36.313637Z","iopub.execute_input":"2023-05-25T10:51:36.314092Z","iopub.status.idle":"2023-05-25T10:51:40.461396Z","shell.execute_reply.started":"2023-05-25T10:51:36.314044Z","shell.execute_reply":"2023-05-25T10:51:40.460227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:51:40.462830Z","iopub.execute_input":"2023-05-25T10:51:40.463773Z","iopub.status.idle":"2023-05-25T10:51:40.518221Z","shell.execute_reply.started":"2023-05-25T10:51:40.463718Z","shell.execute_reply":"2023-05-25T10:51:40.517186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.status.busy":"2023-05-25T10:51:40.519746Z","iopub.execute_input":"2023-05-25T10:51:40.520132Z","iopub.status.idle":"2023-05-25T10:51:40.567539Z","shell.execute_reply.started":"2023-05-25T10:51:40.520088Z","shell.execute_reply":"2023-05-25T10:51:40.566342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}