import os

# Improve CPU inference speed and prevent dataloader timeouts (for local runs)
os.environ["MKL_NUM_THREADS"] = "4"
os.environ["NUMEXPR_NUM_THREADS"] = "4"
os.environ["OMP_NUM_THREADS"] = "4"
os.environ["OMP_SCHEDULE"] = "STATIC"

import time
import numpy as np
import pandas as pd
from pathlib import Path
import soundfile as sf
import librosa
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from PIL import Image
from joblib import Parallel, delayed
import multiprocessing
import json

inferenceStartTime = time.time()
checkTimePerBatchTime = 6000 # 100 minutes
checkTimePerBatch = False
maxInferenceTime = 7110 # 118.5 min 7140 # 119 min 7080 # 118 min 7020 # 117 min
maxInferenceTimeReached = False

useSecondPass = False
nCheckpointsFirstPass = 5
percentagePartsToKeepForSecondPass = 0.7

runLocal = False # True False
outputPrecision = None # None or number of decimal places for results in submission file
useCpu = True # True False
nWorkers = None


rootDir = '../input/'
checkpointRootDir = rootDir + 'bc2023models/'

ensembleMethod = 'mean' # mean max, meanexp
batchSizeForTesting = 8


if runLocal and useCpu:
    os.environ["CUDA_VISIBLE_DEVICES"] = "-1"

submissionMode = len(list(Path(rootDir + 'birdclef-2023/test_soundscapes/').glob('*.ogg'))) > 1

dataDir = rootDir + 'birdclef-2023/test_soundscapes/'

if runLocal:
    nWorkers = 4
    dataDir = rootDir + 'testdata/02_4x_testsoundscape/'
    
    
print('submissionMode', submissionMode)
print('dataDir', dataDir)



checkpoints = []

# Checkpoints orderd by public LB score

# 0, v2, f5 (168)
checkpoints.append({'path': checkpointRootDir + '57_1_53_9_V2_Fold5_NoIc_AddReverb20_200e/Checkpoint188.pt'})
# 1, v2, f3 (144)
checkpoints.append({'path': checkpointRootDir + '53_7_53_3_V2_Fold3_200e/Checkpoint186.pt'})
# 2, resnet50, f5 (203)
checkpoints.append({'path': checkpointRootDir + '60_6_53_9_ResNet50/Checkpoint199.pt'})
# 3, v2, f5 (182)
checkpoints.append({'path': checkpointRootDir + '59_5_57_1_V2_Fold5_Ic80_AddReverb20_200e/Checkpoint178.pt'})

# 4, v2, f5 (217)
checkpoints.append({'path': checkpointRootDir + '62_3_Cont_62_1_LessTestsoundscape/Checkpoint197.pt'})
# 5, v2, f5 (213)
checkpoints.append({'path': checkpointRootDir + '62_1_57_1_AddXcDl02Data/Checkpoint161.pt'})


# 6, v2, f3 (129)
checkpoints.append({'path': checkpointRootDir + '30_8_30_7_400epochs/Checkpoint475.pt'})
# 7, b0, f5 (148)
checkpoints.append({'path': checkpointRootDir + '53_9_B0_Fold5_NoIc_AddReverb20_200e/Checkpoint197.pt'})
# 8, v2, f3 (116)
checkpoints.append({'path': checkpointRootDir + '30_7_30_5_Fold3_IC80_EffNetV2/Checkpoint159.pt'})

# 9, resnet152, f3 (208)
checkpoints.append({'path': checkpointRootDir + '61_5_53_7_ResNet152/Checkpoint190.pt'})

# 10, v2, f0 (180)
checkpoints.append({'path': checkpointRootDir + '59_0_51_9_V2_Fold0_Ic80_AddReverb20_160e/Checkpoint159.pt'})
# 11, v2, f3 (223)
checkpoints.append({'path': checkpointRootDir + '63_3_53_7_MoreNoiseAndIRs/Checkpoint210.pt'})
# 12, b0, f0 (132)
checkpoints.append({'path': checkpointRootDir + '51_8_51_7_AddReverb20/Checkpoint154.pt'})
# 13, v2, f0 (173)
checkpoints.append({'path': checkpointRootDir + '57_3_V2_Fold0_NoIc_AddReverbBeforeMix_240e/Checkpoint231.pt'})
# 14, b2, f5 (186)
checkpoints.append({'path': checkpointRootDir + '59_7_B2_Fold5_Ic75_AddReverb10plus10_200e/Checkpoint199.pt'})



