import torch
import torchaudio
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from i_dataset import StimuliSounds
from ii_structure import SimpleLSTM

## Configuration
# -------------------------------------------
CSV_FILE = "../../dataset.csv"
AUDIO_DIR = "../../sound/"
SAMPLES = 6394

mel_transform = torch.nn.Sequential(
    torchaudio.transforms.MelSpectrogram(
        sample_rate=16000, n_fft=1024, win_length=320,
        hop_length=44, window_fn=torch.hamming_window, n_mels=128
    ),
    torchaudio.transforms.AmplitudeToDB()
)

## Load pytorch file
# -------------------------------------------
model = SimpleLSTM(input_size=128, hidden_size=64)
model.load_state_dict(torch.load("model.pt"))   # read saved weights

## Load from lightning logs
# -------------------------------------------
# best val_loss:
# model = SimpleLSTM.load_from_checkpoint("lightning_logs/version_53/checkpoints/epoch=49-step=10100.ckpt")
# model = model.to('cpu')

## Load dataset and get unique sitmuli
# -------------------------------------------
dataset = StimuliSounds(CSV_FILE, AUDIO_DIR, SAMPLES, mel_transform)
csv = pd.read_csv(CSV_FILE)
unique_files = csv.iloc[:, 0].unique()

## Predict ratings for each unique stimuli
# -------------------------------------------
results = []
with torch.no_grad():
    for stimulus in unique_files:
        idx = csv[csv.iloc[:, 0] == stimulus].index[0]
        mel, _ = dataset[idx]   # (146, 128), rating
        mel = mel.unsqueeze(0)  # (1, time, 128)
        pred = model.forward(mel).item()    # get prediction for a file
        results.append((stimulus, pred))

# convert the results tuple to dataframe of two columsn
model_results = pd.DataFrame(results, columns=['audio_file', 'predicted'])

## Parse metadata from file name
# -------------------------------------------
# {continuum}_{accent}-eng_{step}.wav
def parse_filename(fname):
    parts = fname.split('_')
    continuum = int(parts[0])           # Cruː or Cɔt
    accent = parts[1].split('-')[0]     # asm or amr
    vot = int(parts[2])                 # vot step
    accent_map = {'amr': 'AmE', 'asm': 'AsE'}
    continuum_map = {1: 'Cɔt', 2: 'Cruː'}
    return vot, accent_map[accent], continuum_map[continuum]

# make metadata columns
model_results[['vot', 'english', 'continuum']] = model_results['audio_file'].apply(
    lambda x: pd.Series(parse_filename(x))
)

# Plot results
# -------------------------------------------
# sort
model_results = model_results.sort_values(['vot', 'english', 'continuum'])

# plot
plt.figure(figsize=(8,8))
g = sns.FacetGrid(model_results, col='continuum', hue='english',
                  palette={'AmE':'steelblue', 'AsE':'coral'})
g.map(plt.plot, 'vot', 'predicted', marker='o')
g.map(plt.scatter, 'vot', 'predicted')
g.set(xticks=range(1, 8), xticklabels=range(1, 8))   # ticks at 1, 2, 3, 4, 5,6,7
for ax in g.axes.flat:
    ax.grid(True, color='grey', linestyle='--', alpha=0.3, linewidth=0.5)
g.add_legend()
g.set_ylabels('Predicted rating')
g.set_xlabels('VOT step')
plt.savefig("predicted_trained.png", dpi=300, bbox_inches='tight')
plt.show()