"""
TTABench Evaluation for semaudio semantic_64dim_large model

使用 Resonate repo 的 TTABench (2999 prompts) 评估 64dim_large 模型:
1. 读取 6 个 prompt JSON (acc/generalization/robustness/fairness/bias/toxicity)
2. 加载 semaudio 推理 pipeline (PE + Flan-T5 + Semantic DiT + VAE DiT + DAC)
3. 批量生成 2999 个音频 → prompt_XXXX.wav (支持多卡并行)
4. 计算 AES (CE/CU/PC/PQ) 和 CLAP scores

用法:
    # 单卡
    python eval_ttabench.py --device cuda:6
    # 双卡并行
    python eval_ttabench.py --devices cuda:6,cuda:7
    # 指定 checkpoint + 输出目录名
    python eval_ttabench.py --devices cuda:6,cuda:7 --semantic_ckpt .../step=195000.ckpt --output_name step195000
    # 仅评测
    python eval_ttabench.py --skip_generation --output_name step195000
"""

import argparse
import os
import sys
import json
import logging
from pathlib import Path
from typing import List, Dict
from concurrent.futures import ThreadPoolExecutor
import threading

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")
RESONATE_DIR = BASE_DIR / "Resonate"
TTABENCH_DIR = RESONATE_DIR / "ttabench"
LOG_BASE = SEMAUDIO_DIR / "logs"

# 默认 checkpoint
DEFAULT_SEMANTIC_CKPT = str(LOG_BASE / "semantic_64dim_large" / "checkpoints" / "last.ckpt")
VAE_CKPT = str(LOG_BASE / "vae_64dim_kl" / "checkpoints" / "last.ckpt")
OUT_DIM = 64

# TTABench prompt files
PROMPT_FILES = [
    "acc_prompt.json",
    "generalization_prompt.json",
    "robustness_prompt.json",
    "fairness_prompt.json",
    "bias_prompt.json",
    "toxicity_prompt.json",
]


def parse_args():
    parser = argparse.ArgumentParser(description='TTABench Evaluation')
    parser.add_argument('--device', type=str, default=None,
                        help='单卡设备 (e.g. cuda:6)')
    parser.add_argument('--devices', type=str, default=None,
                        help='多卡设备，逗号分隔 (e.g. cuda:6,cuda:7)')
    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('--semantic_ckpt', type=str, default=None,
                        help='Semantic DiT checkpoint 路径')
    parser.add_argument('--output_name', type=str, default=None,
                        help='输出子目录名 (默认根据 ckpt 名自动生成)')
    parser.add_argument('--skip_generation', action='store_true',
                        help='跳过生成，仅运行评测')
    parser.add_argument('--skip_eval', action='store_true',
                        help='跳过评测，仅生成音频')
    parser.add_argument('--shard', type=str, default=None,
                        help='分片: "0/2" 表示共2片取第0片 (仅生成，不评测)')
    parser.add_argument('--subset', type=str, default=None,
                        help='仅使用指定子集 (acc/generalization/robustness/fairness/bias/toxicity)')
    parser.add_argument('--text_features', type=str, default=None,
                        help='预提取的文本特征 .pt 文件路径 (跳过加载 PE + Flan-T5)')
    parser.add_argument('--vae_ckpt', type=str, default=None,
                        help='VAE checkpoint 路径 (默认: vae_64dim_kl/last.ckpt)')
    parser.add_argument('--seed', type=int, default=None,
                        help='随机种子，固定后可复现生成结果')
    parser.add_argument('--out_dim', type=int, default=64,
                        help='semantic/VAE 输出维度 (64 或 128)')
    args = parser.parse_args()

    # 处理设备参数
    if args.devices:
        args.device_list = [d.strip() for d in args.devices.split(',')]
    elif args.device:
        args.device_list = [args.device]
    else:
        args.device_list = ['cuda:6']

    # 处理 checkpoint
    if args.semantic_ckpt is None:
        args.semantic_ckpt = DEFAULT_SEMANTIC_CKPT

    # 处理输出目录名
    if args.output_name is None:
        ckpt_name = Path(args.semantic_ckpt).stem  # e.g., "step=195000" or "last"
        args.output_name = f"semantic_64dim_large_{ckpt_name}"

    # 处理分片
    if args.shard:
        parts = args.shard.split('/')
        args.shard_idx = int(parts[0])
        args.shard_total = int(parts[1])
    else:
        args.shard_idx = None
        args.shard_total = None

    return args