# # Get checkpoint selection in specific order

# Test
#cpIxs = [7,12]

# # 224 (best ens so far 222 but with 0.65 parts kept for second pass)
# cpIxs = [0,1,6,7,12, 3,8,13,10]
# useSecondPass = True
# nCheckpointsFirstPass = 5
# percentagePartsToKeepForSecondPass = 0.65 #0.7

# 225 (best ens so far 222 but with 4 instead of 8)
cpIxs = [0,1,6,7,12, 3,4,13,10]
useSecondPass = True
nCheckpointsFirstPass = 5
percentagePartsToKeepForSecondPass = 0.7


# # 226 (best ens so far 222 but with 4 instead of 8, 14 instead of 13)
# cpIxs = [0,1,6,7,12, 3,4,14,10]
# useSecondPass = True
# nCheckpointsFirstPass = 5
# percentagePartsToKeepForSecondPass = 0.65

# # 227 (best ens so far 222 but with 14 instead of 13, 2 instead of 10)
# cpIxs = [0,1,6,7,12, 3,8,14,2]
# useSecondPass = True
# nCheckpointsFirstPass = 5
# percentagePartsToKeepForSecondPass = 0.6




checkpoints = [checkpoints[i] for i in cpIxs]


nCheckpoints = len(checkpoints)
print('nCheckpoints', nCheckpoints)






timeStart = time.time()
if not nWorkers:
    nWorkers = multiprocessing.cpu_count()


def get_test_df():

    # Get dataframe from test files (filename, path, row_id, start_time, end_time)
    
    filenames = []
    paths = []
    row_ids = []
    start_times = []
    end_times = []
        
    files = os.listdir(dataDir)
    for file in files:
        if file.endswith('.ogg'):
            path = dataDir + file
            filename = os.path.splitext(file)[0]
            # ToDo: maybe check/verify with actual file length and/or row_ids in test.csv
            for end_time in range(5, 605, 5):
                row_id = filename + '_'  + str(end_time) 
                start_time = end_time - 5
                filenames.append(filename)
                paths.append(path)
                row_ids.append(row_id)
                start_times.append(start_time)
                end_times.append(end_time)

    test_df = pd.DataFrame({
            'filename': filenames,
            'path': paths,
            'row_id': row_ids,
            'start_time': start_times,
            'end_time': end_times
        })

    return test_df


def getDefaultSpecImagesPerFile(
        path,
        sampleRate=32000,
        fftSizeInSamples=2048,
        fftHopSizeInSamples=512,
        segmentDuration=5.0,
        nMelBands=128,
        melStartFreq=40.0,
        melEndFreq=15000.0,
        nLowFreqsInPixelToCutMax=2,
        nHighFreqsInPixelToCutMax=4,
        interpolationMethod=Image.Resampling.LANCZOS,
        imageSize=(312, 128),
        ):

    # Read audio file (assume mono, 32000 Hz)
    sampleVecFile, sampleRateSrc = sf.read(path)
    
    specImages = []

    #durationFile = len(sampleVecFile)/float(sampleRate)
    durationFile = 600.0 # Assume 10min files, get duration maybe from file?
    durationFileInSamples = int(durationFile * sampleRate)
    segmentDurationInSamples = int(segmentDuration * sampleRate)
    
    for startSample in range(0, durationFileInSamples, segmentDurationInSamples):

        endSample = startSample + segmentDurationInSamples
        
        sampleVec = sampleVecFile[startSample:endSample]

        # Get mel spectrogram
        melSpec = librosa.feature.melspectrogram(y=sampleVec, sr=sampleRate, n_fft=fftSizeInSamples, hop_length=fftHopSizeInSamples, n_mels=nMelBands, fmin=melStartFreq, fmax=melEndFreq, power=2.0)

        # Convert power spec to dB scale (compute dB relative to peak power)
        melSpec = librosa.power_to_db(melSpec, ref=np.max, top_db=100)

        nLowFreqsInPixelToCut = int(nLowFreqsInPixelToCutMax/2.0)
        nHighFreqsInPixelToCut = int(nHighFreqsInPixelToCutMax/2.0)

        if nHighFreqsInPixelToCut:
            melSpec = melSpec[nLowFreqsInPixelToCut:-nHighFreqsInPixelToCut]
        else:
            melSpec = melSpec[nLowFreqsInPixelToCut:]

        # Flip spectrum vertically (only for better visialization, low freq. at bottom)
        melSpec = melSpec[::-1, ...]

        # Normalize values between 0 and 1 (& prevent divide by zero)
        melSpec -= melSpec.min()
        melSpecMax = melSpec.max()
        if melSpecMax: 
            melSpec /= melSpecMax

        maxVal = 255.9
        melSpec *= maxVal
        melSpec = maxVal-melSpec

        # Resize
        specImagePil = Image.fromarray(melSpec.astype(np.uint8))
        specImagePil = specImagePil.resize(imageSize, interpolationMethod)

        # Expand to 3 channels
        specImage = specImagePil.convert('RGB')

        specImages.append(specImage)


    return specImages


