# import os
# import cv2
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms, models
from PIL import Image
# import numpy as np

# ===================== 全局配置 =====================
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
BATCH_SIZE = 16
EPOCHS = 20
LEARNING_RATE = 0.001
IMAGE_SIZE = 224  # ResNet输入尺寸
MODEL_SAVE_PATH = "id_card_detector.pth"  # 模型保存路径

# 数据集根目录(必须按照以下结构放置图片)
# dataset_2/
#   ├── normal/    # 真实身份证图片
#   └── abnormal/  # 翻拍/P图/篡改图片
DATASET_PATH = "./dataset_2"


# ===================== 1. 数据预处理与加载 =====================
def get_data_transforms():
    """训练/验证数据增强与预处理"""
    train_transform = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.RandomHorizontalFlip(p=0.3),  # 随机水平翻转
        transforms.RandomRotation(10),  # 随机旋转
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])  # ImageNet均值方差
    ])

    val_transform = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    ])
    return train_transform, val_transform


def load_dataset():
    """加载数据集,划分训练集/验证集"""
    train_transform, val_transform = get_data_transforms()

    # 加载完整数据集
    full_dataset = datasets.ImageFolder(DATASET_PATH, transform=train_transform)
    # 划分 80%训练,20%验证
    train_size = int(0.8 * len(full_dataset))
    val_size = len(full_dataset) - train_size
    train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])

    # 验证集使用验证预处理
    val_dataset.dataset.transform = val_transform

    # DataLoader
    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)
    val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)

    print(f"数据集类别: {full_dataset.class_to_idx}")  # 输出:{'abnormal': 1, 'normal': 0}
    print(f"训练集数量: {train_size}, 验证集数量: {val_size}")
    return train_loader, val_loader, full_dataset.class_to_idx


# ===================== 2. 构建ResNet18模型 =====================
def build_resnet18_model(num_classes=2):
    """加载预训练ResNet18,修改全连接层适配二分类"""
    model = models.resnet18(pretrained=True)  # 加载预训练权重
    # 冻结主干网络(只训练最后一层,小数据集更稳定)
    for param in model.parameters():
        param.requires_grad = False

    # 替换最后一层全连接层(二分类输出2个神经元)
    in_features = model.fc.in_features
    model.fc = nn.Linear(in_features, num_classes)
    model = model.to(DEVICE)
    return model


# ===================== 3. 模型训练 =====================
def train_model(model, train_loader, val_loader):
    # 损失函数 + 优化器
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.fc.parameters(), lr=LEARNING_RATE)
    best_acc = 0.0

    print(f"\n开始训练,使用设备: {DEVICE}")
    for epoch in range(EPOCHS):
        # 训练阶段
        model.train()
        train_loss = 0.0
        train_correct = 0

        for images, labels in train_loader:
            images, labels = images.to(DEVICE), labels.to(DEVICE)
            optimizer.zero_grad()

            outputs = model(images)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()

            train_loss += loss.item() * images.size(0)
            _, preds = torch.max(outputs, 1)
            train_correct += torch.sum(preds == labels.data)

        # 验证阶段
        model.eval()
        val_loss = 0.0
        val_correct = 0

        with torch.no_grad():
            for images, labels in val_loader:
                images, labels = images.to(DEVICE), labels.to(DEVICE)
                outputs = model(images)
                loss = criterion(outputs, labels)

                val_loss += loss.item() * images.size(0)
                _, preds = torch.max(outputs, 1)
                val_correct += torch.sum(preds == labels.data)

        # 计算指标
        train_loss = train_loss / len(train_loader.dataset)
        train_acc = train_correct.double() / len(train_loader.dataset)
        val_loss = val_loss / len(val_loader.dataset)
        val_acc = val_correct.double() / len(val_loader.dataset)

        print(f"Epoch {epoch + 1}/{EPOCHS} | "
              f"训练损失: {train_loss:.4f} 训练准确率: {train_acc:.4f} | "
              f"验证损失: {val_loss:.4f} 验证准确率: {val_acc:.4f}")

        # 保存最优模型
        if val_acc > best_acc:
            best_acc = val_acc
            torch.save(model.state_dict(), MODEL_SAVE_PATH)
            print(f"✅ 最优模型已保存: {MODEL_SAVE_PATH}")

    print(f"\n训练完成!最佳验证准确率: {best_acc:.4f}")


# ===================== 4. 推理预测(单张图片) =====================
def predict_image(image_path, model_path=MODEL_SAVE_PATH):
    """
    推理函数:输入图片路径,返回是否P图/翻拍
    返回:"否"(正常身份证) or "是"(P图/翻拍)
    """
    # 加载模型
    model = build_resnet18_model()
    model.load_state_dict(torch.load(model_path, map_location=DEVICE))
    model.eval()

    # 图片预处理
    transform = transforms.Compose([
        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
    ])

    # 读取图片
    image = Image.open(image_path).convert("RGB")
    image = transform(image).unsqueeze(0).to(DEVICE)

    # 预测
    with torch.no_grad():
        outputs = model(image)
        _, pred = torch.max(outputs, 1)

    # 结果映射
    class_names = {0: "否", 1: "是"}
    result = class_names[pred.item()]
    confidence = torch.softmax(outputs, dim=1)[0][pred].item()

    print(f"\n🔍 图片路径: {image_path}")
    print(f"📊 检测结果: {result}(是否P图/翻拍)")
    print(f"✅ 置信度: {confidence:.4f}")
    return result


# ===================== 主函数 =====================
if __name__ == "__main__":
    '''
        你的项目文件夹/
    ├── dataset_2/           # 数据集根目录
    │   ├── normal/        # 存放【真实拍摄】的身份证图片(几百张即可)
    │   └── abnormal/      # 存放【P图/翻拍/截图/篡改】的身份证图片
    ├── id_card_detector.pth  # 训练后自动生成的模型文件
    └── classify_2.py            # 上面的代码文件
    '''

    # ============= 1. 训练模型 =============
    # 请先准备好数据集!!!
    train_loader, val_loader, class_to_idx = load_dataset()
    model = build_resnet18_model()
    train_model(model, train_loader, val_loader)

    # ============= 2. 推理测试 =============
    # 训练完成后,替换为你的测试图片路径
    # test_img_path = "./test_id_card.jpg"
    # result = predict_image(test_img_path)

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