def load_all_prompts(subset=None) -> List[Dict]:
    """加载 TTABench prompts, 可选子集"""
    SUBSET_MAP = {
        "acc": ["acc_prompt.json"],
        "generalization": ["generalization_prompt.json"],
        "robustness": ["robustness_prompt.json"],
        "fairness": ["fairness_prompt.json"],
        "bias": ["bias_prompt.json"],
        "toxicity": ["toxicity_prompt.json"],
    }
    if subset:
        files = SUBSET_MAP.get(subset, PROMPT_FILES)
    else:
        files = PROMPT_FILES

    all_prompts = []
    prompts_dir = TTABENCH_DIR / "prompts"
    for fname in files:
        fpath = prompts_dir / fname
        with open(fpath, 'r', encoding='utf-8') as f:
            data = json.load(f)
        print(f"  {fname}: {len(data)} prompts")
        all_prompts.extend(data)
    all_prompts.sort(key=lambda x: int(x['id'].replace('prompt_', '')))
    return all_prompts


def load_precomputed_features(text_features_path: str, prompt_ids: List[str], device: str):
    """从预提取的 .pt 文件加载指定 prompt 的文本特征"""
    data = torch.load(text_features_path, map_location='cpu')
    all_ids = data['prompt_ids']
    id_to_idx = {pid: i for i, pid in enumerate(all_ids)}

    indices = [id_to_idx[pid] for pid in prompt_ids]
    indices_t = torch.tensor(indices, dtype=torch.long)

    pe = data['pe_text_embeds'][indices_t].to(device)
    flan_feat = data['flan_text_feature'][indices_t].to(device)
    flan_mask = data['flan_text_mask'][indices_t].to(device)
    return pe, flan_feat, flan_mask


@torch.inference_mode()
def generate_on_device(
    device: str,
    prompts: List[Dict],
    semantic_ckpt: str,
    pred_audio_dir: Path,
    batch_size: int = 16,
    cfg_scale: float = 3.0,
    num_steps_semantic: int = 50,
    num_steps_vae: int = 50,
    progress_lock=None,
    shared_counter=None,
    text_features_path: str = None,
    vae_ckpt: str = None,
    out_dim: int = 64,
    seed: int = None,
):
    """单卡生成: 加载模型 + 批量推理

    text_features_path: 预提取的文本特征 .pt 路径，有则跳过 PE + Flan-T5 加载
    """
    # 固定随机种子（如有）
    if seed is not None:
        import random
        import numpy as np
        torch.manual_seed(seed)
        torch.cuda.manual_seed_all(seed)
        random.seed(seed)
        np.random.seed(seed)
        torch.backends.cudnn.deterministic = True
        torch.backends.cudnn.benchmark = False
        print(f"[{device}] 随机种子已固定: {seed}")

    print(f"[{device}] 加载模型...")

    # 预提取模式: 一次性加载所有文本特征，不加载文本编码器
    precomputed_pe = None
    precomputed_flan = None
    precomputed_mask = None
    pe_model = pe_transform = flan_model = flan_tokenizer = None

    if text_features_path:
        print(f"[{device}] 加载预提取文本特征: {text_features_path}")
        prompt_ids = [p['id'] for p in prompts]
        precomputed_pe, precomputed_flan, precomputed_mask = load_precomputed_features(
            text_features_path, prompt_ids, device
        )
        print(f"[{device}] 文本特征加载完成 ({precomputed_pe.shape[0]} 条)")
    else:
        pe_model, pe_transform, flan_model, flan_tokenizer = load_text_encoders(device)

    # 加载 DAC
    from dacvae import DACVAE
    dac_model = DACVAE.load("facebook/dacvae-watermarked").to(device).eval()
    sample_rate = dac_model.sample_rate

    # 加载 Semantic DiT
    _out_dim = out_dim
    _vae_ckpt = vae_ckpt or VAE_CKPT
    sem_model = load_semantic_dit(semantic_ckpt, _out_dim, device)

    # 加载 VAE DiT
    vae_model = load_vae_model(_vae_ckpt, _out_dim, device)

    print(f"[{device}] 模型加载完成, 开始生成 {len(prompts)} 条...")

    generated = 0
    prompt_id_to_local_idx = {p['id']: i for i, p in enumerate(prompts)}

    for batch_start in range(0, len(prompts), batch_size):
        batch_end = min(batch_start + batch_size, len(prompts))
        batch_prompts = prompts[batch_start:batch_end]

        # 跳过已生成
        all_exist = all(
            (pred_audio_dir / f"{p['id']}.wav").exists()
            for p in batch_prompts
        )
        if all_exist:
            generated += len(batch_prompts)
            if shared_counter is not None:
                with progress_lock:
                    shared_counter[0] += len(batch_prompts)
            continue

        # 文本特征
        if precomputed_pe is not None:
            # 从预提取的特征中取对应 batch
            local_indices = [prompt_id_to_local_idx[p['id']] for p in batch_prompts]
            idx_t = torch.tensor(local_indices, dtype=torch.long)
            pe_text_embeds = precomputed_pe[idx_t]
            flan_text_feature = precomputed_flan[idx_t]
            flan_text_mask = precomputed_mask[idx_t]
        else:
            captions = [p['prompt_text'] for p in batch_prompts]
            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,
        )

        # Stage 2: VAE DiT
        gen_dac = sample_vae_from_semantic(
            vae_model, semantic_z,
            num_steps=num_steps_vae,
        )

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

        for i, p in enumerate(batch_prompts):
            wav_path = pred_audio_dir / f"{p['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 += len(batch_prompts)
        if shared_counter is not None:
            with progress_lock:
                shared_counter[0] += len(batch_prompts)
                total_done = shared_counter[0]
                total_all = shared_counter[1]
                print(f"  进度: {total_done}/{total_all} ({100*total_done/total_all:.1f}%)")

    print(f"[{device}] 完成: 生成 {generated} 条")

    # 释放显存
    del sem_model, vae_model, dac_model
    if pe_model is not None:
        del pe_model, flan_model
    if precomputed_pe is not None:
        del precomputed_pe, precomputed_flan, precomputed_mask
    torch.cuda.empty_cache()


