# this is just a python/utility script version of my training notebook:
# https://www.kaggle.com/code/max1mum/pytorch-custom-model-train-5-min-train

import os
import gc
import cv2
import math
import numpy as np
import pandas as pd
from tqdm.notebook import tqdm
import matplotlib.pyplot as plt
import librosa
from scipy import signal as sci_signal
import sklearn
from datetime import datetime

import torch
from torch import nn

import kaggle_metric_utilities

import albumentations as albu

class config:
    SEED = 14
    DEVICE = 'cpu'
    MIXED_PRECISION = False
    OUTPUT_DIR = 'out_dir'
    
    DATA_ROOT = '/kaggle/input/birdclef-2024'
    PREPROCESSED_DATA_ROOT = '/kaggle/input/5-min-checkpoint'
    LOAD_DATA = True
    FS = 32000
    N_FFT = 1095
    WIN_SIZE = 412
    WIN_LAP = 100
    MIN_FREQ = 40
    MAX_FREQ = 15000 
    
    BATCH_SIZE = 32
    N_WORKERS = 12
    
    USE_XYMASKING = True
    
    FOLDS = 10
    EPOCHS = 5
    LR = 8e-5
    WEIGHT_DECAY = 1e-5
    
    VISUALIZE = True

def set_seed(seed=14):
    os.environ["PYTHONHASHSEED"] = str(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False
    
set_seed()

label_list = sorted(os.listdir(os.path.join(config.DATA_ROOT, 'train_audio')))
label_id_list = list(range(len(label_list)))
label2id = dict(zip(label_list, label_id_list))
id2label = dict(zip(label_id_list, label_list))

metadata_df = pd.read_csv(f'{config.DATA_ROOT}/train_metadata.csv')
metadata_df.head()

train_df = metadata_df[['primary_label', 'rating', 'filename']].copy()

train_df['target'] = train_df.primary_label.map(label2id)
train_df['filepath'] = config.DATA_ROOT + '/train_audio/' + train_df.filename
train_df['samplename'] = train_df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])

train_df.head()

def scipy_spectro(audio_data):
    mean_signal = np.nanmean(audio_data)
    audio_data = np.nan_to_num(audio_data, nan=mean_signal) if np.isnan(audio_data).mean() < 1 else np.zeros_like(audio_data)
    
    frequencies, times, spec_data = sci_signal.spectrogram(
        audio_data, 
        fs=config.FS, 
        nfft=config.N_FFT, 
        nperseg=config.WIN_SIZE, 
        noverlap=config.WIN_LAP, 
        window='hann'
    )
    
    valid_freq = (frequencies >= config.MIN_FREQ) & (frequencies <= config.MAX_FREQ)
    spec_data = spec_data[valid_freq, :]
    
    spec_data = np.log10(spec_data + 1e-20)
    
    spec_data = spec_data - spec_data.min()
    spec_data = spec_data / spec_data.max()
    
    return spec_data

def cupy_spectro(audio_data):    
    audio_data = cp.array(audio_data)
    
    mean_signal = cp.nanmean(audio_data)
    audio_data = cp.nan_to_num(audio_data, nan=mean_signal) if cp.isnan(audio_data).mean() < 1 else cp.zeros_like(audio_data)
    
    frequencies, times, spec_data = cupy_signal.spectrogram(
        audio_data, 
        fs=config.FS, 
        nfft=config.N_FFT, 
        nperseg=config.WIN_SIZE, 
        noverlap=config.WIN_LAP, 
        window='hann'
    )
    
    valid_freq = (frequencies >= config.MIN_FREQ) & (frequencies <= config.MAX_FREQ)
    spec_data = spec_data[valid_freq, :]
    
    spec_data = cp.log10(spec_data + 1e-20)
    
    spec_data = spec_data - spec_data.min()
    spec_data = spec_data / spec_data.max()
    
    return spec_data.get()

if config.LOAD_DATA:
    all_bird_data = np.load(f'/kaggle/input/birdclef-preprocessed/spec_center_5sec_256_256.npy', allow_pickle=True).item()
