"""
AudioCaps Test Set Evaluation

在 AudioCaps test set (957条) 上评估 semaudio 的两个 64dim 模型，计算标准音频生成评估指标。

模型:
  - semantic_64dim (小模型 282M): logs/semantic_64dim + logs/vae_64dim_kl
  - semantic_64dim_large (大模型 610M): logs/semantic_64dim_large + logs/vae_64dim_kl

流程:
  1. 读取 test-audiocaps.tsv (957 条 caption)
  2. 加载文本编码器 (PE pe-a-frame-large + Flan-T5-large)，批量提取文本特征
  3. 对每个模型:
     - Stage 1: Semantic DiT 采样 → semantic latent
     - Stage 2: VAE DiT 采样 → DAC latent
     - DAC 解码 → wav
  4. 准备评测目录 + 调用 run_evaluation_flowedit.py 计算 FD/ISC/CLAP

用法:
    python eval_audiocaps.py --device cuda:6
    python eval_audiocaps.py --device cuda:6 --batch_size 8
    python eval_audiocaps.py --device cuda:6 --skip_generation  # 仅评测
"""

import argparse
import os
import sys
import csv
import json
import logging
from pathlib import Path
from typing import List

import torch
import torchaudio
from tqdm import tqdm

# 添加 semaudio 项目路径
SEMAUDIO_DIR = Path("/data/semaudio")
sys.path.insert(0, str(SEMAUDIO_DIR))

from infer_semantic import (
    load_text_encoders,
    load_semantic_dit,
    load_vae_model,
    extract_text_features,
    sample_semantic,
    sample_vae_from_semantic,
    decode_dac,
)


# ============ 路径配置 ============
BASE_DIR = Path("/data")
DATA_DIR = BASE_DIR / "data"
PEAUDIO_DIR = BASE_DIR / "peaudio"
LOG_BASE = SEMAUDIO_DIR / "logs"

TSV_PATH = DATA_DIR / "test-audiocaps.tsv"
GT_WAV_DIR = DATA_DIR / "test_eval" / "test"

# Pipeline 配置: (semantic_ckpt, vae_ckpt, out_dim, name)
PIPELINES = [
    (
        str(LOG_BASE / "semantic_64dim" / "checkpoints" / "last.ckpt"),
        str(LOG_BASE / "vae_64dim_kl" / "checkpoints" / "last.ckpt"),
        64,
        "semantic_64dim",
    ),
    (
        str(LOG_BASE / "semantic_64dim_large" / "checkpoints" / "last.ckpt"),
        str(LOG_BASE / "vae_64dim_kl" / "checkpoints" / "last.ckpt"),
        64,
        "semantic_64dim_large",
    ),
]

OUTPUT_BASE = SEMAUDIO_DIR / "eval_audiocaps_output"


def parse_args():
    parser = argparse.ArgumentParser(description='AudioCaps Test Set Evaluation')
    parser.add_argument('--device', type=str, default='cuda:0')
    parser.add_argument('--batch_size', type=int, default=16)
    parser.add_argument('--cfg_scale', type=float, default=3.0)
    parser.add_argument('--num_steps_semantic', type=int, default=50)
    parser.add_argument('--num_steps_vae', type=int, default=50)
    parser.add_argument('--models', type=str, nargs='+', default=None,
                        help='指定要评估的模型名 (e.g. semantic_64dim semantic_64dim_large)')
    parser.add_argument('--skip_generation', action='store_true',
                        help='跳过生成，仅运行评测（需已有 pred_audio）')
    parser.add_argument('--skip_eval', action='store_true',
                        help='跳过评测，仅生成音频')

    # 通用 pipeline（支持任意 dim）
    parser.add_argument('--semantic_ckpt', type=str, default=None,
                        help='Semantic DiT checkpoint')
    parser.add_argument('--vae_ckpt', type=str, default=None,
                        help='VAE DiT checkpoint')
    parser.add_argument('--out_dim', type=int, default=64,
                        help='out_dim（如 32, 64, 128）')
    parser.add_argument('--pipeline_name', type=str, default=None,
                        help='Pipeline 名称（用作输出子目录名）')
    parser.add_argument('--output_base', type=str, default=None,
                        help='输出根目录（默认 eval_audiocaps_output）')

    # Shard
    parser.add_argument('--shard', type=str, default=None,
                        help='分片: "0/6" 表示共6片取第0片')

    # Seed
    parser.add_argument('--seed', type=int, default=None,
                        help='随机种子，固定后可复现生成结果')

    return parser.parse_args()


