# %%
import os
import glob
import random
from suno_utils.audio import Audio
from suno_utils.audio.midi import Midi
import shutil
from tqdm import tqdm

import torchcrepe
import tempfile
import torch
import pesto
import torchaudio
import numpy as np
import librosa
import copy
import matplotlib.pyplot as plt

os.environ["CUDA_VISIBLE_DEVICES"] = "4,5,6,7"

# %%
trombone_champ_dir = "/app2/suno/data/victor/trombone_champ_vocals"

# %%
class MidiPair:
    def __init__(self, midi_file: str, audio_file: str):
        self.midi_file = midi_file
        self.audio_file = audio_file
        self.midi = None
        self.audio = None

    def load_midi(self):
        if self.midi is not None:
            return self.midi
        self.midi = Midi.from_path(self.midi_file)
        return self.midi

    def load_audio(self):
        if self.audio is not None:
            return self.audio
        self.audio = Audio.from_file(self.audio_file)
        return self.audio

    def load_all(self):
        self.load_midi()
        self.load_audio()

    def __str__(self):
        return f"MidiPair(midi_file={self.midi_file}, audio_file={self.audio_file})"

    def __repr__(self):
        return self.__str__()

    def play(self):
        self.load_all()
        stereo_audio = self.midi.make_stereo_comparison(self.audio)
        stereo_audio.play()

def load_pairs(dir: str, audio_ext: str = "mp3", midi_ext: str = "mid"):
    midi_files = glob.glob(os.path.join(dir, "**", f"*.{midi_ext}"), recursive=True)
    audio_files = glob.glob(os.path.join(dir, "**", f"*.{audio_ext}"), recursive=True)

    # Create a mapping of base filenames to audio files
    audio_map = {}
    for audio_file in audio_files:
        base_name = os.path.splitext(os.path.basename(audio_file))[0]
        audio_map[base_name] = audio_file

    pairs = []
    for midi_file in midi_files:
        base_name = os.path.splitext(os.path.basename(midi_file))[0]
        if base_name in audio_map:
            pairs.append(MidiPair(midi_file, audio_map[base_name]))

    return pairs


# %% [markdown]
# ## trombone champ

# %%
trombone_champ_pairs = load_pairs(trombone_champ_dir, audio_ext="opus")
print(f"Loaded {len(trombone_champ_pairs)} trombone champ pairs")

# %% [markdown]
# ## fix octave shifts

# %%
def plot_f0_comparison(results, stem, midi):

    # Convert frequencies to MIDI note numbers (pitches)
    frequencies = torchcrepe.filter.median(results, 30)[0]

    # zero out sections that are silent
    hz = int(len(frequencies) / stem.duration_s)
    for i in range(int(stem.duration_s)):
        if stem.get_slice(i, i + 1).loudness < -35:
            frequencies[i * hz : (i + 1) * hz] = 20

    pitches = librosa.hz_to_midi(frequencies)

    # Create time axis
    time_axis = [t / hz for t in range(len(frequencies))]

    plt.figure(figsize=(12, 6))
    plt.plot(time_axis, pitches, label="Vocal F0")

    # filter out notes that are silent
    for i in reversed(range(len(midi.pmidi.instruments[0].notes))):
        audio_slice = stem.get_slice(
            midi.pmidi.instruments[0].notes[i].start, midi.pmidi.instruments[0].notes[i].end + 1
        )
        if np.mean(audio_slice.array_float**2) < 1e-6:
            midi.pmidi.instruments[0].notes.pop(i)
            

    # Overlay MIDI notes
    for note in midi.pmidi.instruments[0].notes:
        start_time = note.start
        end_time = note.end
        pitch = note.pitch
        plt.hlines(
            pitch,
            start_time,
            end_time,
            colors="red",
            linewidth=2,
            alpha=0.7,
            label="MIDI Notes" if note == midi.pmidi.instruments[0].notes[0] else "",
        )

    plt.xlabel("Time (s)")
    plt.ylabel("MIDI Note Number")
    plt.legend()
    plt.show()

# %%
def get_f0_pesto(audio):
    with tempfile.NamedTemporaryFile(suffix=".wav") as f:
        audio.write_wav(f.name)

        x, sr = torchaudio.load(f.name)
        x = x.mean(dim=0)

        _, pitch, confidence, _ = pesto.predict(x, sr)

    return pitch.clone(), confidence.clone()