def getDefaultSpecImages(test_df):

    paths = test_df.path.unique()

    nFilesToProcess = len(paths)
    print('nFilesToProcess', nFilesToProcess)

    specImages = []
    fileIxStart = 0
    fileIxEnd = fileIxStart + nWorkers
    with Parallel(n_jobs=nWorkers) as parallel:

        while fileIxStart < nFilesToProcess:

            pathsToProcessInParallel = paths[fileIxStart:fileIxEnd]
            print('Preprocess', pathsToProcessInParallel)
            
            specImagesPerFile = parallel(delayed(getDefaultSpecImagesPerFile)(path) for path in pathsToProcessInParallel)
            # Flatten list
            specImagesPerFile = [item for sublist in specImagesPerFile for item in sublist]
            specImages += specImagesPerFile

            fileIxStart = fileIxEnd
            fileIxEnd = fileIxStart + nWorkers

    return specImages


def getModelAndConfigParams(checkpointDict):


    # Load model and cfg from torchscript checkpoint file
    extra_files = {'cfg.json': ''}
    model = torch.jit.load(checkpointDict['path'], _extra_files=extra_files)
    cfg = json.loads(extra_files['cfg.json'])

    print('Loaded:', cfg['runId'], 'with encoder:', cfg['baseEncoder'])

    # Add model and cfg to checkpointDict
    checkpointDict['model'] = model
    checkpointDict['cfg'] = cfg

    # Add relevant (BirdClef2023) indices
    modelClassIds = cfg['classIds']
    checkpointDict['birdclef2023_ixs'] = np.where(np.isin(modelClassIds, classIdsBirdClef2023))[0]

    return checkpointDict


def getPredictionsPerModelForSelectedFilePartsDefault(rowIxs, nRowsTotal, checkpoints):

    global checkTimePerBatch
    global maxInferenceTimeReached

    nCheckpoints = len(checkpoints)
    nClasses = 659
    birdclef2023Ixs = np.array([0, 1, 2, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 27, 33, 36, 43, 45, 46, 51, 53, 61, 62, 66, 67, 68, 69, 71, 73, 76, 77, 78, 81, 82, 83, 85, 86, 87, 90, 98, 99, 100, 102, 104, 108, 109, 110, 112, 113, 114, 117, 118, 121, 123, 131, 134, 141, 143, 149, 150, 157, 158, 159, 162, 163, 164, 172, 174, 176, 177, 185, 190, 194, 195, 198, 199, 200, 201, 205, 209, 213, 214, 215, 217, 218, 219, 222, 223, 226, 228, 230, 233, 236, 240, 241, 244, 245, 246, 252, 256, 258, 260, 261, 266, 273, 274, 277, 279, 280, 281, 282, 283, 284, 285, 286, 288, 289, 290, 296, 298, 302, 303, 307, 309, 311, 312, 315, 318, 324, 325, 327, 330, 331, 333, 336, 338, 340, 343, 344, 345, 348, 350, 351, 354, 357, 358, 364, 368, 369, 370, 372, 373, 377, 381, 395, 397, 398, 405, 406, 410, 418, 420, 422, 423, 425, 426, 428, 429, 431, 434, 435, 437, 439, 440, 441, 442, 444, 445, 446, 448, 452, 454, 458, 460, 475, 476, 482, 491, 493, 495, 497, 500, 501, 503, 504, 505, 509, 513, 514, 516, 517, 519, 520, 521, 522, 523, 524, 528, 530, 537, 539, 540, 542, 545, 548, 549, 550, 554, 558, 566, 569, 571, 572, 573, 575, 576, 577, 587, 588, 589, 590, 593, 594, 595, 597, 600, 601, 604, 605, 608, 617, 619, 621, 622, 625, 627, 628, 630, 632, 634, 638, 639, 640, 645, 648, 650, 652, 653, 655, 657])
    predsPerModel = torch.full((nCheckpoints, nRowsTotal, nClasses), -1, dtype=torch.float32)

    # Default normalize, dataset and dataloader
    normalize = transforms.Normalize(mean=[0.5, 0.4, 0.3], std=[0.5, 0.3, 0.1])
    testDatasetDefault = AudioDatasetDefaultSpecImage(rowIxs, transform=transforms.Compose([transforms.ToTensor(),normalize]))
    testLoaderDefault = DataLoader(testDatasetDefault, batch_size=batchSizeForTesting, shuffle=False, num_workers=nWorkers, pin_memory=True)

    with torch.no_grad():
        for i, sample_batched in enumerate(testLoaderDefault):

            if (time.time() - inferenceStartTime) > maxInferenceTime:
                maxInferenceTimeReached = True
                print('maxInferenceTimeReached')
                break

            input = sample_batched['specImage']
            rowIxsBatched = sample_batched['rowIx']
            
            for cpIx in range(nCheckpoints):
                model = checkpoints[cpIx]['model']
                output = model(input)
                predsPerModel[cpIx, rowIxsBatched] = output


    predsPerModel = predsPerModel.numpy()

    # Filter prediction columns to BirdClef2023 classes
    predsPerModel = predsPerModel[:,:,birdclef2023Ixs]

    return predsPerModel