else:
    all_bird_data = dict()
    for i, row_metadata in tqdm(train_df.iterrows()):
        audio_data, _ = librosa.load(row_metadata.filepath, sr=config.FS)

        n_copy = math.ceil(5 * config.FS / len(audio_data))
        if n_copy > 1: audio_data = np.concatenate([audio_data]*n_copy)

        start_idx = int(len(audio_data) / 2 - 2.5 * config.FS)
        end_idx = int(start_idx + 5.0 * config.FS)
        input_audio = audio_data[start_idx:end_idx]

        input_spec = cupy_spectro(input_audio)
        
        input_spec = cv2.resize(input_spec, (256, 256), interpolation=cv2.INTER_AREA)

        all_bird_data[row_metadata.samplename] = input_spec.astype(np.float32)

    np.save(os.path.join(config.PREPROCESSED_DATA_ROOT, f'spec_center_5sec_256_256.npy'), all_bird_data)

class BirdDataset(torch.utils.data.Dataset):
    def __init__(
        self,
        metadata,
        augmentation=None,
        mode='train'
    ):
        super().__init__()
        self.metadata = metadata
        self.augmentation = augmentation
        self.mode = mode
    
    def __len__(self):
        return len(self.metadata)
    
    def __getitem__(self, index):
        row_metadata = self.metadata.iloc[index]
        
        input_spec = all_bird_data[row_metadata.samplename]
        
        if self.augmentation is not None:
            input_spec = self.augmentation(image=input_spec)['image']
        
        target = row_metadata.target
        
        return torch.tensor(input_spec, dtype=torch.float32).unsqueeze(0), torch.tensor(target, dtype=torch.long)
    
def get_transforms(_type):
    if _type == 'train':
        return albu.Compose([
            albu.HorizontalFlip(p=0.5),
            albu.XYMasking(
                p=0.3,
                num_masks_x=(1, 3),
                num_masks_y=(1, 3),
                mask_x_length=(1, 10),
                mask_y_length=(1, 20),
            ) if config.USE_XYMASKING else albu.NoOp()
        ])
    elif _type == 'valid' or _type == 'test':
        return albu.Compose([])
    
dummy_dataset = BirdDataset(train_df, get_transforms('train'))

def show_batch(ds, row=3, col=3):
    fig = plt.figure(figsize=(10, 10))
    img_index = np.random.randint(0, len(ds)-1, row*col)
    
    for i in range(len(img_index)):
        img, label = dummy_dataset[img_index[i]]
        
        if isinstance(img, torch.Tensor):
            img = img.squeeze().detach().numpy()
        
        ax = fig.add_subplot(row, col, i + 1, xticks=[], yticks=[])
        ax.imshow(img)
        ax.set_title(f'ID: {img_index[i]}; Target: {label}')
    
    plt.tight_layout()
    plt.show()

class ConvBlock(nn.Module):
    def __init__(
        self, in_channels, out_channels
    ) -> None:
        super().__init__()
        self.same_channels = in_channels==out_channels

        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels, out_channels, 3, 2, 1),
            nn.BatchNorm2d(out_channels),
            nn.GELU(),
        )

    def forward(self, x):
        x = self.conv1(x)
        return x
    
class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Sequential(
            ConvBlock(1, 2),
            ConvBlock(2, 4),
            ConvBlock(4, 8),
            ConvBlock(8, 8),
            ConvBlock(8, 16),
            ConvBlock(16, 16),
            ConvBlock(16, 32),
            ConvBlock(32, 182),
        )

    def forward(self, x):
        return nn.Softmax()(self.conv(x))
    
model = Model()
# model.load_state_dict(torch.load("/kaggle/input/birdclef-preprocessed/spec_center_5sec_256_256.npy"))
# model.eval()

train_df = train_df.sample(frac=1).reset_index(drop=True)

split = int(train_df.shape[0] * 14/15)

train_df_ = train_df[:split].copy()
valid_df = train_df[split:].copy()
train_df = train_df_

print(f'Train Samples: {len(train_df)}')
print(f'Valid Samples: {len(valid_df)}')

train_ds = BirdDataset(train_df, get_transforms('train'), 'train')
val_ds = BirdDataset(valid_df, get_transforms('valid'), 'valid')
# test_ds = BirdDataset(test_df, get_transforms('valid'), 'valid')