def generate_audio_multi_gpu(
    prompts: List[Dict],
    device_list: List[str],
    semantic_ckpt: str,
    output_dir: Path,
    batch_size: int = 16,
    cfg_scale: float = 3.0,
    num_steps_semantic: int = 50,
    num_steps_vae: int = 50,
):
    """多卡并行生成"""
    pred_audio_dir = output_dir / "pred_audio"
    pred_audio_dir.mkdir(parents=True, exist_ok=True)

    num_gpus = len(device_list)
    total = len(prompts)

    if num_gpus == 1:
        # 单卡直接跑
        print(f"\n🖥️ 单卡生成: {device_list[0]}")
        generate_on_device(
            device=device_list[0],
            prompts=prompts,
            semantic_ckpt=semantic_ckpt,
            pred_audio_dir=pred_audio_dir,
            batch_size=batch_size,
            cfg_scale=cfg_scale,
            num_steps_semantic=num_steps_semantic,
            num_steps_vae=num_steps_vae,
        )
    else:
        # 多卡: 按 prompt 均分
        chunk_size = (total + num_gpus - 1) // num_gpus
        prompt_chunks = []
        for i in range(num_gpus):
            start = i * chunk_size
            end = min(start + chunk_size, total)
            prompt_chunks.append(prompts[start:end])

        print(f"\n🖥️ {num_gpus} 卡并行生成:")
        for i, (dev, chunk) in enumerate(zip(device_list, prompt_chunks)):
            print(f"  {dev}: {len(chunk)} prompts")

        # 共享进度
        progress_lock = threading.Lock()
        shared_counter = [0, total]  # [done, total]

        # 多线程启动
        with ThreadPoolExecutor(max_workers=num_gpus) as executor:
            futures = []
            for dev, chunk in zip(device_list, prompt_chunks):
                f = executor.submit(
                    generate_on_device,
                    device=dev,
                    prompts=chunk,
                    semantic_ckpt=semantic_ckpt,
                    pred_audio_dir=pred_audio_dir,
                    batch_size=batch_size,
                    cfg_scale=cfg_scale,
                    num_steps_semantic=num_steps_semantic,
                    num_steps_vae=num_steps_vae,
                    progress_lock=progress_lock,
                    shared_counter=shared_counter,
                )
                futures.append(f)

            # 等待所有完成
            for f in futures:
                f.result()

    wav_count = len(list(pred_audio_dir.glob("*.wav")))
    print(f"\n✅ 生成完成: {wav_count} 个音频 → {pred_audio_dir}")
    return pred_audio_dir


