"""VAE Generator 训练脚本（Scaled 数据版）

使用 PyTorch Lightning 进行训练，TensorBoard 记录 loss

输入特征（新 bytes+shape parquet 格式）:
- dac_mean + dac_scale → reparameterize → dac_encoded_sampled: (128, 250)
- pe_audio_embeds: (250, 1024) - PE 音频嵌入

AudioEmbeddingProjector 输出 VAE mean/logvar → reparameterize → zsem
总 loss = flow_matching_loss + kl_weight * kl_loss
"""

import argparse
import os
import sys
from datetime import datetime
from typing import Optional

import torch
import torch.nn as nn
import pytorch_lightning as pl
from pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor
from pytorch_lightning.loggers import TensorBoardLogger
from torch.utils.data import DataLoader

# 添加项目路径
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

# 导入项目模块
from generator.vaegen import VAEGenerator, VAEGeneratorConfig
from flow_matching import VAEFlowMatchingWrapper
from data_loaders.scaled_dataset import ScaledFeatureDataset, scaled_collate_fn
from callbacks import AudioSampleCallback


class VAEGeneratorLightningModule(pl.LightningModule):
    """
    VAE Generator 训练模块（含 KL 正则）

    总 loss = flow_matching_loss + kl_weight * kl_loss
    """

    def __init__(
        self,
        # 模型配置
        model_dim: int = 1024,
        n_layers: int = 16,
        n_heads: int = 16,
        pe_audio_in_dim: int = 1024,
        pe_audio_out_dim: int = 64,
        dac_dim: int = 128,
        dropout: float = 0.1,
        # Flow Matching 配置
        inference_mode: str = 'euler',
        num_steps: int = 25,
        # KL 配置
        kl_weight: float = 1e-4,
        no_reparam: bool = False,
        # 训练配置
        learning_rate: float = 1e-4,
        weight_decay: float = 0.01,
        warmup_steps: int = 1000,
        freeze_projector: bool = False,
        **kwargs
    ):
        super().__init__()
        self.save_hyperparameters()

        # 创建 VAEGenerator 配置
        config = VAEGeneratorConfig(
            model_dim=model_dim,
            n_layers=n_layers,
            n_heads=n_heads,
            pe_audio_in_dim=pe_audio_in_dim,
            pe_audio_out_dim=pe_audio_out_dim,
            dac_dim=dac_dim,
            dropout=dropout,
        )

        # 创建 VAEGenerator 模型
        self.vae_generator = VAEGenerator(config)

        # 使用 VAE Flow Matching Wrapper 包装模型
        self.flow_wrapper = VAEFlowMatchingWrapper(
            model=self.vae_generator,
            inference_mode=inference_mode,
            num_steps=num_steps,
        )

        # KL 权重
        self.kl_weight = kl_weight
        self.no_reparam = no_reparam

        # 冻结 AudioEmbeddingProjector
        if freeze_projector:
            for p in self.vae_generator.audio_projector.parameters():
                p.requires_grad = False
            print("=" * 60)
            print("AudioEmbeddingProjector 已冻结（requires_grad=False）")
            proj_params = sum(p.numel() for p in self.vae_generator.audio_projector.parameters())
            print(f"  冻结参数量: {proj_params:,} ({proj_params / 1e6:.2f}M)")
            print("=" * 60)

        # 训练参数
        self.learning_rate = learning_rate
        self.weight_decay = weight_decay
        self.warmup_steps = warmup_steps

    def training_step(self, batch, batch_idx):
        """训练步骤"""
        # dac_encoded_sampled: (batch, 128, 250) -> (batch, 250, 128)
        x1 = batch['dac_encoded_sampled'].transpose(1, 2)
        pe_audio_embeds = batch['pe_audio_embeds']  # (batch, 250, 1024)

        # ---- Step 1: AudioEmbeddingProjector forward ----
        if self.no_reparam:
            # 无 reparameterize: 只用 mean，纯确定性投影
            projector = self.vae_generator.audio_projector
            h = projector.encoder(pe_audio_embeds)
            mean = projector.fc_mean(h)
            pe_audio_cond = projector.norm(mean)
            kl_loss = torch.tensor(0.0, device=x1.device)
        else:
            pe_audio_cond, kl_loss = self.vae_generator.audio_projector(
                pe_audio_embeds, return_kl=True
            )  # pe_audio_cond: (B, 250, out_dim), kl_loss: scalar

        # ---- Step 2: Flow Matching loss ----
        # 手动调用 flow matching（传入已降维的 pe_audio_cond）
        batch_size = x1.shape[0]
        device = x1.device
        dtype = x1.dtype

        t = torch.rand(batch_size, device=device, dtype=dtype)
        x0, xt, t = self.flow_wrapper.flow_matching.sample_xt(x1, t)

        # 通过 VAEDiT（跳过 audio_projector，直接用 pe_audio_cond）
        predicted_v = self.vae_generator.dit(
            x=xt,
            time=t,
            pe_audio_cond=pe_audio_cond,
        )

        fm_loss = self.flow_wrapper.flow_matching.loss(predicted_v, x0, x1, reduction='mean')

        # ---- 总 loss ----
        total_loss = fm_loss + self.kl_weight * kl_loss

        # 记录损失
        self.log('train/loss', total_loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)
        self.log('train/fm_loss', fm_loss, on_step=True, on_epoch=True, prog_bar=False, logger=True)
        self.log('train/kl_loss', kl_loss, on_step=True, on_epoch=True, prog_bar=False, logger=True)

        # 记录学习率
        lr = self.trainer.optimizers[0].param_groups[0]['lr']
        self.log('train/lr', lr, on_step=True, on_epoch=False, prog_bar=False, logger=True)

        return total_loss

    def validation_step(self, batch, batch_idx):
        """验证步骤"""
        x1 = batch['dac_encoded_sampled'].transpose(1, 2)
        pe_audio_embeds = batch['pe_audio_embeds']

        if self.no_reparam:
            projector = self.vae_generator.audio_projector
            h = projector.encoder(pe_audio_embeds)
            mean = projector.fc_mean(h)
            pe_audio_cond = projector.norm(mean)
            kl_loss = torch.tensor(0.0, device=x1.device)
        else:
            pe_audio_cond, kl_loss = self.vae_generator.audio_projector(
                pe_audio_embeds, return_kl=True
            )

        batch_size = x1.shape[0]
        device = x1.device
        dtype = x1.dtype
        t = torch.rand(batch_size, device=device, dtype=dtype)
        x0, xt, t = self.flow_wrapper.flow_matching.sample_xt(x1, t)

        predicted_v = self.vae_generator.dit(
            x=xt,
            time=t,
            pe_audio_cond=pe_audio_cond,
        )

        fm_loss = self.flow_wrapper.flow_matching.loss(predicted_v, x0, x1, reduction='mean')
        total_loss = fm_loss + self.kl_weight * kl_loss

        self.log('val/loss', total_loss, on_step=False, on_epoch=True, prog_bar=True, logger=True, sync_dist=True)
        self.log('val/fm_loss', fm_loss, on_step=False, on_epoch=True, prog_bar=False, logger=True, sync_dist=True)
        self.log('val/kl_loss', kl_loss, on_step=False, on_epoch=True, prog_bar=False, logger=True, sync_dist=True)

        return total_loss

    def configure_optimizers(self):
        """配置优化器和学习率调度器"""
        optimizer = torch.optim.AdamW(
            filter(lambda p: p.requires_grad, self.parameters()),
            lr=self.learning_rate,
            weight_decay=self.weight_decay,
            betas=(0.9, 0.95)
        )

        def lr_lambda(current_step):
            if current_step < self.warmup_steps:
                return float(current_step) / float(max(1, self.warmup_steps))
            else:
                progress = float(current_step - self.warmup_steps) / float(
                    max(1, self.trainer.estimated_stepping_batches - self.warmup_steps)
                )
                return max(0.1, 0.5 * (1.0 + torch.cos(torch.tensor(progress * 3.14159)).item()))

        scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

        return {
            'optimizer': optimizer,
            'lr_scheduler': {
                'scheduler': scheduler,
                'interval': 'step',
                'frequency': 1
            }
        }


