import argparse
import os
import sys
import time
import logging

import torch
import torch.nn as nn

from src.data import create_dataloaders
from src.models import create_model
from src.engine import train_one_epoch, evaluate
from src.utils import set_seed, ensure_dir, save_checkpoint


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="最简猫狗分类训练脚本（自动切分数据）")
    p.add_argument("--data-dir", type=str, default="dataset", help="数据集根目录")
    p.add_argument("--epochs", type=int, default=15, help="训练轮数")
    p.add_argument("--batch-size", type=int, default=32, help="批大小")
    p.add_argument("--lr", type=float, default=1e-3, help="学习率")
    p.add_argument("--weight-decay", type=float, default=1e-4, help="权重衰减")
    p.add_argument("--img-size", type=int, default=224, help="输入图像边长")
    p.add_argument("--workers", type=int, default=4, help="DataLoader 线程数")
    p.add_argument("--output-dir", type=str, default="runs/exp1", help="输出目录（日志与模型）")
    p.add_argument("--log-file", type=str, default="train.log", help="日志文件名或绝对路径")
    p.add_argument("--seed", type=int, default=42, help="随机种子")
    p.add_argument("--no-amp", action="store_true", help="禁用混合精度训练")
    return p.parse_args()


def main() -> None:
    args = parse_args()
    set_seed(args.seed)

    # 输出与日志
    ensure_dir(args.output_dir)
    log_path = args.log_file if os.path.isabs(args.log_file) else os.path.join(args.output_dir, args.log_file)
    os.makedirs(os.path.dirname(log_path) or ".", exist_ok=True)
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s | %(levelname)s | %(message)s",
        handlers=[logging.StreamHandler(sys.stdout), logging.FileHandler(log_path, mode="w")],
    )
    logging.info("日志文件: %s", log_path)

    # 数据
    loaders, class_names = create_dataloaders(
        args.data_dir, img_size=args.img_size, batch_size=args.batch_size, workers=args.workers
    )
    for sp, ld in loaders.items():
        logging.info("数据划分 %s: %d 张", sp, len(ld.dataset))
    logging.info("类别: %s | 类别数: %d", class_names, len(class_names))

    # 模型
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = create_model(num_classes=len(class_names))
    model.to(device)
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    gpu_name = torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU"
    try:
        import torchvision
        tv_ver = torchvision.__version__
    except Exception:
        tv_ver = "unknown"
    logging.info(
        "环境 | PyTorch: %s | TorchVision: %s | 设备: %s | AMP: %s",
        torch.__version__, tv_ver, gpu_name, str(torch.cuda.is_available() and not args.no_amp),
    )
    logging.info("模型: EfficientNet-B0 | 总参数: %d | 可训练参数: %d", total_params, trainable_params)
    logging.info(
        "训练超参 | epoch: %d | batch: %d | lr: %.6f | weight_decay: %.6f | img_size: %d | workers: %d | 随机种子: %d",
        args.epochs,
        args.batch_size,
        args.lr,
        args.weight_decay,
        args.img_size,
        args.workers,
        args.seed,
    )

    # 损失与优化
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.SGD(
        [p for p in model.parameters() if p.requires_grad],
        lr=args.lr,
        momentum=0.9,
        weight_decay=args.weight_decay,
    )
    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)
    scaler = torch.cuda.amp.GradScaler(enabled=(torch.cuda.is_available() and not args.no_amp))

    # 训练循环
    best_acc = 0.0
    for epoch in range(1, args.epochs + 1):
        start = time.perf_counter()
        cur_lr = optimizer.param_groups[0].get("lr", args.lr)
        logging.info("开始第 %d/%d 轮 | 学习率: %.6f", epoch, args.epochs, cur_lr)
        train_loss, train_acc = train_one_epoch(
            model, criterion, optimizer, loaders["train"], device, scaler, scaler.is_enabled()
        )
        val_loss, val_acc = evaluate(model, criterion, loaders["val"], device)
        scheduler.step()
        elapsed = time.perf_counter() - start
        logging.info(
            "第 %d 轮完成 | 训练: 损失 %.4f, 准确率 %.2f%% | 验证: 损失 %.4f, 准确率 %.2f%% | 用时: %.1fs",
            epoch,
            train_loss,
            train_acc * 100.0,
            val_loss,
            val_acc * 100.0,
            elapsed,
        )

        is_best = val_acc > best_acc
        best_acc = max(best_acc, val_acc)
        save_checkpoint(
            {
                "epoch": epoch,
                "model_state": model.state_dict(),
                "optimizer_state": optimizer.state_dict(),
                "scheduler_state": scheduler.state_dict(),
                "scaler_state": scaler.state_dict(),
                "best_acc": best_acc,
                "class_names": class_names,
                "img_size": args.img_size,
            },
            is_best=is_best,
            output_dir=args.output_dir,
        )
        logging.info(
            "已保存检查点: %s | 是否刷新最佳: %s | 当前最佳验证准确率: %.2f%%",
            os.path.join(args.output_dir, "last.pth"),
            str(is_best),
            best_acc * 100.0,
        )

    logging.info("训练完成 | 最佳验证准确率: %.2f%% | 输出目录: %s", best_acc * 100.0, args.output_dir)

    # 可选：若存在测试集，则进行一次评估
    if "test" in loaders:
        test_loss, test_acc = evaluate(model, criterion, loaders["test"], device)
        logging.info(
            "测试集结果 | 损失: %.4f | 准确率: %.2f%% | 样本数: %d",
            test_loss,
            test_acc * 100.0,
            len(loaders["test"].dataset),
        )


if __name__ == "__main__":
    main()
