import os

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 = 7140 # 119 minutes
maxInferenceTimeReached = False

useSecondPass = False
nCheckpointsFirstPass = 5
percentagePartsToKeepForSecondPass = 0.7
iterModeSecondPass = 'data' # 'data' 'checkpoints'

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


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

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


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

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

if runLocal:
    nWorkers = 4
    
    
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'})

# NNN

# 15, b0, f5, SED (205/199)
checkpoints.append({'path': checkpointRootDir + '61_0_53_9_SED/Checkpoint195.pt'})
# 16, b0, f0, MultiLabelSoftMarginLoss
checkpoints.append({'path': checkpointRootDir + '60_0_53_9_MultiLabelSoftMarginLoss/Checkpoint195.pt'})
# 17, b0, f5, AddBc2022DataAsNoise
checkpoints.append({'path': checkpointRootDir + '60_3_53_9_AddBc2022DataAsNoise/Checkpoint197.pt'})

# 18, v2, f3, latest
checkpoints.append({'path': checkpointRootDir + '63_5_63_3_AddXcDl2First10Sec_LessTestSc/Checkpoint215.pt'})



# Get checkpoint selection in specific order

# 3rd place solution
cpIxs = [0,1,6,7,12, 3,4,13,10]
useSecondPass = True
nCheckpointsFirstPass = 5
percentagePartsToKeepForSecondPass = 0.7
iterModeSecondPass = 'data'



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


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


if not nWorkers:
    nWorkers = multiprocessing.cpu_count()


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}


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

    return checkpointDict

def getPredictionsPerModelForSelectedFileParts(rowIxs, nRowsTotal, checkpoints, iterMode='data'):

    global checkTimePerBatch
    global maxInferenceTimeReached

    nCheckpoints = len(checkpoints)
    predsPerModel = torch.full((nCheckpoints, nRowsTotal, nClassesBirdClef2021Plus2023), -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)

    print('iterMode', iterMode)
    
    if iterMode == 'data':

        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


    if iterMode == 'checkpoints':

        for cpIx in range(len(checkpoints)):

            checkpoint = checkpoints[cpIx]
            model = checkpoint['model']

            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)
                    predsPerModel[cpIx, rowIxsBatched] = output

            print('cpIx', cpIx)
            

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

            if timeSinceInferenceStart > checkTimePerBatchTime:
                print('checkTimePerBatchTime reached')
                checkTimePerBatch = True



    predsPerModel = predsPerModel.numpy()

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

    return predsPerModel


# # Main

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

