"""
FlowEdit Final - 完全按照peaudio的正确方式实现

关键设计（对照peaudio/flowedit_semantic_eval_v2_uncond_src.py）：
- v_src: force_uncond=True（真正的无条件速度）
- v_tar: 目标文本 + CFG
- n_min=0: 纯FlowEdit，不使用最后几步trick
- 只调 tar_cfg_scale

用法:
    python flowedit_final.py --tar_cfg_scale 3.5 --device cuda:0
    python flowedit_final.py --tar_cfg_scale 5.0 --device cuda:1
"""

import argparse
import os
import sys
import json
from pathlib import Path
from datetime import datetime
from typing import Optional, Tuple

import torch
import torch.nn.functional as F
import torchaudio
import numpy as np
from tqdm import tqdm

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from dit.dit import DiT
from dit.config import TransformerConfig
from flow_matching.flow_matching import FlowMatchingWrapper
from flow_matching.vae_flow_matching import VAEFlowMatchingWrapper
from generator.vaegen import VAEGenerator, VAEGeneratorConfig, AudioEmbeddingProjector


# ============================================================================
# FlowEdit 核心 - 完全匹配peaudio的uncond版本
# ============================================================================

class FlowEditUncond:
    """
    v_src = model(zt_src, t, force_uncond=True)   # 无条件速度
    v_tar = model(zt_tar, t, tar_cond, cfg)       # 目标条件速度+CFG
    zt_edit += dt * (v_tar - v_src)

    n_min=0: 纯FlowEdit，不切换到标准采样
    """

    def __init__(self, semantic_model, audio_projector, vae_gen, device='cuda'):
        self.semantic_model = semantic_model
        self.audio_projector = audio_projector
        self.vae_gen = vae_gen
        self.device = device
        self.semantic_model.eval()
        self.audio_projector.eval()
        self.vae_gen.eval()

    def get_velocity_uncond(self, x, t, global_cond, cross_cond, memory_mask=None):
        """无条件速度 - force_uncond=True"""
        B = x.shape[0]
        if t.dim() == 0:
            t = t.expand(B)
        v = self.semantic_model.model(
            x=x, time=t, global_cond=global_cond, cross_cond=cross_cond,
            memory_padding_mask=memory_mask, force_uncond=True,
        )
        return v

    def get_velocity_cond(self, x, t, global_cond, cross_cond, memory_mask=None, cfg_scale=3.5):
        """有条件速度 + CFG"""
        B = x.shape[0]
        if t.dim() == 0:
            t = t.expand(B)

        if cfg_scale == 1.0:
            v = self.semantic_model.model(
                x=x, time=t, global_cond=global_cond, cross_cond=cross_cond,
                memory_padding_mask=memory_mask, force_uncond=False,
            )
        else:
            v_cond = self.semantic_model.model(
                x=x, time=t, global_cond=global_cond, cross_cond=cross_cond,
                memory_padding_mask=memory_mask, force_uncond=False,
            )
            v_uncond = self.semantic_model.model(
                x=x, time=t, global_cond=global_cond, cross_cond=cross_cond,
                memory_padding_mask=memory_mask, force_uncond=True,
            )
            v = v_uncond + cfg_scale * (v_cond - v_uncond)
        return v

    @torch.no_grad()
    def edit(self, x_src, tar_global_cond, tar_cross_cond, tar_memory_mask=None,
             num_steps=50, n_avg=1, tar_cfg_scale=3.5):
        """
        纯FlowEdit编辑（n_min=0）

        v_delta = v_tar(zt_tar, tar_cond, cfg) - v_src(zt_src, uncond)
        """
        device = x_src.device
        dtype = x_src.dtype

        timesteps = torch.linspace(1.0, 0.0, num_steps + 1, device=device, dtype=dtype)
        zt_edit = x_src.clone()

        for i in range(num_steps):
            t_i = timesteps[i]
            dt = timesteps[i + 1] - t_i

            v_delta_avg = torch.zeros_like(zt_edit)

            for _ in range(n_avg):
                noise = torch.randn_like(x_src)
                zt_src = (1 - t_i) * x_src + t_i * noise
                zt_tar = zt_edit + (zt_src - x_src)

                # 源速度：force_uncond=True（peaudio的做法）
                v_src = self.get_velocity_uncond(
                    x=zt_src, t=t_i,
                    global_cond=tar_global_cond,  # 会被force_uncond忽略
                    cross_cond=tar_cross_cond,
                    memory_mask=tar_memory_mask,
                )

                # 目标速度：使用目标文本 + CFG
                v_tar = self.get_velocity_cond(
                    x=zt_tar, t=t_i,
                    global_cond=tar_global_cond,
                    cross_cond=tar_cross_cond,
                    memory_mask=tar_memory_mask,
                    cfg_scale=tar_cfg_scale,
                )

                v_delta_avg += (v_tar - v_src)

            v_delta_avg = v_delta_avg / n_avg
            zt_edit = zt_edit + dt * v_delta_avg

        return zt_edit

    @torch.no_grad()
    def edit_and_decode(self, x_src, tar_global_cond, tar_cross_cond, tar_memory_mask=None,
                        num_steps=50, n_avg=1, tar_cfg_scale=3.5, vae_num_steps=25):
        """FlowEdit + VAE decode to DAC latent"""
        edited_semantic = self.edit(
            x_src=x_src,
            tar_global_cond=tar_global_cond,
            tar_cross_cond=tar_cross_cond,
            tar_memory_mask=tar_memory_mask,
            num_steps=num_steps,
            n_avg=n_avg,
            tar_cfg_scale=tar_cfg_scale,
        )

        # VAE decode: semantic → DAC latent
        pe_audio_cond = edited_semantic  # already projected
        B = pe_audio_cond.shape[0]
        T_seq, dac_dim = 250, 128
        x = torch.randn(B, T_seq, dac_dim, device=self.device, dtype=pe_audio_cond.dtype)
        dt_val = -1.0 / vae_num_steps

        for step_i in range(vae_num_steps):
            t_val = 1.0 - step_i / vae_num_steps
            t = torch.full((B,), t_val, device=self.device, dtype=x.dtype)
            v = self.vae_gen.dit(x=x, time=t, pe_audio_cond=pe_audio_cond)
            x = x + v * dt_val

        pred_dac = x.transpose(1, 2).float()  # (B, 128, 250)
        return pred_dac