def getPredictionsPerModelForSelectedFileParts(rowIxs, nRowsTotal, checkpoints):

    global checkTimePerBatch
    global maxInferenceTimeReached

    nCheckpoints = len(checkpoints)
    predsPerModel = np.full((nCheckpoints, nRowsTotal, nClassesBirdClef2023), -1, dtype=np.float32)

    # Default normalize, dataset and dataloader
    normalize = transforms.Normalize(mean=[0.5, 0.4, 0.3], std=[0.5, 0.3, 0.1])
    testDatasetDefault = AudioDatasetDefaultSpecImage(rowIxs, transform=transforms.Compose([transforms.ToTensor(),normalize]))
    testLoaderDefault = DataLoader(testDatasetDefault, batch_size=batchSizeForTesting, shuffle=False, num_workers=nWorkers, pin_memory=True)

    for cpIx in range(len(checkpoints)):

        checkpoint = checkpoints[cpIx]
        model = checkpoint['model']
        cfg = checkpoint['cfg']
        nClasses = len(cfg['classIds']) # 659

        # Get predictions

        # Create output matrix tensor of size nRows, nClasses
        outputs = torch.full((nRowsTotal, nClasses), -1, dtype=torch.float32)

        with torch.no_grad():
            for i, sample_batched in enumerate(testLoaderDefault):

                if checkTimePerBatch:
                    if (time.time() - inferenceStartTime) > maxInferenceTime:
                        maxInferenceTimeReached = True
                        print('maxInferenceTimeReached')
                        break

                input = sample_batched['specImage']
                rowIxsBatched = sample_batched['rowIx']
                output = model(input)
                outputs[rowIxsBatched] = output

        preds = outputs.cpu().data.numpy()

        # Filter prediction columns to BirdClef2023 classes
        preds = preds[:, checkpoints[cpIx]['birdclef2023_ixs']]

        predsPerModel[cpIx] = preds

        print('cpIx', cpIx, 'preds.shape', preds.shape)
        

        if maxInferenceTimeReached:
            break

        if (time.time() - inferenceStartTime) > checkTimePerBatchTime:
            print('checkTimePerBatchTime reached')
            checkTimePerBatch = True

    
    return predsPerModel


# # Dataset
class AudioDatasetDefaultSpecImage(Dataset):
    
    def __init__(self, rowIxs, transform=None):
        print('Use default spec image')
        self.rowIxs = rowIxs
        self.transform = transform

    def __len__(self):
        return len(self.rowIxs)

    def __getitem__(self, segmentIx):
        rowIx = self.rowIxs[segmentIx]
        specImage = specImages[rowIx]
        if self.transform: specImage = self.transform(specImage)
        return {'specImage': specImage, 'rowIx': rowIx}

# # Main

