CoolFace
Apppublic

nchdlhbctm/TraceDetect-AI

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes
train_image_model.py81 linesDownload Raw Back to root
1import os
2import torch
3import torch.nn as nn
4import torch.optim as optim
5from torchvision import datasets, models, transforms
6from torch.utils.data import DataLoader
7from tqdm import tqdm  # 引入了实时进度条神器
8
9
10def train_model():
11    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
12    print(f"当前使用的计算设备: {device}")
13
14    if device.type == 'cpu':
15        print("⚠️ 警告:当前正在使用 CPU 训练,4000张图片预计每轮需要 15-30 分钟,请保持耐心!")
16
17    data_transforms = transforms.Compose([
18        transforms.Resize((224, 224)),
19        transforms.RandomHorizontalFlip(),
20        transforms.ToTensor(),
21        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
22    ])
23
24    data_dir = './data'
25    image_dataset = datasets.ImageFolder(data_dir, data_transforms)
26    # CPU 训练比较慢,我们把 batch_size 稍微调大一点点到 16
27    dataloader = DataLoader(image_dataset, batch_size=16, shuffle=True)
28
29    print(f"总计训练图片数量: {len(image_dataset)} 张\n")
30
31    print("正在加载 MobileNetV2 模型...")
32    model = models.mobilenet_v2(weights=models.MobileNet_V2_Weights.IMAGENET1K_V1)
33
34    num_ftrs = model.classifier[1].in_features
35    model.classifier[1] = nn.Linear(num_ftrs, 2)
36    model = model.to(device)
37
38    criterion = nn.CrossEntropyLoss()
39    optimizer = optim.Adam(model.parameters(), lr=1e-4)
40
41    # 为了用 CPU 能快点看到结果,我们先设为 3 轮
42    num_epochs = 3
43    print("\n--- 开始模型微调 ---")
44
45    for epoch in range(num_epochs):
46        model.train()
47        running_loss = 0.0
48        corrects = 0
49
50        # 【核心修改】:用 tqdm 包装 dataloader,生成实时进度条
51        progress_bar = tqdm(dataloader, desc=f"第 {epoch + 1}/{num_epochs} 轮", leave=False, colour='green')
52
53        for inputs, labels in progress_bar:
54            inputs = inputs.to(device)
55            labels = labels.to(device)
56
57            optimizer.zero_grad()
58            outputs = model(inputs)
59            _, preds = torch.max(outputs, 1)
60            loss = criterion(outputs, labels)
61
62            loss.backward()
63            optimizer.step()
64
65            running_loss += loss.item() * inputs.size(0)
66            corrects += torch.sum(preds == labels.data)
67
68            # 让进度条实时显示当前的误差值
69            progress_bar.set_postfix({'loss': f"{loss.item():.4f}"})
70
71        epoch_loss = running_loss / len(image_dataset)
72        epoch_acc = corrects.double() / len(image_dataset)
73        print(f"✅ 第 {epoch + 1}/{num_epochs} 轮完成 | 平均损失: {epoch_loss:.4f} | 准确率: {epoch_acc:.4f}")
74
75    save_path = 'mobilenet_finetuned.pth'
76    torch.save(model.state_dict(), save_path)
77    print(f"\n🎉 训练完成!模型权重已保存至: {save_path}")
78
79
80if __name__ == '__main__':
81    train_model()