class ScaledDataModule(pl.LightningDataModule):
    """Scaled Feature 数据模块"""

    def __init__(
        self,
        train_data_dir: str,
        val_data_dir: Optional[str] = None,
        train_pattern: str = "*_train_*.parquet",
        val_pattern: str = "*_validation_*.parquet",
        train_arrow_cache: Optional[str] = None,
        val_arrow_cache: Optional[str] = None,
        batch_size: int = 32,
        num_workers: int = 4,
        pin_memory: bool = True,
    ):
        super().__init__()
        self.train_data_dir = train_data_dir
        self.val_data_dir = val_data_dir or train_data_dir
        self.train_pattern = train_pattern
        self.val_pattern = val_pattern
        self.train_arrow_cache = train_arrow_cache
        self.val_arrow_cache = val_arrow_cache
        self.batch_size = batch_size
        self.num_workers = num_workers
        self.pin_memory = pin_memory

    def setup(self, stage: Optional[str] = None):
        if stage == 'fit' or stage is None:
            self.train_dataset = ScaledFeatureDataset(
                data_dir=self.train_data_dir,
                pattern=self.train_pattern,
                arrow_cache_path=self.train_arrow_cache,
            )

            try:
                self.val_dataset = ScaledFeatureDataset(
                    data_dir=self.val_data_dir,
                    pattern=self.val_pattern,
                    arrow_cache_path=self.val_arrow_cache,
                )
            except FileNotFoundError:
                print("[ScaledDataModule] 未找到验证集 parquet，从训练集划分 5%")
                total_size = len(self.train_dataset)
                val_size = max(1, int(0.05 * total_size))
                train_size = total_size - val_size
                self.train_dataset, self.val_dataset = torch.utils.data.random_split(
                    self.train_dataset, [train_size, val_size]
                )

    def train_dataloader(self):
        return DataLoader(
            self.train_dataset,
            batch_size=self.batch_size,
            shuffle=True,
            num_workers=self.num_workers,
            pin_memory=self.pin_memory,
            collate_fn=scaled_collate_fn,
            drop_last=True,
            persistent_workers=self.num_workers > 0,
        )

    def val_dataloader(self):
        return DataLoader(
            self.val_dataset,
            batch_size=self.batch_size,
            shuffle=False,
            num_workers=min(2, self.num_workers),
            pin_memory=self.pin_memory,
            collate_fn=scaled_collate_fn,
            drop_last=False,
        )