def get_f0(audio):
    # Process in 60s chunks to avoid OOM
    chunk_duration = 60
    total_duration = audio.duration_s #min(240, audio.duration_s)
    all_results = []
    all_period_results = []

    for start_time in range(0, int(total_duration), chunk_duration):
        end_time = min(start_time + chunk_duration, total_duration)

        with tempfile.NamedTemporaryFile(suffix=".wav") as f:
            audio.get_slice(start_time, end_time).write_wav(f.name)
            chunk_results, period_results = torchcrepe.predict_from_file(
                f.name, device="cuda:3", decoder=torchcrepe.decode.weighted_argmax, pad=False, return_periodicity=True
            )
        all_results.append(chunk_results)
        all_period_results.append(period_results)

    # Concatenate results
    results = torch.cat([r[0] for r in all_results], dim=0).unsqueeze(0)
    period_final_results = torch.cat([r[0] for r in all_period_results], dim=0).unsqueeze(0)
    return results, period_final_results


def octave_dist(note1, note2):
    diff = abs(note1 - note2)
    return min(diff, 12 - diff)

def fix_midi_octave_errors(pair, verbose=False, pitch_algo="crepe", λ=500, confidence_threshold=0.05):
    midi = pair.load_midi()
    stem = pair.load_audio().resample(22050)

    if verbose:
        print(f"MIDI has {len(midi.pmidi.instruments)} instruments")
        if midi.pmidi.instruments:
            print(f"First instrument has {len(midi.pmidi.instruments[0].notes)} notes")
        else:
            print("No instruments found!")

    # get vocal f0
    if pitch_algo == "pesto":
        frequencies, confidences = get_f0_pesto(stem)
        results = frequencies.unsqueeze(0)
    else:
        results, confidences = get_f0(stem)  
        confidences = confidences[0]
        frequencies = torchcrepe.filter.median(results, 30)[0]

    # zero out sections that are silent OR low confidence
    hz = int(len(frequencies) / stem.duration_s)
    device = frequencies.device
    for i in range(int(stem.duration_s)):
        if stem.get_slice(i, i + 1).loudness < -35:
            frequencies[i * hz : (i + 1) * hz] = 20
            if confidences is not None:
                confidences[i * hz : (i + 1) * hz] = 0

    # Validate MIDI structure
    if not midi.pmidi.instruments:
        raise ValueError("MIDI file has no instruments")
    
    if not midi.pmidi.instruments[0].notes:
        raise ValueError("MIDI file has no notes in the first instrument")

    # viterbi octave correction
    Ks = torch.arange(-3, 4, device=device)  # [-3, -2, -1, 0, 1, 2, 3]
    N = len(midi.pmidi.instruments[0].notes)

    dp = torch.full((N, len(Ks)), float('inf'), device=device)
    prev = torch.zeros((N, len(Ks)), dtype=torch.long, device=device)
    note_estimates = []
    note_confidences = []
    
    for note in midi.pmidi.instruments[0].notes:
        start_idx, end_idx = int(note.start * hz), int(note.end * hz)
        note_pitches = frequencies[start_idx:end_idx]
        
        if confidences is not None:
            note_confs = confidences[start_idx:end_idx]
            # filter by confidence and silence
            valid_mask = (note_pitches > 20) & (note_confs > confidence_threshold)
            valid_pitches = note_pitches[valid_mask]
            valid_confs = note_confs[valid_mask]
            
            if len(valid_pitches) > 0:
                # weighted median
                median_pitch = weighted_median_torch(valid_pitches, valid_confs)
                # average confidence for this note
                avg_confidence = torch.mean(valid_confs)
            else:
                median_pitch = torch.tensor(float('nan'), device=device)
                avg_confidence = torch.tensor(0.0, device=device)
        else:
            # fallback for CREPE
            valid_pitches = note_pitches[note_pitches > 20]
            if len(valid_pitches) > 0:
                median_pitch = torch.median(valid_pitches)
            else:
                median_pitch = torch.tensor(float('nan'), device=device)
            avg_confidence = torch.tensor(1.0, device=device)
            
        # Convert Hz to MIDI
        if torch.isnan(median_pitch):
            note_estimates.append(median_pitch)
        else:
            midi_pitch = 69 + 12 * torch.log2(median_pitch / 440)
            note_estimates.append(midi_pitch)
        note_confidences.append(avg_confidence)

    notes = midi.pmidi.instruments[0].notes
    
    # init with confidence weighting
    for ik, k in enumerate(Ks):
        if torch.isnan(note_estimates[0]):
            dp[0, ik] = 0
        else:
            error = (note_estimates[0] - (notes[0].pitch + 12 * k)) ** 2
            confidence_weight = 1.0 / torch.clamp(note_confidences[0], min=0.1)
            dp[0, ik] = error * confidence_weight

    # fill with confidence weighting
    for i in range(1, N):
        for ik, k in enumerate(Ks):
            if torch.isnan(note_estimates[i]):
                obs = torch.tensor(0.0, device=device)
            else:
                error = (note_estimates[i] - (notes[i].pitch + 12 * k)) ** 2
                confidence_weight = 1.0 / torch.clamp(note_confidences[i], min=0.1)
                obs = error * confidence_weight
                
            # transition costs
            transition_costs = λ * torch.abs(k - Ks)
            costs = dp[i - 1, :] + transition_costs
            dp[i, ik] = obs + torch.min(costs)
            prev[i, ik] = torch.argmin(costs)

    # backtrack
    best_path = torch.zeros(N, dtype=torch.long, device=device)
    best_path[-1] = torch.argmin(dp[-1, :])
    for i in range(N - 2, -1, -1):
        best_path[i] = prev[i + 1, best_path[i + 1]]

    if verbose:
        print(best_path.cpu().numpy())

    cost = torch.tensor(0.0, device=device)
    for i in range(N):
        if torch.isnan(note_estimates[i]):
            cost += 0
        else:
            diff = note_estimates[i] - (notes[i].pitch + 12 * Ks[best_path[i]])
            cost += diff**2
    normalized_cost = cost / N
    if verbose:
        print(f"Cost: {cost.item()}, Normalized cost: {normalized_cost.item()}")

    new_midi = copy.deepcopy(midi)
    # apply shifts - convert back to CPU for MIDI manipulation
    best_path_cpu = best_path.cpu()
    Ks_cpu = Ks.cpu()
    for i, note in enumerate(notes):
        note.pitch += 12 * Ks_cpu[best_path_cpu[i]].item()
    
    return results, new_midi


