from __future__ import annotations

import argparse
import csv
from pathlib import Path

import librosa
import torch
import torch.nn.functional as functional
from transformers import AutoFeatureExtractor, WavLMForXVector

PROJECT_ROOT = Path(__file__).resolve().parent


MODEL_ID = "microsoft/wavlm-base-plus-sv"


def load_audio(path: Path) -> torch.Tensor:
    waveform, _ = librosa.load(path, sr=16000, mono=True)
    return torch.from_numpy(waveform)


def embedding(
    waveform: torch.Tensor,
    extractor: AutoFeatureExtractor,
    model: WavLMForXVector,
    device: torch.device,
) -> torch.Tensor:
    inputs = extractor(
        waveform.numpy(),
        sampling_rate=16000,
        return_tensors="pt",
        padding=True,
    )
    inputs = {name: value.to(device) for name, value in inputs.items()}
    with torch.inference_mode():
        result = model(**inputs).embeddings
    return functional.normalize(result, dim=-1).cpu()


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--root",
        type=Path,
        default=PROJECT_ROOT / "audio",
    )
    parser.add_argument(
        "--output",
        type=Path,
        default=PROJECT_ROOT / "recomputed-similarity.csv",
    )
    parser.add_argument(
        "--reference",
        type=Path,
        required=True,
        help="Reference WAV to compare against; record this choice with your results",
    )
    args = parser.parse_args()

    root = args.root.resolve()
    reference_path = (
        args.reference.resolve()
        if args.reference is not None
        else root / "00_reference" / "00_original_reference.wav"
    )
    if not reference_path.is_file():
        raise FileNotFoundError(reference_path)

    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    extractor = AutoFeatureExtractor.from_pretrained(MODEL_ID)
    model = WavLMForXVector.from_pretrained(MODEL_ID).eval().to(device)
    reference_embedding = embedding(load_audio(reference_path), extractor, model, device)

    rows = []
    for path in sorted(root.glob("0[1-9]_*/*.wav")):
        candidate_embedding = embedding(load_audio(path), extractor, model, device)
        score = float(functional.cosine_similarity(reference_embedding, candidate_embedding).item())
        row = {
            "category": path.parent.name,
            "file": path.name,
            "cosine_similarity": round(score, 6),
            "model_id": MODEL_ID,
        }
        rows.append(row)
        print(f"similarity={score:.4f} file={path}", flush=True)

    output_path = args.output.resolve()
    output_path.parent.mkdir(parents=True, exist_ok=True)
    with output_path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)
    print(f"report={output_path}")


if __name__ == "__main__":
    main()