# ============================================================================
# 模型加载
# ============================================================================

PE_MODULE_PATH = "/data/peaudio/perception_models"


def load_models(semantic_ckpt, vae_ckpt, out_dim=128, device='cuda'):
    """加载所有模型"""
    print(f"Loading models (dim={out_dim})...")

    # PE audio encoder
    sys.path.insert(0, PE_MODULE_PATH)
    from core.audio_visual_encoder import PEAudioFrame, PEAudioFrameTransform
    pe_model = PEAudioFrame.from_config("pe-a-frame-large", pretrained=True).to(device).eval()
    pe_transform = PEAudioFrameTransform.from_config("pe-a-frame-large")

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

    # Flan-T5
    from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
    flan_model = AutoModelForSeq2SeqLM.from_pretrained("google/flan-t5-large").to(device).eval()
    flan_tokenizer = AutoTokenizer.from_pretrained("google/flan-t5-large")

    # Semantic model
    sem_ckpt = torch.load(semantic_ckpt, map_location='cpu', weights_only=False)
    sem_hp = sem_ckpt.get('hyper_parameters', {})

    config = TransformerConfig(
        dim=sem_hp.get('dim', 1152), n_layers=sem_hp.get('n_layers', 28),
        n_heads=sem_hp.get('n_heads', 16),
        in_channels=out_dim, out_channels=out_dim,
        global_cond_dim=sem_hp.get('global_cond_dim', 1024),
        cross_cond_dim=sem_hp.get('cross_cond_dim', 1024),
        use_global_cond=True, use_cross_cond=True,
        max_positions=sem_hp.get('max_positions', 1024),
        cfg_dropout=0.0,
    )
    dit_model = DiT(config)
    semantic_model = FlowMatchingWrapper(
        model=dit_model, inference_mode='euler', num_steps=50, reverse_flow=True,
    )
    state = {k[len('model.'):]: v for k, v in sem_ckpt['state_dict'].items() if k.startswith('model.')}
    semantic_model.load_state_dict(state, strict=True)
    semantic_model = semantic_model.to(device).eval()

    # VAE model
    from infer_vae import load_vae_model
    vae_wrapper, vae_generator = load_vae_model(vae_ckpt, out_dim, device)

    # Build editor
    editor = FlowEditUncond(
        semantic_model=semantic_model,
        audio_projector=vae_generator.audio_projector,
        vae_gen=vae_generator,
        device=device,
    )

    return editor, pe_model, pe_transform, dac_model, flan_model, flan_tokenizer


