-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathtrain_model.py
More file actions
255 lines (214 loc) · 9.92 KB
/
Copy pathtrain_model.py
File metadata and controls
255 lines (214 loc) · 9.92 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
import os
import sys
import json
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader, random_split
from torchvision.models import resnet50, ResNet50_Weights
from torchvision.transforms import transforms
from PIL import Image
import argparse
class CalligraphyDataset(Dataset):
def __init__(self, data_dir, transform=None):
self.data_dir = data_dir
self.transform = transform
self.image_paths = []
self.labels = []
self.char_to_idx = {}
self.idx_to_char = {}
supported_exts = ('.jpg', '.jpeg', '.png', '.gif', '.bmp')
char_idx = 0
for char_name in sorted(os.listdir(data_dir)):
char_dir = os.path.join(data_dir, char_name)
if os.path.isdir(char_dir):
if char_name not in self.char_to_idx:
self.char_to_idx[char_name] = char_idx
self.idx_to_char[char_idx] = char_name
char_idx += 1
for root, _, files in os.walk(char_dir):
for file in files:
if file.lower().endswith(supported_exts):
self.image_paths.append(os.path.join(root, file))
self.labels.append(self.char_to_idx[char_name])
def __len__(self):
return len(self.image_paths)
def __getitem__(self, idx):
img_path = self.image_paths[idx]
try:
image = Image.open(img_path).convert('RGB')
except Exception as e:
print(f"Warning: Skipping corrupted image {img_path}: {e}")
return self.__getitem__((idx + 1) % len(self))
label = self.labels[idx]
if self.transform:
image = self.transform(image)
return image, label
def get_model(num_classes):
weights = ResNet50_Weights.DEFAULT
model = resnet50(weights=weights)
# 解冻后半部分参数
for name, param in model.named_parameters():
if "layer4" in name or "fc" in name:
param.requires_grad = True
else:
param.requires_grad = False
num_ftrs = model.fc.in_features
model.fc = nn.Sequential(
nn.Dropout(0.3),
nn.Linear(num_ftrs, num_classes)
)
return model
def save_checkpoint(state, filename='checkpoint.pth'):
"""保存训练断点"""
torch.save(state, filename)
print(f"Checkpoint saved to {filename}")
def load_checkpoint(checkpoint_path, model, optimizer):
"""加载训练断点"""
if os.path.isfile(checkpoint_path):
print(f"Loading checkpoint from {checkpoint_path}")
checkpoint = torch.load(checkpoint_path)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
start_epoch = checkpoint['epoch']
best_acc = checkpoint['best_acc']
loss = checkpoint['loss']
print(f"Resuming training from epoch {start_epoch + 1}, best accuracy: {best_acc:.4f}")
return start_epoch, best_acc, loss
else:
print(f"No checkpoint found at {checkpoint_path}, starting training from scratch.")
return 0, 0.0, float('inf')
def train_one_epoch(model, dataloader, criterion, optimizer, device, start_batch_idx=0):
"""修改后的train_one_epoch,支持从指定batch开始"""
model.train()
running_loss = 0.0
correct_predictions = 0
total_samples = 0
# 跳过前面的批次
for i, (inputs, labels) in enumerate(dataloader):
if i < start_batch_idx:
continue
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
correct_predictions += torch.sum(preds == labels.data)
total_samples += inputs.size(0)
if i % 50 == 49:
print(f" Batch {i+1}/{len(dataloader)}, Loss: {loss.item():.4f}")
epoch_loss = running_loss / total_samples if total_samples > 0 else 0.0
epoch_acc = correct_predictions.double() / total_samples if total_samples > 0 else 0.0
return epoch_loss, epoch_acc
def evaluate(model, dataloader, criterion, device):
model.eval()
running_loss = 0.0
correct_predictions = 0
total_samples = 0
with torch.no_grad():
for inputs, labels in dataloader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = model(inputs)
loss = criterion(outputs, labels)
running_loss += loss.item() * inputs.size(0)
_, preds = torch.max(outputs, 1)
correct_predictions += torch.sum(preds == labels.data)
total_samples += inputs.size(0)
epoch_loss = running_loss / total_samples
epoch_acc = correct_predictions.double() / total_samples
return epoch_loss, epoch_acc
def main():
parser = argparse.ArgumentParser(description='Train a calligraphy recognition model.')
parser.add_argument('--data-dir', type=str, required=True, help='Path to the chinese_fonts directory.')
parser.add_argument('--epochs', type=int, default=50, help='Number of training epochs.')
parser.add_argument('--batch-size', type=int, default=32, help='Batch size for training.')
parser.add_argument('--lr', type=float, default=0.0005, help='Learning rate.')
parser.add_argument('--resume', type=str, default='', help='Path to checkpoint to resume from')
parser.add_argument('--checkpoint-freq', type=int, default=5, help='Save checkpoint every N epochs')
args = parser.parse_args()
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
if not torch.cuda.is_available():
print("Warning: CUDA not available. Training on CPU will be very slow.")
data_transforms = {
'train': transforms.Compose([
transforms.Resize((224, 224)),
transforms.RandomRotation(15),
transforms.RandomAffine(degrees=0, translate=(0.2, 0.2), scale=(0.7, 1.3)),
transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3),
transforms.RandomGrayscale(p=0.1),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406],
[0.229, 0.224, 0.225])
]),
'val': transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406],
[0.229, 0.224, 0.225])
]),
}
full_dataset = CalligraphyDataset(args.data_dir, transform=data_transforms['train'])
num_classes = len(full_dataset.char_to_idx)
print(f"Found {len(full_dataset)} images belonging to {num_classes} classes.")
# 检查字符映射文件是否存在,如果不存在则创建
if not os.path.exists('char_map.json'):
with open('char_map.json', 'w', encoding='utf-8') as f:
json.dump(full_dataset.char_to_idx, f, ensure_ascii=False, indent=4)
print("Character map saved to char_map.json")
else:
print("Character map already exists, skipping creation.")
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 = CalligraphyDataset(args.data_dir, transform=data_transforms['val'])
dataloaders = {
'train': DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=4),
'val': DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=4)
}
model = get_model(num_classes).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=args.lr)
# 断点续训相关变量
start_epoch = 0
best_acc = 0.0
# 如果指定了断点文件,则加载
if args.resume:
start_epoch, best_acc, _ = load_checkpoint(args.resume, model, optimizer)
# 训练循环
for epoch in range(start_epoch, args.epochs):
print(f"\nEpoch {epoch+1}/{args.epochs}")
print('-' * 10)
train_loss, train_acc = train_one_epoch(model, dataloaders['train'], criterion, optimizer, device)
print(f"Train Loss: {train_loss:.4f} Acc: {train_acc:.4f}")
val_loss, val_acc = evaluate(model, dataloaders['val'], criterion, device)
print(f"Val Loss: {val_loss:.4f} Acc: {val_acc:.4f}")
# 保存最佳模型
if val_acc > best_acc:
best_acc = val_acc
torch.save(model.state_dict(), 'best_model.pth')
print("Best model saved to best_model.pth")
# 定期保存断点
if (epoch + 1) % args.checkpoint_freq == 0:
checkpoint_path = f'checkpoint_epoch_{epoch+1}.pth'
save_checkpoint({
'epoch': epoch + 1,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'best_acc': best_acc,
'loss': val_loss,
}, checkpoint_path)
# 始终保存最新的断点
save_checkpoint({
'epoch': epoch + 1,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'best_acc': best_acc,
'loss': val_loss,
}, 'latest_checkpoint.pth')
print(f"\nTraining complete. Best validation accuracy: {best_acc:.4f}")
if __name__ == '__main__':
main()