train_dl = torch.utils.data.DataLoader(
    train_ds,
    batch_size=config.BATCH_SIZE,
    shuffle=True,
    num_workers=config.N_WORKERS,
    pin_memory=True,
    persistent_workers=True
)

val_dl = torch.utils.data.DataLoader(
    val_ds,
    batch_size=1,
    shuffle=False,
    num_workers=config.N_WORKERS,
    pin_memory=True,
    persistent_workers=True
)

epochs = 5    
loss_fn = nn.CrossEntropyLoss()
optim = torch.optim.Adam(
    filter(lambda p: p.requires_grad, model.parameters()),
    lr=config.LR
)
scaler = torch.cuda.amp.GradScaler()

epoch_i = 0
for epoch in range(epochs):
    i = 0
    sum = 0
    for batch in train_dl:
        image, target = batch
#         image = image.to(config.DEVICE)
#         target = target.to(config.DEVICE)
#         model = model.to(config.DEVICE)

        y_pred = model(image).squeeze()

        loss = loss_fn(y_pred, target)

        loss.backward()
        optim.step()

        # with torch.cuda.amp.autocast():
        #     scaler.scale(train_loss).backward()
        #     # grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=10)
        #     scaler.step(optim)
        #     scaler.update()

        i += 1
        sum += loss
    
    print("Epoch #" + str(epoch) + " | Loss: " + str((sum / i).item()))
    epoch_i+=1
    
    if epoch_i % 100 == 0 and epoch_i > 0:
        now = datetime.now()

        torch.save(model.state_dict(), os.path.join(
            config.PREPROCESSED_DATA_ROOT, f'model_save_time=' + str(now) + ' | epoch: ' + str(epoch_i) + '.npy'))
        
def predict(data_loader, model):
#     model.to(config.DEVICE)
    model.eval()
    pred = []
    for batch in data_loader:
        with torch.no_grad():
            x, _ = batch
#             x = x.to(config.DEVICE)
            outputs = model(x)
        pred.append(outputs.detach().cpu())
    
    pred = torch.cat(pred, dim=0).cpu().detach()
    
    return pred.numpy().astype(np.float32)

predictions = []

predictions.append(predict(val_dl, model))
gc.collect()

predictions = np.mean(predictions, axis=0).squeeze()

def score(solution, submission):
    if not pd.api.types.is_numeric_dtype(submission.values):
        bad_dtypes = {x: submission[x].dtype  for x in submission.columns
                      if not pd.api.types.is_numeric_dtype(submission[x])}
        raise TypeError(f'Invalid submission data types found: {bad_dtypes}')

    solution_sums = solution.sum(axis=0)
    scored_columns = list(solution_sums[solution_sums > 0].index.values)
    assert len(scored_columns) > 0

    return kaggle_metric_utilities.safe_call_score( sklearn.metrics.roc_auc_score,
                                                    solution[scored_columns].values,
                                                    submission[scored_columns].values,
                                                    average='macro'
                                                  )

i = 0
num_correct_1 = 0
num_correct_9 = 0
num_correct_18 = 0
num_correct_91 = 0
avg_correct_pos = 0
target_list = []
for _, target in val_dl:
    target_list_ = [0] * 182
    for ii in range(182):
        if ii == target[0].item():
            target_list_[ii] = 1

        if predictions[ii].argmax() == target[0]:
            if ii == 0:
                num_correct_1 += 1

            if ii <= 9:
                num_correct_9 += 1
            
            if ii <= 18:
                num_correct_18 += 1
            
            if ii <= 91:
                num_correct_91 += 1

            avg_correct_pos += ii
            
            break

    i+=1
    target_list.append(target_list_)

target_list_pd = pd.DataFrame(target_list)
predictions_pd = pd.DataFrame(predictions)

print("Correct: " + str(num_correct_1/i))
print("Correct Top 5%: " + str(num_correct_9/i))
print("Correct Top 10%: " + str(num_correct_18/i))
print("Correct Top 50%: " + str(num_correct_91/i))
print("Avg Correct Pos: " + str(avg_correct_pos/i))
# print("Score:", score(target_list_pd, predictions_pd))