def load_tsv(tsv_path: str):
    """读取 test-audiocaps.tsv，返回 [{tsv_id, caption}, ...]"""
    samples = []
    with open(tsv_path, 'r', encoding='utf-8') as f:
        reader = csv.DictReader(f, delimiter='\t')
        for row in reader:
            samples.append({
                'tsv_id': row['id'],
                'caption': row['caption'],
            })
    return samples


@torch.inference_mode()
def generate_audio_batch(
    samples: List[dict],
    sem_model,
    vae_model,
    dac_model,
    pe_model,
    pe_transform,
    flan_model,
    flan_tokenizer,
    out_dim: int,
    output_dir: str,
    batch_size: int = 16,
    cfg_scale: float = 3.0,
    num_steps_semantic: int = 50,
    num_steps_vae: int = 50,
    device: str = 'cuda:0',
):
    """批量推理: text → semantic → vae → dac → wav"""
    pred_audio_dir = Path(output_dir) / "pred_audio"
    pred_audio_dir.mkdir(parents=True, exist_ok=True)

    sample_rate = dac_model.sample_rate
    total = len(samples)
    generated_count = 0

    for batch_start in tqdm(range(0, total, batch_size), desc="生成音频"):
        batch_end = min(batch_start + batch_size, total)
        batch_samples = samples[batch_start:batch_end]

        # 跳过已生成的 batch
        all_exist = all(
            (pred_audio_dir / f"{s['tsv_id']}.wav").exists()
            for s in batch_samples
        )
        if all_exist:
            generated_count += len(batch_samples)
            continue

        # 提取文本特征 (PE pe-a-frame-large + Flan-T5-large)
        captions = [s['caption'] for s in batch_samples]
        pe_text_embeds, flan_text_feature, flan_text_mask = extract_text_features(
            captions, pe_model, pe_transform, flan_model, flan_tokenizer, device
        )

        # Stage 1: Semantic DiT 采样
        semantic_z = sample_semantic(
            sem_model, pe_text_embeds, flan_text_feature, flan_text_mask,
            out_dim=out_dim,
            num_steps=num_steps_semantic,
            cfg_scale=cfg_scale,
        )  # (B, 250, out_dim)

        # Stage 2: VAE DiT 采样
        gen_dac = sample_vae_from_semantic(
            vae_model, semantic_z,
            num_steps=num_steps_vae,
        )  # (B, 250, 128)

        # Stage 3: DAC 解码 + 保存
        gen_audio = decode_dac(gen_dac, dac_model, device)

        for i, s in enumerate(batch_samples):
            wav_path = pred_audio_dir / f"{s['tsv_id']}.wav"
            if wav_path.exists():
                continue
            wav = gen_audio[i].cpu().float()
            wav = wav / (wav.abs().max() + 1e-8)
            torchaudio.save(str(wav_path), wav, sample_rate)

        generated_count += len(batch_samples)

    print(f"✅ 生成完成: {generated_count} 个音频 → {pred_audio_dir}")
    return pred_audio_dir