def extract_text_features(caption, pe_model, pe_transform, flan_model, flan_tokenizer, device):
    """提取文本特征"""
    pe_inputs = pe_transform(text=[caption])
    pe_inputs = pe_inputs.to(device)
    with torch.no_grad():
        text_out = pe_model._get_text_output(pe_inputs['input_ids'], pe_inputs['attention_mask'])
        pe_text_embeds = pe_model.text_head(text_out.pooler_output)

    flan_inputs = flan_tokenizer(caption, return_tensors="pt", padding=True, truncation=True).to(device)
    with torch.no_grad():
        flan_out = flan_model.encoder(**flan_inputs)
    flan_feature = flan_out.last_hidden_state
    flan_mask = flan_inputs['attention_mask']

    return pe_text_embeds, flan_feature, flan_mask


def encode_audio_semantic(wav_path, pe_model, pe_transform, editor, device, target_sr=48000):
    """音频 → PE semantic → projected semantic"""
    wav, sr = torchaudio.load(str(wav_path))
    if sr != target_sr:
        wav = torchaudio.functional.resample(wav, sr, target_sr)
    if wav.shape[0] > 1:
        wav = wav.mean(dim=0, keepdim=True)
    # pad/trim to 10s
    target_len = target_sr * 10
    if wav.shape[1] > target_len:
        wav = wav[:, :target_len]
    elif wav.shape[1] < target_len:
        wav = F.pad(wav, (0, target_len - wav.shape[1]))

    wav = wav.to(device)

    # PE encode (use no_grad, NOT inference_mode, so tensors can flow to projector)
    pe_inputs = pe_transform(audio=[wav], sampling_rate=target_sr)
    pe_inputs = {k: v.to(device) if torch.is_tensor(v) else v for k, v in pe_inputs.items()}
    with torch.no_grad():
        audio_output = pe_model.audio_model(pe_inputs['input_values'])
        pe_audio_embeds = pe_model.audio_head(audio_output.last_hidden_state)

    # Trim/pad to 250 frames
    T = pe_audio_embeds.shape[1]
    if T > 250:
        pe_audio_embeds = pe_audio_embeds[:, :250, :]
    elif T < 250:
        pe_audio_embeds = F.pad(pe_audio_embeds, (0, 0, 0, 250 - T))

    # Project through audio_projector (no KL)
    projected = editor.audio_projector(pe_audio_embeds, return_kl=False)
    return projected, wav


# ============================================================================
# Main
# ============================================================================