# BirdClef2023 Classes
classIdsBirdClef2023 = ['abethr1', 'abhori1', 'abythr1', 'afbfly1', 'afdfly1', 'afecuc1', 'affeag1', 'afgfly1', 'afghor1', 'afmdov1', 'afpfly1', 'afpkin1', 'afpwag1', 'afrgos1', 'afrgrp1', 'afrjac1', 'afrthr1', 'amesun2', 'augbuz1', 'bagwea1', 'barswa', 'bawhor2', 'bawman1', 'bcbeat1', 'beasun2', 'bkctch1', 'bkfruw1', 'blacra1', 'blacuc1', 'blakit1', 'blaplo1', 'blbpuf2', 'blcapa2', 'blfbus1', 'blhgon1', 'blhher1', 'blksaw1', 'blnmou1', 'blnwea1', 'bltapa1', 'bltbar1', 'bltori1', 'blwlap1', 'brcale1', 'brcsta1', 'brctch1', 'brcwea1', 'brican1', 'brobab1', 'broman1', 'brosun1', 'brrwhe3', 'brtcha1', 'brubru1', 'brwwar1', 'bswdov1', 'btweye2', 'bubwar2', 'butapa1', 'cabgre1', 'carcha1', 'carwoo1', 'categr', 'ccbeat1', 'chespa1', 'chewea1', 'chibat1', 'chtapa3', 'chucis1', 'cibwar1', 'cohmar1', 'colsun2', 'combul2', 'combuz1', 'comsan', 'crefra2', 'crheag1', 'crohor1', 'darbar1', 'darter3', 'didcuc1', 'dotbar1', 'dutdov1', 'easmog1', 'eaywag1', 'edcsun3', 'egygoo', 'equaka1', 'eswdov1', 'eubeat1', 'fatrav1', 'fatwid1', 'fislov1', 'fotdro5', 'gabgos2', 'gargan', 'gbesta1', 'gnbcam2', 'gnhsun1', 'gobbun1', 'gobsta5', 'gobwea1', 'golher1', 'grbcam1', 'grccra1', 'grecor', 'greegr', 'grewoo2', 'grwpyt1', 'gryapa1', 'grywrw1', 'gybfis1', 'gycwar3', 'gyhbus1', 'gyhkin1', 'gyhneg1', 'gyhspa1', 'gytbar1', 'hadibi1', 'hamerk1', 'hartur1', 'helgui', 'hipbab1', 'hoopoe', 'huncis1', 'hunsun2', 'joygre1', 'kerspa2', 'klacuc1', 'kvbsun1', 'laudov1', 'lawgol', 'lesmaw1', 'lessts1', 'libeat1', 'litegr', 'litswi1', 'litwea1', 'loceag1', 'lotcor1', 'lotlap1', 'luebus1', 'mabeat1', 'macshr1', 'malkin1', 'marsto1', 'marsun2', 'mcptit1', 'meypar1', 'moccha1', 'mouwag1', 'ndcsun2', 'nobfly1', 'norbro1', 'norcro1', 'norfis1', 'norpuf1', 'nubwoo1', 'pabspa1', 'palfly2', 'palpri1', 'piecro1', 'piekin1', 'pitwhy', 'purgre2', 'pygbat1', 'quailf1', 'ratcis1', 'raybar1', 'rbsrob1', 'rebfir2', 'rebhor1', 'reboxp1', 'reccor', 'reccuc1', 'reedov1', 'refbar2', 'refcro1', 'reftin1', 'refwar2', 'rehblu1', 'rehwea1', 'reisee2', 'rerswa1', 'rewsta1', 'rindov', 'rocmar2', 'rostur1', 'ruegls1', 'rufcha2', 'sacibi2', 'sccsun2', 'scrcha1', 'scthon1', 'shesta1', 'sichor1', 'sincis1', 'slbgre1', 'slcbou1', 'sltnig1', 'sobfly1', 'somgre1', 'somtit4', 'soucit1', 'soufis1', 'spemou2', 'spepig1', 'spewea1', 'spfbar1', 'spfwea1', 'spmthr1', 'spwlap1', 'squher1', 'strher', 'strsee1', 'stusta1', 'subbus1', 'supsta1', 'tacsun1', 'tafpri1', 'tamdov1', 'thrnig1', 'trobou1', 'varsun2', 'vibsta2', 'vilwea1', 'vimwea1', 'walsta1', 'wbgbir1', 'wbrcha2', 'wbswea1', 'wfbeat1', 'whbcan1', 'whbcou1', 'whbcro2', 'whbtit5', 'whbwea1', 'whbwhe3', 'whcpri2', 'whctur2', 'wheslf1', 'whhsaw1', 'whihel1', 'whrshr1', 'witswa1', 'wlwwar', 'wookin1', 'woosan', 'wtbeat1', 'yebapa1', 'yebbar1', 'yebduc1', 'yebere1', 'yebgre1', 'yebsto1', 'yeccan1', 'yefcan', 'yelbis1', 'yenspu1', 'yertin1', 'yesbar1', 'yespet1', 'yetgre1', 'yewgre1']
nClassesBirdClef2023 = 264