def prepare_eval_dirs(samples, output_dir: str, gt_wav_dir: str):
    """准备评测目录: reference_audio/ (symlink GT), captions.csv"""
    output_path = Path(output_dir)
    ref_audio_dir = output_path / "reference_audio"
    ref_audio_dir.mkdir(parents=True, exist_ok=True)

    gt_wav_path = Path(gt_wav_dir)
    linked = 0
    for s in samples:
        gt_wav = gt_wav_path / f"{s['tsv_id']}.wav"
        link_path = ref_audio_dir / f"{s['tsv_id']}.wav"
        if not link_path.exists() and gt_wav.exists():
            os.symlink(str(gt_wav.resolve()), str(link_path))
            linked += 1

    print(f"📎 Symlinked {linked} GT wavs → {ref_audio_dir}")

    captions_csv = output_path / "captions.csv"
    with open(captions_csv, 'w', encoding='utf-8', newline='') as f:
        writer = csv.writer(f)
        writer.writerow(['name', 'caption'])
        for s in samples:
            writer.writerow([s['tsv_id'], s['caption']])
    print(f"📝 写入 captions.csv: {len(samples)} 条 → {captions_csv}")


def run_metrics(output_dir: str, audio_length: float = 10.0):
    """调用 peaudio/metrics/run_evaluation_flowedit.py 计算指标（subprocess 方式）"""
    import subprocess
    script = str(PEAUDIO_DIR / "metrics" / "run_evaluation_flowedit.py")
    cmd = [
        sys.executable, script,
        "--samples_dir", output_dir,
        "--audio_length", str(audio_length),
        "--batch_size", "64",
        "--num_workers", "8",
    ]
    print(f"  运行: {' '.join(cmd)}")
    result = subprocess.run(cmd, capture_output=True, text=True, cwd=str(PEAUDIO_DIR))
    print(result.stdout)
    if result.returncode != 0:
        print(f"⚠️ 评测脚本失败 (returncode={result.returncode})")
        print(result.stderr[-2000:] if len(result.stderr) > 2000 else result.stderr)
        return {}

    # 读取结果
    metrics_file = Path(output_dir) / "metrics_flowedit.json"
    if metrics_file.exists():
        with open(metrics_file) as f:
            return json.load(f)
    return {}