def prepare_jsonl(pred_audio_dir: Path, output_dir: Path) -> str:
    """准备 JSONL 文件 (audio-aes 和 CLAP 评测输入)"""
    jsonl_path = output_dir / "all_samples.jsonl"
    wav_files = sorted(pred_audio_dir.glob("*.wav"))
    with open(jsonl_path, 'w') as f:
        for wav in wav_files:
            json.dump({"path": str(wav.resolve())}, f)
            f.write('\n')
    print(f"📝 JSONL: {len(wav_files)} files → {jsonl_path}")
    return str(jsonl_path)


def run_aes(jsonl_path: str, output_dir: Path):
    """运行 AES 评测 (CE/CU/PC/PQ)"""
    import subprocess

    aes_result_jsonl = str(output_dir / "aes_scores.jsonl")
    aes_summary = str(output_dir / "aes_summary.txt")

    # Step 1: audio-aes CLI
    print("📊 计算 AES scores...")
    cmd = f"audio-aes {jsonl_path} --batch-size 4 > {aes_result_jsonl}"
    print(f"  运行: {cmd}")
    ret = os.system(cmd)
    if ret != 0:
        print(f"⚠️ audio-aes 失败 (returncode={ret})")
        return {}

    # Step 2: 汇总
    total_ce = total_cu = total_pc = total_pq = count = 0
    with open(aes_result_jsonl, 'r') as f:
        for line in f:
            line = line.strip()
            if not line.startswith("{"):
                continue
            data = json.loads(line)
            total_ce += data.get('CE', 0)
            total_cu += data.get('CU', 0)
            total_pc += data.get('PC', 0)
            total_pq += data.get('PQ', 0)
            count += 1

    if count > 0:
        metrics = {
            "AES_CE": total_ce / count,
            "AES_CU": total_cu / count,
            "AES_PC": total_pc / count,
            "AES_PQ": total_pq / count,
            "AES_count": count,
        }
    else:
        metrics = {}

    # 保存汇总
    with open(aes_summary, 'w') as f:
        for k, v in metrics.items():
            f.write(f"{k}: {v}\n")

    print(f"  AES 结果: CE={metrics.get('AES_CE', 0):.4f}, CU={metrics.get('AES_CU', 0):.4f}, "
          f"PC={metrics.get('AES_PC', 0):.4f}, PQ={metrics.get('AES_PQ', 0):.4f} (n={count})")
    return metrics


def run_clap(pred_audio_dir: Path, prompts: List[Dict], output_dir: Path):
    """运行 CLAP 评测"""
    from msclap import CLAP

    clap_result_jsonl = str(output_dir / "clap_scores.jsonl")
    clap_summary = str(output_dir / "clap_summary.txt")

    # Build prompt id → text mapping
    prompt_map = {p['id']: p['prompt_text'] for p in prompts}

    print("📊 计算 CLAP scores...")
    clap_model = CLAP(version='2023', use_cuda=True)

    results = []
    wav_files = sorted(pred_audio_dir.glob("*.wav"))

    for wav_path in tqdm(wav_files, desc="CLAP scoring"):
        prompt_id = wav_path.stem  # e.g., prompt_0001
        prompt_text = prompt_map.get(prompt_id, "")
        if not prompt_text:
            continue

        audio_emb = clap_model.get_audio_embeddings([str(wav_path)])
        text_emb = clap_model.get_text_embeddings([prompt_text])
        score = torch.nn.functional.cosine_similarity(audio_emb, text_emb).item()
        results.append({
            "prompt_id": prompt_id,
            "prompt_text": prompt_text,
            "clap_score": score,
        })

    # 保存详细结果
    with open(clap_result_jsonl, 'w', encoding='utf-8') as f:
        for r in results:
            json.dump(r, f)
            f.write('\n')

    # 计算均值
    if results:
        avg_clap = sum(r['clap_score'] for r in results) / len(results)
    else:
        avg_clap = 0

    # 按维度统计
    dim_ranges = {
        "acc": (1, 1500),
        "generalization": (1501, 1800),
        "robustness": (1801, 2100),
        "fairness": (2101, 2400),
        "bias": (2401, 2700),
        "toxicity": (2701, 3000),
    }
    dim_scores = {}
    for dim_name, (lo, hi) in dim_ranges.items():
        scores = [r['clap_score'] for r in results
                  if lo <= int(r['prompt_id'].replace('prompt_', '')) <= hi]
        if scores:
            dim_scores[f"CLAP_{dim_name}"] = sum(scores) / len(scores)

    metrics = {"CLAP_overall": avg_clap, **dim_scores, "CLAP_count": len(results)}

    with open(clap_summary, 'w') as f:
        for k, v in metrics.items():
            f.write(f"{k}: {v}\n")

    print(f"  CLAP 总均值: {avg_clap:.4f} (n={len(results)})")
    for k, v in dim_scores.items():
        print(f"    {k}: {v:.4f}")

    return metrics


