nchdlhbctm/TraceDetect-AI
0
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()