def weighted_median_torch(values, weights):
    """Compute weighted median using torch operations"""
    if len(values) == 0:
        return torch.tensor(float('nan'), device=values.device)
    
    sorted_indices = torch.argsort(values)
    sorted_values = values[sorted_indices]
    sorted_weights = weights[sorted_indices]
    
    cumsum = torch.cumsum(sorted_weights, dim=0)
    total_weight = cumsum[-1]
    
    median_pos = total_weight / 2
    median_idx = torch.searchsorted(cumsum, median_pos)
    
    # Clamp to valid range
    median_idx = torch.clamp(median_idx, 0, len(sorted_values) - 1)
    return sorted_values[median_idx]

# %%
#midi_pair = random.choice(trombone_champ_pairs)

# %%
#stem = midi_pair.load_audio().resample(22050)
#midi = midi_pair.load_midi()
#pitches, new_midi = fix_midi_octave_errors(midi_pair,λ=500)
#plot_f0_comparison(pitches, midi_pair.load_audio(), new_midi)

# %%
#stem = midi_pair.load_audio().resample(22050)
#midi = midi_pair.load_midi()
#pitches, new_midi = fix_midi_octave_errors(midi_pair, pitch_algo="pesto",λ=350)
#new_midi.make_stereo_comparison(stem).play()
#plot_f0_comparison(pitches, midi_pair.load_audio(), new_midi)



# %%
pitch_algo = "crepe"
λ = 500

output_dir = f"/app2/suno/data/sara/trombone_champ_vocals_octaved_{pitch_algo}_{λ}"
os.makedirs(output_dir, exist_ok=True)

print(f"Writing to {output_dir}")
for midi_pair in tqdm(trombone_champ_pairs):
    filename = midi_pair.audio_file.split("/")[-1].split(".")[0]
    try:
        pitches, new_midi = fix_midi_octave_errors(midi_pair, pitch_algo=pitch_algo, λ = λ)
    except Exception as e:
        print(f"Skipping {filename}: {e}")
        continue
    midi_path = os.path.join(output_dir, filename + ".mid")
    audio_path = os.path.join(output_dir, filename + ".opus")
    new_midi.write(midi_path)
    shutil.copy(midi_pair.audio_file, audio_path)
    try:
        new_midi.write(midi_path)
        shutil.copy(midi_pair.audio_file, audio_path)
    except Exception as e:
        print(f"Failed to write {midi_path}: {e}")