import argparse
import os
import sys
import logging

import torch
import torch.nn as nn

from src.data import create_dataloaders
from src.engine import evaluate as eval_fn
from src.models import create_model


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="在验证/测试集上评估检查点")
    p.add_argument("--data-dir", type=str, default="dataset", help="数据集根目录")
    p.add_argument("--checkpoint", type=str, required=True, help="模型检查点路径")
    p.add_argument("--split", type=str, default="val", choices=["val", "test"], help="评估数据集")
    p.add_argument("--batch-size", type=int, default=64, help="批大小")
    p.add_argument("--img-size", type=int, default=224, help="输入图像边长")
    p.add_argument("--workers", type=int, default=4, help="DataLoader 线程数")
    p.add_argument("--log-file", type=str, default=None, help="可选日志文件路径")
    return p.parse_args()


def main() -> None:
    args = parse_args()

    # 日志
    handlers = [logging.StreamHandler(sys.stdout)]
    if args.log_file:
        os.makedirs(os.path.dirname(args.log_file) or ".", exist_ok=True)
        handlers.append(logging.FileHandler(args.log_file, mode="w"))
    logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s", handlers=handlers)

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    logging.info("加载检查点: %s", args.checkpoint)
    ckpt = torch.load(args.checkpoint, map_location=device)
    class_names = tuple(ckpt.get("class_names", ("cats", "dogs")))
    num_classes = len(class_names)
    logging.info("类别: %s | 类别数: %d", class_names, num_classes)

    loaders, _ = create_dataloaders(
        args.data_dir, img_size=args.img_size, batch_size=args.batch_size, workers=args.workers
    )
    if args.split not in loaders:
        raise FileNotFoundError(f"评估 split '{args.split}' 不存在: {args.data_dir}")
    logging.info("评估数据集: %s | 样本数: %d", args.split, len(loaders[args.split].dataset))

    model = create_model(num_classes=num_classes)
    model.load_state_dict(ckpt["model_state"], strict=True)
    model.to(device)

    criterion = nn.CrossEntropyLoss()
    loss, acc = eval_fn(model, criterion, loaders[args.split], device)
    logging.info("结果 | 损失: %.4f | 准确率: %.2f%% | 数据集: %s", loss, acc * 100.0, args.split)


if __name__ == "__main__":
    main()