def parse_args():
    parser = argparse.ArgumentParser(description='VAE Generator 训练（Scaled 数据）')

    # 数据参数
    parser.add_argument('--train_data', type=str,
                        default='/data/data/parquet_features/',
                        help='训练数据 parquet 目录')
    parser.add_argument('--val_data', type=str, default=None,
                        help='验证数据 parquet 目录（默认同 train_data）')
    parser.add_argument('--train_pattern', type=str, default='*_train_*.parquet',
                        help='训练数据 glob 模式')
    parser.add_argument('--val_pattern', type=str, default='*_validation_*.parquet',
                        help='验证数据 glob 模式')
    parser.add_argument('--train_arrow_cache', type=str, default=None,
                        help='训练集 Arrow 缓存路径 (save_to_disk 生成)')
    parser.add_argument('--val_arrow_cache', type=str, default=None,
                        help='验证集 Arrow 缓存路径 (save_to_disk 生成)')
    parser.add_argument('--batch_size', type=int, default=32)
    parser.add_argument('--num_workers', type=int, default=8)

    # 模型参数
    parser.add_argument('--model_dim', type=int, default=1024)
    parser.add_argument('--n_layers', type=int, default=16)
    parser.add_argument('--n_heads', type=int, default=16)
    parser.add_argument('--pe_audio_in_dim', type=int, default=1024)
    parser.add_argument('--pe_audio_out_dim', type=int, default=64,
                        help='AudioEmbeddingProjector 输出维度')
    parser.add_argument('--in_channels', type=int, default=64,
                        help='VAEDiT pe_audio_dim（应与 pe_audio_out_dim 一致）')
    parser.add_argument('--dac_dim', type=int, default=128)
    parser.add_argument('--dropout', type=float, default=0.1)

    # KL 参数
    parser.add_argument('--kl_weight', type=float, default=1e-4,
                        help='KL 损失权重')
    parser.add_argument('--no_reparam', action='store_true',
                        help='关闭 reparameterize，projector 只输出 mean（纯确定性投影）')

    # Flow Matching 参数
    parser.add_argument('--inference_mode', type=str, default='euler')
    parser.add_argument('--num_steps', type=int, default=25)

    # 训练参数
    parser.add_argument('--learning_rate', type=float, default=1e-4)
    parser.add_argument('--weight_decay', type=float, default=0.01)
    parser.add_argument('--warmup_steps', type=int, default=1000)
    parser.add_argument('--max_epochs', type=int, default=1000)
    parser.add_argument('--max_steps', type=int, default=-1)
    parser.add_argument('--gradient_clip_val', type=float, default=1.0)
    parser.add_argument('--accumulate_grad_batches', type=int, default=1)

    # 日志和检查点
    parser.add_argument('--log_dir', type=str, default='./logs')
    parser.add_argument('--exp_name', type=str, default=None)
    parser.add_argument('--save_top_k', type=int, default=3)
    parser.add_argument('--save_every_n_steps', type=int, default=5000,
                        help='每 N 步保存一次 checkpoint')
    parser.add_argument('--audio_sample_every_n_steps', type=int, default=2000,
                        help='每 N 步生成音频样本')
    parser.add_argument('--val_check_interval', type=float, default=1.0)
    parser.add_argument('--log_every_n_steps', type=int, default=50)

    # 硬件
    parser.add_argument('--gpu_ids', type=str, default='0')
    parser.add_argument('--precision', type=str, default='bf16-mixed',
                        choices=['32', '16-mixed', 'bf16-mixed'])
    parser.add_argument('--strategy', type=str, default='auto')

    # 其他
    parser.add_argument('--seed', type=int, default=42)
    parser.add_argument('--resume_from', type=str, default=None)
    parser.add_argument('--freeze_projector', action='store_true',
                        help='冻结 AudioEmbeddingProjector 权重，只训练 DiT')

    return parser.parse_args()


