"""
Semantic DiT 训练脚本 (txt2semantic latent)

从文本（pe_text_embeds + flan_text_feature）生成 semantic latent。

与 peaudio/train_semantic.py 的关键区别：
- AudioEmbeddingProjector 是 VAE 式（mean/logvar → reparameterize）
- projector 冻结，默认保持 train 模式，每次 step 采样 z = mean + std * eps
- --no_reparam 模式：projector 设为 eval，手动 encoder → fc_mean → norm，跳过 reparameterize
  （配合 no-KL VAE 使用，避免 std≈1.0 的噪声污染训练目标）
- flan_text_feature 维度是 1024（peaudio 是 512）
- 使用 ScaledFeatureDataset（带 Arrow 缓存）

训练目标：
  --no_reparam: x1 = encoder(pe_audio_embeds) → fc_mean → norm  (确定性)
  默认:         x1 = AudioEmbeddingProjector.forward(pe_audio_embeds)  (随机采样)
"""

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 dit.dit import DiT
from dit.config import TransformerConfig
from flow_matching import FlowMatchingWrapper
from data_loaders.scaled_dataset import ScaledFeatureDataset, scaled_collate_fn
from generator.vaegen import AudioEmbeddingProjector


class SemanticDiTLightningModule(pl.LightningModule):
    """
    Semantic DiT 训练模块 (txt2semantic latent)

    训练目标：pe_audio_embeds 经过冻结的 VAE AudioEmbeddingProjector 采样后的 z
    projector 冻结但保持 train 模式，每次 training_step 都会产生不同的采样
    """

    def __init__(
        self,
        # 模型配置
        dim: int = 1024,
        n_layers: int = 16,
        n_heads: int = 16,
        in_channels: int = 64,        # 降维后的 semantic 特征维度 (8 或 64)
        out_channels: int = 64,
        global_cond_dim: int = 1024,   # pe_text_embeds 维度
        cross_cond_dim: int = 1024,    # flan_text_feature 维度（semaudio 是 1024）
        max_positions: int = 1024,
        # Audio Projector 配置
        pe_audio_in_dim: int = 1024,
        pe_audio_out_dim: int = 64,    # 8 或 64
        pe_audio_hidden_dim: int = 0,  # 0 = 自动计算 (in_dim + out_dim) // 2
        # 预训练权重路径
        vae_checkpoint_path: Optional[str] = None,
        # Flow Matching 配置
        inference_mode: str = 'euler',
        num_steps: int = 25,
        # CFG 配置
        cfg_dropout: float = 0.1,
        # 训练配置
        learning_rate: float = 1e-4,
        weight_decay: float = 0.01,
        warmup_steps: int = 1000,
        **kwargs
    ):
        super().__init__()
        self.save_hyperparameters()

        # 1. 创建并加载 Audio Projector（VAE 式降维 MLP）
        hidden_dim = pe_audio_hidden_dim if pe_audio_hidden_dim > 0 else None
        self.audio_projector = AudioEmbeddingProjector(
            in_dim=pe_audio_in_dim,
            out_dim=pe_audio_out_dim,
            hidden_dim=hidden_dim,
            dropout=0.1,
        )

        # 加载预训练的 audio_projector 权重
        if vae_checkpoint_path is not None:
            self._load_audio_projector(vae_checkpoint_path)

        # no_reparam 标志
        self.no_reparam = kwargs.get('no_reparam', False)

        # 冻结 audio_projector 参数（不参与训练）
        for param in self.audio_projector.parameters():
            param.requires_grad = False

        if self.no_reparam:
            # no_reparam 模式：projector 设为 eval，不做 reparameterize
            # _get_semantic_target 会手动走 encoder → fc_mean → norm
            self.audio_projector.eval()
            print("🔧 Audio Projector 设置为 eval 模式（no_reparam: 直接用 mean）")
        else:
            # 原始模式：保持 train 模式，reparameterize 采样（mean + std * eps）
            pass

        # 2. 创建 DiT 模型配置
        config = TransformerConfig(
            dim=dim,
            n_layers=n_layers,
            n_heads=n_heads,
            in_channels=in_channels,
            out_channels=out_channels,
            global_cond_dim=global_cond_dim,
            cross_cond_dim=cross_cond_dim,
            use_global_cond=True,
            use_cross_cond=True,
            max_positions=max_positions,
            cfg_dropout=cfg_dropout,
        )

        # 创建 DiT 模型
        dit_model = DiT(config)

        # 3. 使用 Flow Matching Wrapper 包装模型
        self.model = FlowMatchingWrapper(
            model=dit_model,
            inference_mode=inference_mode,
            num_steps=num_steps,
        )

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

        # 保存 pe_audio_out_dim 用于验证
        self.pe_audio_out_dim = pe_audio_out_dim

    def _load_audio_projector(self, checkpoint_path: str):
        """
        从 VAE Generator checkpoint 中加载 audio_projector 权重

        VAE ckpt 中权重名称格式：
        - 'vae_generator.audio_projector.xxx' (Lightning state_dict)
        - 或 'audio_projector.xxx'
        """
        print(f"📂 加载 Audio Projector 权重: {checkpoint_path}")

        checkpoint = torch.load(checkpoint_path, map_location='cpu')

        if 'state_dict' in checkpoint:
            state_dict = checkpoint['state_dict']
        else:
            state_dict = checkpoint

        # 提取 audio_projector 相关的权重
        audio_projector_state_dict = {}
        for key, value in state_dict.items():
            if 'audio_projector' in key:
                new_key = key.split('audio_projector.')[-1]
                audio_projector_state_dict[new_key] = value

        if len(audio_projector_state_dict) == 0:
            print("⚠️ 警告: 未找到 audio_projector 权重，使用随机初始化")
        else:
            self.audio_projector.load_state_dict(audio_projector_state_dict)
            print(f"✅ 成功加载 {len(audio_projector_state_dict)} 个 audio_projector 参数")
            # 打印加载的权重键名
            for key in sorted(audio_projector_state_dict.keys()):
                shape = audio_projector_state_dict[key].shape
                print(f"   {key}: {shape}")

    def on_train_start(self):
        """训练开始时设置 audio_projector 模式"""
        if self.no_reparam:
            self.audio_projector.eval()
            print("🔧 Audio Projector 保持 eval 模式（no_reparam）")
        else:
            self.audio_projector.train()
            print("🔧 Audio Projector 设置为 train 模式（冻结但保持采样）")

    def on_train_batch_start(self, batch, batch_idx):
        """每个 batch 开始时确保 audio_projector 模式正确
        （防止 Lightning 自动调用 eval/train）"""
        if self.no_reparam:
            if self.audio_projector.training:
                self.audio_projector.eval()
        else:
            if not self.audio_projector.training:
                self.audio_projector.train()

    def forward(self, x, t, global_cond=None, cross_cond=None, memory_padding_mask=None):
        """前向传播"""
        return self.model.model(
            x, t,
            global_cond=global_cond,
            cross_cond=cross_cond,
            memory_padding_mask=memory_padding_mask
        )

    def _get_semantic_target(self, pe_audio_embeds: torch.Tensor) -> torch.Tensor:
        """
        获取 semantic 训练目标

        - no_reparam=False (原始模式): 通过冻结的 audio_projector（train 模式）采样 z
          每次调用产生不同的采样结果（reparameterization trick）
        - no_reparam=True: 手动走 encoder → fc_mean → norm，跳过 reparameterize
          与 train_vae_scaled.py --no_reparam 的行为一致

        Args:
            pe_audio_embeds: (batch, 250, 1024)

        Returns:
            semantic_target: (batch, 250, pe_audio_out_dim)
        """
        with torch.no_grad():
            if self.no_reparam:
                # 手动走 encoder → fc_mean → norm，跳过 reparameterize
                projector = self.audio_projector
                h = projector.encoder(pe_audio_embeds)
                mean = projector.fc_mean(h)
                z = projector.norm(mean)
            else:
                # 原始模式：projector 在 train 模式，自动 reparameterize
                z = self.audio_projector(pe_audio_embeds, return_kl=False)
        return z

    def training_step(self, batch, batch_idx):
        """训练步骤"""
        # 提取数据
        pe_audio_embeds = batch['pe_audio_embeds']      # (batch, 250, 1024)
        global_cond = batch['pe_text_embeds']           # (batch, 1024)
        cross_cond = batch['flan_text_feature']         # (batch, max_seq_len, 1024)
        cross_cond_mask = batch['flan_text_mask']       # (batch, max_seq_len)

        # 获取 semantic 训练目标: pe_audio_embeds -> VAE 采样后的 z
        # x1: (batch, 250, pe_audio_out_dim)
        x1 = self._get_semantic_target(pe_audio_embeds)

        # 计算 Flow Matching 损失
        loss_dict = self.model.compute_loss(
            x1=x1,
            global_cond=global_cond,
            cross_cond=cross_cond,
            memory_padding_mask=cross_cond_mask
        )
        loss = loss_dict['loss']

        # 记录损失
        self.log('train/loss', loss, on_step=True, on_epoch=True, prog_bar=True, 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 loss

    def validation_step(self, batch, batch_idx):
        """验证步骤"""
        pe_audio_embeds = batch['pe_audio_embeds']
        global_cond = batch['pe_text_embeds']
        cross_cond = batch['flan_text_feature']
        cross_cond_mask = batch['flan_text_mask']

        # 验证时也使用采样目标（与训练一致）
        x1 = self._get_semantic_target(pe_audio_embeds)

        loss_dict = self.model.compute_loss(
            x1=x1,
            global_cond=global_cond,
            cross_cond=cross_cond,
            memory_padding_mask=cross_cond_mask
        )
        loss = loss_dict['loss']

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

        return loss

    def configure_optimizers(self):
        """配置优化器和学习率调度器 — 只优化 DiT 参数"""
        optimizer = torch.optim.AdamW(
            self.model.parameters(),  # 只包含 DiT 参数（audio_projector 已冻结）
            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 SemanticDataModule(pl.LightningDataModule):
    """
    Semantic DiT 数据模块

    使用 ScaledFeatureDataset（带 Arrow 缓存支持）
    """

    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 = 8,
        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("[SemanticDataModule] 未找到验证集 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='Semantic DiT 训练 (txt2semantic latent)')

    # 数据参数
    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 缓存路径')
    parser.add_argument('--val_arrow_cache', type=str, default=None,
                        help='验证集 Arrow 缓存路径')
    parser.add_argument('--batch_size', type=int, default=32)
    parser.add_argument('--num_workers', type=int, default=8)

    # 模型参数
    parser.add_argument('--dim', type=int, default=1024,
                        help='DiT 模型隐藏层维度')
    parser.add_argument('--n_layers', type=int, default=16,
                        help='Transformer 层数')
    parser.add_argument('--n_heads', type=int, default=16,
                        help='注意力头数')
    parser.add_argument('--in_channels', type=int, default=64,
                        help='输入通道数 (降维后的 semantic 特征维度，应与 pe_audio_out_dim 一致)')
    parser.add_argument('--out_channels', type=int, default=64,
                        help='输出通道数')
    parser.add_argument('--global_cond_dim', type=int, default=1024,
                        help='全局条件维度 (pe_text_embeds)')
    parser.add_argument('--cross_cond_dim', type=int, default=1024,
                        help='交叉条件维度 (flan_text_feature，semaudio 是 1024)')

    # Audio Projector 参数
    parser.add_argument('--pe_audio_in_dim', type=int, default=1024,
                        help='PE 音频嵌入输入维度')
    parser.add_argument('--pe_audio_out_dim', type=int, default=64,
                        help='PE 音频嵌入输出维度（降维后），应与 in_channels 一致')
    parser.add_argument('--pe_audio_hidden_dim', type=int, default=0,
                        help='Audio Projector 隐藏层维度 (0=自动计算)')
    parser.add_argument('--vae_checkpoint', type=str, default=None,
                        help='预训练的 VAE checkpoint 路径，用于加载 audio_projector 权重')

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

    # CFG 参数
    parser.add_argument('--cfg_dropout', type=float, default=0.1,
                        help='训练时条件 dropout 概率，用于 CFG 训练')

    # 训练参数
    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('--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',
                        help='指定使用的 GPU ID (如 "0" 或 "0,1")')
    parser.add_argument('--precision', type=str, default='bf16-mixed',
                        choices=['32', '16-mixed', 'bf16-mixed'])
    parser.add_argument('--strategy', type=str, default='auto')

    # no_reparam 模式
    parser.add_argument('--no_reparam', action='store_true',
                        help='不做 reparameterize，直接用 mean（配合 no-KL VAE）')

    # 其他
    parser.add_argument('--seed', type=int, default=42)
    parser.add_argument('--resume_from', type=str, default=None,
                        help='从检查点恢复训练')

    return parser.parse_args()


def main():
    args = parse_args()

    # 强制 in_channels / out_channels 与 pe_audio_out_dim 一致
    if args.in_channels != args.pe_audio_out_dim:
        print(f"⚠️ 警告: in_channels ({args.in_channels}) 与 pe_audio_out_dim ({args.pe_audio_out_dim}) 不一致")
        args.in_channels = args.pe_audio_out_dim
    args.out_channels = args.pe_audio_out_dim
    print(f"📐 in_channels = out_channels = pe_audio_out_dim = {args.pe_audio_out_dim}")

    pl.seed_everything(args.seed)

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

    print("=" * 60)
    print("Semantic DiT 训练 (txt2semantic latent)")
    print("=" * 60)
    print(f"实验名称: {args.exp_name}")
    print(f"训练数据: {args.train_data}")
    print(f"批次大小: {args.batch_size}")
    print()
    print("DiT 模型配置:")
    print(f"  模型维度: {args.dim}")
    print(f"  层数: {args.n_layers}")
    print(f"  注意力头数: {args.n_heads}")
    print(f"  输入/输出通道数: {args.in_channels}")
    print(f"  全局条件维度 (pe_text_embeds): {args.global_cond_dim}")
    print(f"  交叉条件维度 (flan_text_feature): {args.cross_cond_dim}")
    print(f"  CFG dropout: {args.cfg_dropout}")
    print()
    print("Audio Projector 配置 (VAE 式, 冻结):")
    print(f"  输入维度: {args.pe_audio_in_dim}")
    print(f"  输出维度: {args.pe_audio_out_dim}")
    print(f"  隐藏维度: {args.pe_audio_hidden_dim} (0=自动)")
    print(f"  no_reparam: {args.no_reparam}")
    print(f"  VAE checkpoint: {args.vae_checkpoint or '无（随机初始化）'}")
    print()
    print("训练配置:")
    print(f"  学习率: {args.learning_rate}")
    print(f"  最大轮数: {args.max_epochs}")
    print(f"  训练精度: {args.precision}")
    print("=" * 60)

    # 数据模块
    data_module = SemanticDataModule(
        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 = SemanticDiTLightningModule(
        dim=args.dim,
        n_layers=args.n_layers,
        n_heads=args.n_heads,
        in_channels=args.in_channels,
        out_channels=args.out_channels,
        global_cond_dim=args.global_cond_dim,
        cross_cond_dim=args.cross_cond_dim,
        pe_audio_in_dim=args.pe_audio_in_dim,
        pe_audio_out_dim=args.pe_audio_out_dim,
        pe_audio_hidden_dim=args.pe_audio_hidden_dim,
        vae_checkpoint_path=args.vae_checkpoint,
        inference_mode=args.inference_mode,
        num_steps=args.num_steps,
        cfg_dropout=args.cfg_dropout,
        learning_rate=args.learning_rate,
        weight_decay=args.weight_decay,
        warmup_steps=args.warmup_steps,
        no_reparam=args.no_reparam,
    )

    # 参数量统计
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    frozen_params = total_params - trainable_params
    print(f"\n模型参数统计:")
    print(f"  总参数量: {total_params:,} ({total_params / 1e6:.2f}M)")
    print(f"  可训练参数量 (DiT): {trainable_params:,} ({trainable_params / 1e6:.2f}M)")
    print(f"  冻结参数量 (Audio Projector): {frozen_params:,} ({frozen_params / 1e6:.2f}M)")

    # 验证 DiT in_channels == pe_audio_out_dim
    print(f"\n✅ 验证: DiT in_channels = {args.in_channels}, pe_audio_out_dim = {args.pe_audio_out_dim}")
    assert args.in_channels == args.pe_audio_out_dim, \
        f"DiT in_channels ({args.in_channels}) != pe_audio_out_dim ({args.pe_audio_out_dim})"

    # 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'),
    ]

    # GPU
    if args.gpu_ids == 'auto':
        # torchrun 模式: Lightning 自动检测设备
        gpu_devices = 'auto'
        strategy = args.strategy
        print(f"\n使用 GPU: auto (由 torchrun 管理)")
    else:
        gpu_ids = [int(x) for x in args.gpu_ids.split(',')]
        gpu_devices = gpu_ids
        strategy = args.strategy if len(gpu_ids) > 1 else 'auto'
        print(f"\n使用 GPU: {gpu_ids}")

    # 数据、模型、callbacks 都准备好了，现在杀掉占卡程序 occ.py，然后立刻创建 Trainer 接管 GPU
    import subprocess
    print(f"\n🔪 杀掉占卡程序 occ.py ...")
    my_pid = os.getpid()
    try:
        result = subprocess.run(['pgrep', '-f', 'occ.py'], capture_output=True, text=True)
        pids = result.stdout.strip().split('\n')
        for pid in pids:
            pid = pid.strip()
            if pid and int(pid) != my_pid:
                os.kill(int(pid), 9)
                print(f"  已杀掉 PID={pid}")
    except Exception as e:
        print(f"  ⚠️ {e}")
    import time; time.sleep(3)  # 等 GPU 显存释放

    # Trainer（立刻接管 GPU）
    trainer = pl.Trainer(
        max_epochs=args.max_epochs,
        max_steps=args.max_steps,
        accelerator='gpu',
        devices=gpu_devices,
        strategy=strategy,
        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}")

    trainer.fit(
        model,
        datamodule=data_module,
        ckpt_path=args.resume_from,
    )

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


if __name__ == '__main__':
    main()