def main():
    parser = argparse.ArgumentParser(description='FlowEdit Final (peaudio-style uncond)')
    parser.add_argument('--semantic_ckpt', type=str,
                        default='./logs/semantic_128dim_large_nokl_vae/checkpoints/step=500000.ckpt')
    parser.add_argument('--vae_ckpt', type=str,
                        default='./logs/vae_128dim_large_nokl_8gpu/checkpoints/step=500000.ckpt')
    parser.add_argument('--out_dim', type=int, default=128)
    parser.add_argument('--edits_json', type=str, required=True,
                        help='JSON file with edit pairs')
    parser.add_argument('--audio_dir', type=str,
                        default='/data/data/test_eval/test')
    parser.add_argument('--tar_cfg_scale', type=float, default=3.5,
                        help='Target CFG scale (main knob)')
    parser.add_argument('--num_steps', type=int, default=50)
    parser.add_argument('--vae_num_steps', type=int, default=25)
    parser.add_argument('--output_dir', type=str, default='./flowedit_eval_results/final')
    parser.add_argument('--device', type=str, default='cuda:0')
    parser.add_argument('--seed', type=int, default=42)
    args = parser.parse_args()

    torch.manual_seed(args.seed)
    np.random.seed(args.seed)
    device = args.device

    # Create output dir
    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
    output_dir = Path(args.output_dir) / f"cfg{args.tar_cfg_scale}_{timestamp}"
    output_dir.mkdir(parents=True, exist_ok=True)
    pred_dir = output_dir / "pred_audio"
    orig_dir = output_dir / "original_audio"
    pred_dir.mkdir(exist_ok=True)
    orig_dir.mkdir(exist_ok=True)

    # Save config
    config = vars(args)
    config['n_min'] = 0
    config['method'] = 'uncond_src (force_uncond=True for v_src, matching peaudio)'
    with open(output_dir / "run_config.json", 'w') as f:
        json.dump(config, f, indent=2)

    # Load models
    editor, pe_model, pe_transform, dac_model, flan_model, flan_tokenizer = \
        load_models(args.semantic_ckpt, args.vae_ckpt, args.out_dim, device)

    # Load edits
    with open(args.edits_json) as f:
        edits = json.load(f)
    print(f"Loaded {len(edits)} edit pairs")
    print(f"Config: tar_cfg={args.tar_cfg_scale}, n_min=0, steps={args.num_steps}, vae_steps={args.vae_num_steps}")

    results = []
    sr = dac_model.sample_rate

    for idx, edit in enumerate(tqdm(edits, desc="Editing")):
        audio_id = edit['audio_id']
        src_caption = edit.get('original_caption', edit.get('src_caption', ''))
        tar_caption = edit.get('target_caption', edit.get('tar_caption', ''))
        edit_type = edit.get('edit_type', 'unknown')

        # Find audio file
        wav_path = Path(args.audio_dir) / f"{audio_id}.wav"
        if not wav_path.exists():
            wav_path = Path(args.audio_dir) / f"Y{audio_id}.wav"
        if not wav_path.exists():
            continue

        sample_name = f"{audio_id}_edit{idx:03d}"

        try:
            # Encode source audio → semantic
            x_src, wav_tensor = encode_audio_semantic(wav_path, pe_model, pe_transform, editor, device)

            # Extract target text features
            tar_global, tar_cross, tar_mask = extract_text_features(
                tar_caption, pe_model, pe_transform, flan_model, flan_tokenizer, device
            )

            # FlowEdit (uncond source, n_min=0)
            pred_dac = editor.edit_and_decode(
                x_src=x_src,
                tar_global_cond=tar_global,
                tar_cross_cond=tar_cross,
                tar_memory_mask=tar_mask,
                num_steps=args.num_steps,
                n_avg=1,
                tar_cfg_scale=args.tar_cfg_scale,
                vae_num_steps=args.vae_num_steps,
            )

            # DAC decode
            pred_audio = dac_model.decode(pred_dac).squeeze(0).cpu()

            # Save
            pred_audio_norm = pred_audio / (pred_audio.abs().max() + 1e-8)
            torchaudio.save(str(pred_dir / f"{sample_name}.wav"), pred_audio_norm, sr)

            # Save original
            orig_wav = wav_tensor.cpu()
            orig_wav = orig_wav / (orig_wav.abs().max() + 1e-8)
            torchaudio.save(str(orig_dir / f"{audio_id}.wav"), orig_wav, sr)

            results.append({
                'sample_name': sample_name,
                'audio_id': audio_id,
                'edit_type': edit_type,
                'src_caption': src_caption,
                'tar_caption': tar_caption,
                'pred_path': str(pred_dir / f"{sample_name}.wav"),
                'original_path': str(orig_dir / f"{audio_id}.wav"),
            })

        except Exception as e:
            torch.cuda.empty_cache()
            print(f"  Error on {audio_id}: {e}")
            continue

    # Save results
    with open(output_dir / "results.json", 'w') as f:
        json.dump(results, f, indent=2)

    print(f"\nDone! {len(results)} edits saved to {output_dir}")
    print(f"  pred_audio: {pred_dir}")
    print(f"  original_audio: {orig_dir}")


if __name__ == '__main__':
    main()