def main():
    args = parse_args()
    pl.seed_everything(args.seed)

    if args.exp_name is None:
        args.exp_name = f"vae_scaled_{datetime.now().strftime('%Y%m%d_%H%M%S')}"

    print("=" * 60)
    print("VAE Generator 训练（Scaled 数据 + KL 正则）")
    print("=" * 60)
    print(f"实验名称: {args.exp_name}")
    print(f"训练数据: {args.train_data}")
    print(f"批次大小: {args.batch_size}")
    print()
    print("模型配置:")
    print(f"  模型维度: {args.model_dim}")
    print(f"  层数: {args.n_layers}")
    print(f"  注意力头数: {args.n_heads}")
    print(f"  音频嵌入: {args.pe_audio_in_dim} -> {args.pe_audio_out_dim}")
    print(f"  KL 权重: {args.kl_weight}")
    print("=" * 60)

    # 数据模块
    data_module = ScaledDataModule(
        train_data_dir=args.train_data,
        val_data_dir=args.val_data,
        train_pattern=args.train_pattern,
        val_pattern=args.val_pattern,
        train_arrow_cache=args.train_arrow_cache,
        val_arrow_cache=args.val_arrow_cache,
        batch_size=args.batch_size,
        num_workers=args.num_workers,
    )

    # 模型
    model = VAEGeneratorLightningModule(
        model_dim=args.model_dim,
        n_layers=args.n_layers,
        n_heads=args.n_heads,
        pe_audio_in_dim=args.pe_audio_in_dim,
        pe_audio_out_dim=args.pe_audio_out_dim,
        dac_dim=args.dac_dim,
        dropout=args.dropout,
        inference_mode=args.inference_mode,
        num_steps=args.num_steps,
        kl_weight=args.kl_weight,
        no_reparam=args.no_reparam,
        learning_rate=args.learning_rate,
        weight_decay=args.weight_decay,
        warmup_steps=args.warmup_steps,
        freeze_projector=args.freeze_projector,
    )

    # 参数量
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print(f"\n模型参数: {total_params:,} ({total_params / 1e6:.2f}M)")
    print(f"可训练参数: {trainable_params:,} ({trainable_params / 1e6:.2f}M)")
    if args.freeze_projector:
        frozen_params = total_params - trainable_params
        print(f"冻结参数: {frozen_params:,} ({frozen_params / 1e6:.2f}M) [AudioEmbeddingProjector]")

    # Logger
    logger = TensorBoardLogger(save_dir=args.log_dir, name=args.exp_name)

    # Callbacks
    callbacks = [
        # 按 val loss 保存最优
        ModelCheckpoint(
            dirpath=os.path.join(args.log_dir, args.exp_name, 'checkpoints'),
            filename='epoch={epoch:02d}-step={step}-val_loss={val/loss:.4f}',
            monitor='val/loss',
            mode='min',
            save_top_k=args.save_top_k,
            save_last=True,
            auto_insert_metric_name=False,
        ),
        # 按 step 定期保存
        ModelCheckpoint(
            dirpath=os.path.join(args.log_dir, args.exp_name, 'checkpoints'),
            filename='step={step}',
            every_n_train_steps=args.save_every_n_steps,
            save_top_k=-1,
            auto_insert_metric_name=False,
        ),
        LearningRateMonitor(logging_interval='step'),
        # 音频采样
        AudioSampleCallback(
            every_n_steps=args.audio_sample_every_n_steps,
            num_samples=4,
            num_inference_steps=25,
        ),
    ]

    # GPU
    gpu_ids = [int(x) for x in args.gpu_ids.split(',')]
    print(f"使用 GPU: {gpu_ids}")

    # Trainer
    trainer = pl.Trainer(
        max_epochs=args.max_epochs,
        max_steps=args.max_steps,
        accelerator='gpu',
        devices=gpu_ids,
        strategy=args.strategy if len(gpu_ids) > 1 else 'auto',
        precision=args.precision,
        gradient_clip_val=args.gradient_clip_val,
        accumulate_grad_batches=args.accumulate_grad_batches,
        val_check_interval=args.val_check_interval,
        log_every_n_steps=args.log_every_n_steps,
        logger=logger,
        callbacks=callbacks,
        enable_progress_bar=True,
        enable_model_summary=True,
        num_sanity_val_steps=0,
    )

    print(f"\n开始训练...")
    print(f"TensorBoard: tensorboard --logdir {args.log_dir}")

    # 冻结 projector 时，只加载模型权重（不恢复 optimizer state，因为参数组已变）
    resume_ckpt = args.resume_from
    if args.freeze_projector and args.resume_from:
        print(f"\n[freeze_projector] 从 {args.resume_from} 加载模型权重（跳过 optimizer state）...")
        ckpt = torch.load(args.resume_from, map_location='cpu')
        # Lightning checkpoint 中模型权重在 'state_dict' key 下
        if 'state_dict' in ckpt:
            model.load_state_dict(ckpt['state_dict'])
        else:
            model.load_state_dict(ckpt)
        print("[freeze_projector] 模型权重加载完成，optimizer 将从头初始化")
        resume_ckpt = None  # 不再通过 trainer resume

    trainer.fit(model, datamodule=data_module, ckpt_path=resume_ckpt)

    print("\n训练完成!")
    print(f"最佳模型: {callbacks[0].best_model_path}")


if __name__ == '__main__':
    main()