# Get dataframe from test files (filename, path, row_id, start_time, end_time)
test_df = get_test_df()

# Get default spec images for all file parts
specImages = getDefaultSpecImages(test_df)

# Get models and config params
for cp in checkpoints:
    cp = getModelAndConfigParams(cp)


# Get list if indices of file parts
rowIxs = list(range(len(test_df)))
nRowsTotal = len(rowIxs)


if useSecondPass:

    # Get predictions per model (first pass) using all file parts
    predsPerModelFirstPass = getPredictionsPerModelForSelectedFilePartsDefault(
        rowIxs,
        nRowsTotal,
        checkpoints[:nCheckpointsFirstPass] # first nCheckpointsFirstPass
    )
    print('predsPerModelFirstPass.shape', predsPerModelFirstPass.shape)


    # Get standard deviation between models
    predsPerModelFirstPassStd = np.std(predsPerModelFirstPass, axis=0)
    # Get mean of standard deviation (per part over all classes)
    predsPerModelFirstPassStdMean = np.mean(predsPerModelFirstPassStd, axis=1)
    # Multiply with weight_factor (mean over models and classes)
    # low weight_factor means predictions for all classes are near zero (no detection)
    weight_factor = np.mean(np.mean(predsPerModelFirstPass, axis=0), axis=1)
    attention = predsPerModelFirstPassStdMean * weight_factor
    # Get indices of parts sorted by attention (descending)
    rowIxsSortedByAttention = list(np.argsort(attention)[::-1])
    # Remove parts with low attention
    nPartsToKeep = int(len(rowIxsSortedByAttention) * percentagePartsToKeepForSecondPass)
    rowIxsSortedByAttention = rowIxsSortedByAttention[:nPartsToKeep]
    
    # Get predictions per model (second pass)
    # Use only selected file parts
    predsPerModelSecondPass = getPredictionsPerModelForSelectedFileParts(
        rowIxsSortedByAttention,
        nRowsTotal,
        checkpoints[nCheckpointsFirstPass:]
    )
    print('predsPerModelSecondPass.shape', predsPerModelSecondPass.shape)

    # Get predictions per model (all passes)
    predsPerModel = np.concatenate((predsPerModelFirstPass, predsPerModelSecondPass), axis=0)

else:

    # Use all fileparts and checkpoints
    predsPerModel = getPredictionsPerModelForSelectedFilePartsDefault(
        rowIxs,
        nRowsTotal,
        checkpoints     
    )


print('predsPerModel.shape', predsPerModel.shape)


# Apply postprocessing per model ?


# Average over checkpoints

# Mask missing values
predsPerModelMasked = np.ma.masked_equal(predsPerModel, -1)

if ensembleMethod == 'mean':
    predictions = np.mean(predsPerModelMasked, axis=0)
if ensembleMethod == 'max':
    predictions = np.max(predsPerModelMasked, axis=0)
if ensembleMethod == 'meanexp':
    predictions = np.mean(predsPerModelMasked ** 2, axis=0)

print('predictions.shape', predictions.shape)


# Apply postprocessing on ensemble ?


# Convert to df
prediction_df = pd.DataFrame(predictions, columns=classIdsBirdClef2023)
prediction_df.insert(loc=0, column='row_id', value=test_df['row_id'].tolist())

# Write csv
if outputPrecision:
    float_format = '%.' + str(outputPrecision) + 'f'
    prediction_df.to_csv('submission.csv', index=False, float_format=float_format)
else:
    prediction_df.to_csv('submission.csv', index=False)


print('ElapsedTime [s]: ', (time.time() - timeStart))
print('Done.')