def main():
    args = parse_args()

    # Seed
    if args.seed is not None:
        import random
        import numpy as np
        torch.manual_seed(args.seed)
        torch.cuda.manual_seed_all(args.seed)
        random.seed(args.seed)
        np.random.seed(args.seed)
        torch.backends.cudnn.deterministic = True
        torch.backends.cudnn.benchmark = False
        print(f"🎲 随机种子已固定: {args.seed}")

    # Output base
    output_base = Path(args.output_base) if args.output_base else OUTPUT_BASE

    print("=" * 70)
    print("🎵 AudioCaps Test Set Evaluation (semaudio)")
    print("=" * 70)
    print(f"Device: {args.device}")
    print(f"Batch size: {args.batch_size}")
    print(f"CFG scale: {args.cfg_scale}")
    print(f"Semantic steps: {args.num_steps_semantic}, VAE steps: {args.num_steps_vae}")
    if args.seed is not None:
        print(f"Seed: {args.seed}")
    print("=" * 70)

    # 1. 读取 TSV
    print(f"\n📂 读取 TSV: {TSV_PATH}")
    samples = load_tsv(str(TSV_PATH))
    print(f"   共 {len(samples)} 条测试样本")

    # Shard
    shard_idx = shard_total = None
    if args.shard:
        shard_idx, shard_total = [int(x) for x in args.shard.split('/')]
        chunk_size = (len(samples) + shard_total - 1) // shard_total
        start = shard_idx * chunk_size
        end = min(start + chunk_size, len(samples))
        samples_to_gen = samples[start:end]
        print(f"   分片 {shard_idx}/{shard_total}: samples[{start}:{end}] = {len(samples_to_gen)} 条")
    else:
        samples_to_gen = samples

    # 2. 筛选 pipeline
    available_pipelines = []

    # 通用 pipeline（优先使用）
    if args.semantic_ckpt and args.vae_ckpt:
        name = args.pipeline_name or f'sem_{args.out_dim}dim_custom'
        available_pipelines.append((args.semantic_ckpt, args.vae_ckpt, args.out_dim, name))
        print(f"  ✅ {name}: semantic={args.semantic_ckpt}")
    else:
        # 使用默认 PIPELINES
        for sem_ckpt, vae_ckpt, out_dim, name in PIPELINES:
            if args.models and name not in args.models:
                continue
            if os.path.isfile(sem_ckpt) and os.path.isfile(vae_ckpt):
                available_pipelines.append((sem_ckpt, vae_ckpt, out_dim, name))
                print(f"  ✅ {name}: semantic={sem_ckpt}")
            else:
                print(f"  ❌ {name}: checkpoint 不存在，跳过")

    if not available_pipelines:
        print("❌ 没有可用的 pipeline!")
        return

    print(f"\n🔧 将评估 {len(available_pipelines)} 个模型: "
          f"{[p[3] for p in available_pipelines]}")

    # 3. 加载共享模型
    if not args.skip_generation:
        print("\n📂 加载文本编码器 (PE pe-a-frame-large + Flan-T5-large)...")
        pe_model, pe_transform, flan_model, flan_tokenizer = load_text_encoders(args.device)

        print("\n📂 加载 DAC 模型...")
        from dacvae import DACVAE
        dac_model = DACVAE.load("facebook/dacvae-watermarked").to(args.device).eval()
        print(f"✅ DAC 加载完成, sample_rate={dac_model.sample_rate}")

    # 4. 对每个 pipeline 推理 + 评测
    all_metrics = {}

    for sem_ckpt, vae_ckpt, out_dim, name in available_pipelines:
        print(f"\n{'='*70}")
        print(f"🚀 Pipeline: {name} (out_dim={out_dim})")
        print(f"{'='*70}")

        output_dir = str(output_base / name)
        os.makedirs(output_dir, exist_ok=True)

        if not args.skip_generation:
            # 加载 pipeline 模型
            sem_model = load_semantic_dit(sem_ckpt, out_dim, args.device)
            vae_model = load_vae_model(vae_ckpt, out_dim, args.device)

            # 生成音频
            generate_audio_batch(
                samples=samples_to_gen,
                sem_model=sem_model,
                vae_model=vae_model,
                dac_model=dac_model,
                pe_model=pe_model,
                pe_transform=pe_transform,
                flan_model=flan_model,
                flan_tokenizer=flan_tokenizer,
                out_dim=out_dim,
                output_dir=output_dir,
                batch_size=args.batch_size,
                cfg_scale=args.cfg_scale,
                num_steps_semantic=args.num_steps_semantic,
                num_steps_vae=args.num_steps_vae,
                device=args.device,
            )

            # 释放 pipeline 模型显存
            del sem_model, vae_model
            torch.cuda.empty_cache()

        # 分片模式只生成不评测
        if shard_idx is not None:
            pred_audio_dir = Path(output_dir) / "pred_audio"
            wav_count = len(list(pred_audio_dir.glob("*.wav")))
            print(f"\n✅ 分片 {shard_idx} 生成完成! 当前共 {wav_count}/957 个音频")
            continue

        # 准备评测目录
        prepare_eval_dirs(samples, output_dir, str(GT_WAV_DIR))

        # 运行评测
        if not args.skip_eval:
            print(f"\n📊 运行评测: {name}")
            metrics = run_metrics(output_dir)
            all_metrics[name] = metrics

    # 5. 汇总
    if all_metrics:
        print(f"\n{'='*70}")
        print("📊 所有模型评测结果汇总")
        print(f"{'='*70}")
        for name, metrics in all_metrics.items():
            print(f"\n  {name}:")
            for k, v in metrics.items():
                if isinstance(v, float):
                    print(f"    {k:<30}: {v:.6f}")
                else:
                    print(f"    {k:<30}: {v}")

        summary_path = output_base / "summary.json"
        summary_path.parent.mkdir(parents=True, exist_ok=True)
        with open(summary_path, 'w') as f:
            json.dump(all_metrics, f, indent=4)
        print(f"\n📁 汇总保存到: {summary_path}")

    print("\n✅ 评测完成!")


if __name__ == '__main__':
    main()