classIdsBirdClef2021Plus2023 = ['abethr1', 'abhori1', 'abythr1', 'acafly', 'acowoo', 'afbfly1', 'afdfly1', 'afecuc1', 'affeag1', 'afgfly1', 'afghor1', 'afmdov1', 'afpfly1', 'afpkin1', 'afpwag1', 'afrgos1', 'afrgrp1', 'afrjac1', 'afrthr1', 'aldfly', 'ameavo', 'amecro', 'amegfi', 'amekes', 'amepip', 'amered', 'amerob', 'amesun2', 'amewig', 'amtspa', 'andsol1', 'annhum', 'astfly', 'augbuz1', 'azaspi1', 'babwar', 'bagwea1', 'baleag', 'balori', 'banana', 'banswa', 'banwre1', 'barant1', 'barswa', 'batpig1', 'bawhor2', 'bawman1', 'bawswa1', 'bawwar', 'baywre1', 'bbwduc', 'bcbeat1', 'bcnher', 'beasun2', 'belkin1', 'belvir', 'bewwre', 'bkbmag1', 'bkbplo', 'bkbwar', 'bkcchi', 'bkctch1', 'bkfruw1', 'bkhgro', 'bkmtou1', 'bknsti', 'blacra1', 'blacuc1', 'blakit1', 'blaplo1', 'blbgra1', 'blbpuf2', 'blbthr1', 'blcapa2', 'blcjay1', 'blctan1', 'blfbus1', 'blhgon1', 'blhher1', 'blhpar1', 'blkpho', 'blksaw1', 'blnmou1', 'blnwea1', 'blsspa1', 'bltapa1', 'bltbar1', 'bltori1', 'blugrb1', 'blujay', 'blwlap1', 'bncfly', 'bnhcow', 'bobfly1', 'bongul', 'botgra', 'brbmot1', 'brbsol1', 'brcale1', 'brcsta1', 'brctch1', 'brcvir1', 'brcwea1', 'brebla', 'brican1', 'brncre', 'brnjay', 'brnthr', 'brobab1', 'broman1', 'brosun1', 'brratt1', 'brrwhe3', 'brtcha1', 'brubru1', 'brwhaw', 'brwpar1', 'brwwar1', 'bswdov1', 'btbwar', 'btnwar', 'btweye2', 'btywar', 'bubwar2', 'bucmot2', 'buggna', 'bugtan', 'buhvir', 'bulori', 'burwar1', 'bushti', 'butapa1', 'butsal1', 'buwtea', 'cabgre1', 'cacgoo1', 'cacwre', 'calqua', 'caltow', 'cangoo', 'canwar', 'carcha1', 'carchi', 'carwoo1', 'carwre', 'casfin', 'caskin', 'caster1', 'casvir', 'categr', 'ccbeat1', 'ccbfin', 'cedwax', 'chbant1', 'chbchi', 'chbwre1', 'chcant2', 'chespa1', 'chewea1', 'chibat1', 'chispa', 'chswar', 'chtapa3', 'chucis1', 'cibwar1', 'cinfly2', 'clanut', 'clcrob', 'cliswa', 'cobtan1', 'cocwoo1', 'cogdov', 'cohmar1', 'colcha1', 'colsun2', 'coltro1', 'combul2', 'combuz1', 'comgol', 'comgra', 'comloo', 'commer', 'compau', 'compot1', 'comrav', 'comsan', 'comyel', 'coohaw', 'cotfly1', 'cowscj1', 'crefra2', 'cregua1', 'creoro1', 'crfpar', 'crheag1', 'crohor1', 'cubthr', 'daejun', 'darbar1', 'darter3', 'didcuc1', 'dotbar1', 'dowwoo', 'ducfly', 'dusfly', 'dutdov1', 'easblu', 'easkin', 'easmea', 'easmog1', 'easpho', 'eastow', 'eawpew', 'eaywag1', 'edcsun3', 'egygoo', 'eletro', 'equaka1', 'eswdov1', 'eubeat1', 'eucdov', 'eursta', 'fatrav1', 'fatwid1', 'fepowl', 'fiespa', 'fislov1', 'flrtan1', 'fotdro5', 'foxspa', 'gabgos2', 'gadwal', 'gamqua', 'gargan', 'gartro1', 'gbbgul', 'gbesta1', 'gbwwre1', 'gcrwar', 'gilwoo', 'gnbcam2', 'gnhsun1', 'gnttow', 'gnwtea', 'gobbun1', 'gobsta5', 'gobwea1', 'gocfly1', 'gockin', 'gocspa', 'goftyr1', 'gohque1', 'golher1', 'goowoo1', 'grasal1', 'grbani', 'grbcam1', 'grbher3', 'grccra1', 'grcfly', 'grecor', 'greegr', 'grekis', 'grepew', 'grethr1', 'gretin1', 'grewoo2', 'greyel', 'grhcha1', 'grhowl', 'grnher', 'grnjay', 'grtgra', 'grwpyt1', 'gryapa1', 'grycat', 'gryhaw2', 'grywrw1', 'gwfgoo', 'gybfis1', 'gycwar3', 'gyhbus1', 'gyhkin1', 'gyhneg1', 'gyhspa1', 'gytbar1', 'hadibi1', 'haiwoo', 'hamerk1', 'hartur1', 'helgui', 'heptan', 'hergul', 'herthr', 'herwar', 'higmot1', 'hipbab1', 'hofwoo1', 'hoopoe', 'houfin', 'houspa', 'houwre', 'huncis1', 'hunsun2', 'hutvir', 'incdov', 'indbun', 'joygre1', 'kebtou1', 'kerspa2', 'killde', 'klacuc1', 'kvbsun1', 'labwoo', 'larspa', 'laudov1', 'laufal1', 'laugul', 'lawgol', 'lazbun', 'leafly', 'leasan', 'lesgol', 'lesgre1', 'lesmaw1', 'lessts1', 'lesvio1', 'libeat1', 'linspa', 'linwoo1', 'litegr', 'litswi1', 'littin1', 'litwea1', 'lobdow', 'lobgna5', 'loceag1', 'logshr', 'lotcor1', 'lotduc', 'lotlap1', 'lotman1', 'lucwar', 'luebus1', 'mabeat1', 'macshr1', 'macwar', 'magwar', 'malkin1', 'mallar3', 'marsto1', 'marsun2', 'marwre', 'mastro1', 'mcptit1', 'meapar', 'melbla1', 'meypar1', 'moccha1', 'monoro1', 'mouchi', 'moudov', 'mouela1', 'mouqua', 'mouwag1', 'mouwar', 'mutswa', 'naswar', 'ndcsun2', 'nobfly1', 'norbro1', 'norcar', 'norcro1', 'norfis1', 'norfli', 'normoc', 'norpar', 'norpuf1', 'norsho', 'norwat', 'nrwswa', 'nubwoo1', 'nutwoo', 'oaktit', 'obnthr1', 'ocbfly1', 'oliwoo1', 'olsfly', 'orbeup1', 'orbspa1', 'orcpar', 'orcwar', 'orfpar', 'osprey', 'ovenbi1', 'pabspa1', 'pabspi1', 'palfly2', 'palpri1', 'paltan1', 'palwar', 'pasfly', 'pavpig2', 'phivir', 'pibgre', 'piecro1', 'piekin1', 'pilwoo', 'pinsis', 'pirfly1', 'pitwhy', 'plawre1', 'plaxen1', 'plsvir', 'plupig2', 'prowar', 'purfin', 'purgal2', 'purgre2', 'putfru1', 'pygbat1', 'pygnut', 'quailf1', 'ratcis1', 'rawwre1', 'raybar1', 'rbsrob1', 'rcatan1', 'rebfir2', 'rebhor1', 'rebnut', 'reboxp1', 'rebsap', 'rebwoo', 'reccor', 'reccuc1', 'redcro', 'reedov1', 'reevir1', 'refbar2', 'refcro1', 'reftin1', 'refwar2', 'rehbar1', 'rehblu1', 'rehwea1', 'reisee2', 'relpar', 'rerswa1', 'reshaw', 'rethaw', 'rewbla', 'rewsta1', 'ribgul', 'rindov', 'rinkin1', 'roahaw', 'robgro', 'rocmar2', 'rocpig', 'rostur1', 'rotbec', 'royter1', 'rthhum', 'rtlhum', 'ruboro1', 'rubpep1', 'rubrob', 'rubwre1', 'ruckin', 'rucspa1', 'rucwar', 'rucwar1', 'rudpig', 'rudtur', 'ruegls1', 'rufcha2', 'rufhum', 'rugdov', 'rumfly1', 'runwre1', 'rutjac1', 'sacibi2', 'saffin', 'sancra', 'sander', 'savspa', 'saypho', 'scamac1', 'scatan', 'scbwre1', 'sccsun2', 'scptyr1', 'scrcha1', 'scrtan1', 'scthon1', 'semplo', 'shesta1', 'shicow', 'sibtan2', 'sichor1', 'sincis1', 'sinwre1', 'slbgre1', 'slcbou1', 'sltnig1', 'sltred', 'smbani', 'snogoo', 'sobfly1', 'sobtyr1', 'socfly1', 'solsan', 'somgre1', 'somtit4', 'sonspa', 'soucit1', 'soufis1', 'soulap1', 'spemou2', 'spepig1', 'spewea1', 'spfbar1', 'spfwea1', 'spmthr1', 'sposan', 'spotow', 'spvear1', 'spwlap1', 'squcuc1', 'squher1', 'stbori', 'stejay', 'sthant1', 'sthwoo1', 'strcuc1', 'strfly1', 'strher', 'strsal1', 'strsee1', 'stusta1', 'stvhum2', 'subbus1', 'subfly', 'sumtan', 'supsta1', 'swaspa', 'swathr', 'tacsun1', 'tafpri1', 'tamdov1', 'tenwar', 'thbeup1', 'thbkin', 'thrnig1', 'thswar1', 'towsol', 'treswa', 'trobou1', 'trogna1', 'trokin', 'tromoc', 'tropar', 'tropew1', 'tuftit', 'tunswa', 'varsun2', 'veery', 'verdin', 'vibsta2', 'vigswa', 'vilwea1', 'vimwea1', 'walsta1', 'warvir', 'wbgbir1', 'wbrcha2', 'wbswea1', 'wbwwre1', 'webwoo1', 'wegspa1', 'wesant1', 'wesblu', 'weskin', 'wesmea', 'westan', 'wewpew', 'wfbeat1', 'whbcan1', 'whbcou1', 'whbcro2', 'whbman1', 'whbnut', 'whbtit5', 'whbwea1', 'whbwhe3', 'whcpar', 'whcpri2', 'whcsee1', 'whcspa', 'whctur2', 'wheslf1', 'whevir', 'whfpar1', 'whhsaw1', 'whihel1', 'whimbr', 'whiwre1', 'whrshr1', 'whtdov', 'whtspa', 'whwbec1', 'whwdov', 'wilfly', 'willet1', 'wilsni1', 'wiltur', 'witswa1', 'wlswar', 'wlwwar', 'wooduc', 'wookin1', 'woosan', 'woothr', 'wrenti', 'wtbeat1', 'y00475', 'yebapa1', 'yebbar1', 'yebcha', 'yebduc1', 'yebela1', 'yebere1', 'yebfly', 'yebgre1', 'yebori1', 'yebsap', 'yebsee1', 'yebsto1', 'yeccan1', 'yefcan', 'yefgra1', 'yegvir', 'yehbla', 'yehcar1', 'yelbis1', 'yelgro', 'yelwar', 'yenspu1', 'yeofly1', 'yertin1', 'yerwar', 'yesbar1', 'yespet1', 'yeteup1', 'yetgre1', 'yetvir', 'yewgre1', 'nocall']
nClassesBirdClef2021Plus2023 = 659

#birdclef2023Ixs = np.where(np.isin(classIdsBirdClef2021Plus2023, classIdsBirdClef2023))[0]
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])



# 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 of indices for all file parts
rowIxs = list(range(len(test_df)))
nRowsTotal = len(rowIxs)


if useSecondPass:

    # Get predictions per model (first pass) using all file parts
    predsPerModelFirstPass = getPredictionsPerModelForSelectedFileParts(
        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:],
        iterMode=iterModeSecondPass
    )
    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 = getPredictionsPerModelForSelectedFileParts(
        rowIxs,
        nRowsTotal,
        checkpoints     
    )


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



# 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)


# 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() - inferenceStartTime))
print('Done.') 