def main():
    args = parse_args()

    print("=" * 70)
    print("🎵 TTABench Evaluation (semaudio semantic_64dim_large)")
    print("=" * 70)
    print(f"Devices: {args.device_list}")
    print(f"Batch size: {args.batch_size}")
    print(f"CFG scale: {args.cfg_scale}")
    print(f"Steps: semantic={args.num_steps_semantic}, vae={args.num_steps_vae}")
    print(f"Semantic ckpt: {args.semantic_ckpt}")
    print(f"Output name: {args.output_name}")
    print("=" * 70)

    # 1. 加载所有 prompts
    print("\n📂 加载 TTABench prompts...")
    prompts = load_all_prompts(subset=args.subset)
    print(f"   共 {len(prompts)} 条 prompts")

    # 分片 (连续切块)
    if args.shard_idx is not None:
        chunk_size = (len(prompts) + args.shard_total - 1) // args.shard_total
        start = args.shard_idx * chunk_size
        end = min(start + chunk_size, len(prompts))
        prompts_to_gen = prompts[start:end]
        print(f"   分片 {args.shard_idx}/{args.shard_total}: prompt[{start}:{end}] = {len(prompts_to_gen)} 条")
    else:
        prompts_to_gen = prompts

    output_dir = SEMAUDIO_DIR / "eval_ttabench_output" / args.output_name
    output_dir.mkdir(parents=True, exist_ok=True)

    # 2. 生成音频 (单进程单卡模式)
    if not args.skip_generation:
        device = args.device_list[0]
        pred_audio_dir = output_dir / "pred_audio"
        pred_audio_dir.mkdir(parents=True, exist_ok=True)

        generate_on_device(
            device=device,
            prompts=prompts_to_gen,
            semantic_ckpt=args.semantic_ckpt,
            pred_audio_dir=pred_audio_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,
            text_features_path=args.text_features,
            vae_ckpt=args.vae_ckpt,
            out_dim=args.out_dim,
            seed=args.seed,
        )
    else:
        pred_audio_dir = output_dir / "pred_audio"
        print(f"\n⏭️ 跳过生成，使用已有音频: {pred_audio_dir}")

    # 如果是分片模式，只生成不评测
    if args.shard_idx is not None:
        wav_count = len(list(pred_audio_dir.glob("*.wav")))
        print(f"\n✅ 分片 {args.shard_idx} 生成完成! 当前共 {wav_count}/2999 个音频")
        print("全部分片完成后，用 --skip_generation 运行评测")
        return

    # 3. 评测
    if not args.skip_eval:
        all_metrics = {}

        # 准备 JSONL
        jsonl_path = prepare_jsonl(pred_audio_dir, output_dir)

        # AES
        print(f"\n{'='*70}")
        print("📊 AES Evaluation (CE/CU/PC/PQ)")
        print(f"{'='*70}")
        aes_metrics = run_aes(jsonl_path, output_dir)
        all_metrics.update(aes_metrics)

        # CLAP
        print(f"\n{'='*70}")
        print("📊 CLAP Evaluation")
        print(f"{'='*70}")
        clap_metrics = run_clap(pred_audio_dir, prompts, output_dir)
        all_metrics.update(clap_metrics)

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

        # 保存
        summary_path = output_dir / "ttabench_summary.json"
        with open(summary_path, 'w') as f:
            json.dump(all_metrics, f, indent=4)
        print(f"\n📁 汇总保存到: {summary_path}")

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


if __name__ == '__main__':
    main()
