You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何构建CNN实现照片与手绘作品的自动分类?

手绘与照片分类的CNN实现方案(适合大一工科生)

一、先搞定数据基础

  • 小规模标注:从12000张图里各挑500-1000张明确的手绘/照片,分别放进train_handdrawn、train_photo(训练集)和val_handdrawn、val_photo(验证集)文件夹,不用全标,先跑通流程再说。
  • 统一图片规格:用Python的PIL库把所有图片缩放到224x224,代码示例:
    from PIL import Image
    import os
    
    def resize_imgs(src_dir, dest_dir, size=(224,224)):
        os.makedirs(dest_dir, exist_ok=True)
        for fname in os.listdir(src_dir):
            fpath = os.path.join(src_dir, fname)
            try:
                with Image.open(fpath) as img:
                    img.resize(size).save(os.path.join(dest_dir, fname))
            except Exception as e:
                print(f"跳过损坏图片:{fname}")
    
  • 简单数据增强:给训练集加翻转、小角度旋转,避免模型过拟合,用torchvision实现:
    from torchvision import transforms
    train_transform = transforms.Compose([
        transforms.Resize((224,224)),
        transforms.RandomHorizontalFlip(p=0.5),
        transforms.RandomRotation(10),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])
    val_transform = transforms.Compose([
        transforms.Resize((224,224)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])
    

二、用迁移学习搭CNN(不用从零造轮子)

大一阶段别自己写复杂的CNN层,直接用预训练模型改,比如ResNet18,代码用PyTorch(入门友好):

import torch
import torch.nn as nn
from torchvision import models

# 加载预训练的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, 2)

三、训练与验证流程

  • 用ImageFolder和DataLoader加载数据:
    from torch.utils.data import DataLoader
    from torchvision.datasets import ImageFolder
    
    train_dataset = ImageFolder('训练集根目录', transform=train_transform)
    val_dataset = ImageFolder('验证集根目录', transform=val_transform)
    train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
    val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
    
  • 定义损失函数和优化器:
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.Adam(model.fc.parameters(), lr=0.001)
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    model.to(device)
    
  • 训练循环(简化版):
    for epoch in range(10):
        # 训练模式
        model.train()
        train_loss = 0.0
        for imgs, labels in train_loader:
            imgs, labels = imgs.to(device), labels.to(device)
            optimizer.zero_grad()
            outputs = model(imgs)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()
            train_loss += loss.item() * imgs.size(0)
        avg_train_loss = train_loss / len(train_dataset)
    
        # 验证模式
        model.eval()
        val_loss = 0.0
        correct = 0
        total = 0
        with torch.no_grad():
            for imgs, labels in val_loader:
                imgs, labels = imgs.to(device), labels.to(device)
                outputs = model(imgs)
                loss = criterion(outputs, labels)
                val_loss += loss.item() * imgs.size(0)
                _, preds = torch.max(outputs, 1)
                total += labels.size(0)
                correct += (preds == labels).sum().item()
        avg_val_loss = val_loss / len(val_dataset)
        val_acc = correct / total
        print(f"第{epoch+1}轮: 训练损失{avg_train_loss:.4f}, 验证损失{avg_val_loss:.4f}, 验证准确率{val_acc:.4f}")
    

四、批量分类所有图片

训练好模型后,批量处理12000张图:

def batch_classify(model, src_dir, dest_handdrawn, dest_photo, transform):
    os.makedirs(dest_handdrawn, exist_ok=True)
    os.makedirs(dest_photo, exist_ok=True)
    model.eval()
    with torch.no_grad():
        for fname in os.listdir(src_dir):
            fpath = os.path.join(src_dir, fname)
            try:
                img = Image.open(fpath).convert('RGB')
                img_tensor = transform(img).unsqueeze(0).to(device)
                output = model(img_tensor)
                _, pred = torch.max(output, 1)
                # 这里的0/1对应你训练集的标签顺序,比如0是手绘,1是照片
                if pred.item() == 0:
                    img.save(os.path.join(dest_handdrawn, fname))
                else:
                    img.save(os.path.join(dest_photo, fname))
            except Exception as e:
                print(f"分类失败:{fname}")

五、大一阶段的实用建议

  • 不用全量标注,先拿1000-2000张图跑通,效果不够再逐步增加标注数据。
  • 本地没GPU的话,用免费云端GPU平台,训练速度能快好几倍。
  • 先把代码跑起来,再慢慢理解CNN各层的作用,不用一开始就啃透底层原理。

内容的提问来源于stack exchange,提问作者Selim Rebai

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 11:07:11