diff --git a/.gitignore b/.gitignore index 06703ae..8546bca 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,4 @@ .idea */img -*/data \ No newline at end of file +*/data +model \ No newline at end of file diff --git a/Lab-3/src/__init__.py b/Lab-3/src/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/Lab-3/src/train.py b/Lab-3/src/train.py new file mode 100644 index 0000000..fd24f03 --- /dev/null +++ b/Lab-3/src/train.py @@ -0,0 +1,158 @@ +import torch +import torchvision +from torchvision import datasets, models, transforms +import os +from torch.autograd import Variable +import matplotlib.pyplot as plt +import time + +model_path = '../model/model_name.pth' +model_params_path = '../model/params_name.pth' + +transform = transforms.Compose( + [transforms.CenterCrop(224), + transforms.ToTensor(), + transforms.Normalize([0.5,0.5,0.5],[0.5,0.5,0.5]) + ] +) + +data_dir = "../data" # 更改为您自己的数据集路径 + +data_transform = { + x:transforms.Compose( + [ + transforms.Resize(256), # Resize图像到256x256 + transforms.CenterCrop(224), # 中心裁剪到224x224 + transforms.ToTensor(), + transforms.Normalize( + mean=[0.485, 0.456, 0.406], # ImageNet数据集的 normalize 设置 + std=[0.229, 0.224, 0.225] + ) + ] + ) + for x in ["train","valid"] +} + +image_datasets = { + x:datasets.ImageFolder( + root=os.path.join(data_dir,x), # 将输入参数中的两个名字拼接成一个完整的文件路径 + transform=data_transform[x] + ) + for x in ["train","valid"] +} + +dataloader = { + x:torch.utils.data.DataLoader( + dataset=image_datasets[x], + batch_size=16, + shuffle=True + ) + for x in ["train","valid"] +} + +X_example, y_example = next(iter(dataloader["train"])) +example_classes = image_datasets["train"].classes # ['cat', 'dog'] +index_classes = image_datasets["train"].class_to_idx #{'cat': 0, 'dog': 1} + +Use_gpu = torch.cuda.is_available() +model = models.resnet50(pretrained=True) +print(model) + +for param in model.parameters(): + param.requires_grad = False + +# 重新设计分类器的结构 +model.fc = torch.nn.Linear(2048,2) +print(model) + +if Use_gpu: + model = model.cuda() + +loss_f = torch.nn.CrossEntropyLoss() +optimizer = torch.optim.Adam(model.fc.parameters(), lr=0.00001) + +has_been_trained = os.path.isfile(model_path) +if has_been_trained: + epoch_n = 0 + model.load_state_dict(torch.load(model_params_path)) +else: + epoch_n = 1 + +time_open = time.time() +for epoch in range(epoch_n, epoch_n + 5): # 训练5个epoch + print("Epoch {}/{}".format(epoch, epoch_n + 4)) + print("-" * 10) + + for phase in ["train", "valid"]: + if phase == "train": + print("Training...") + model.train(True) # 启用 BatchNormalization 和 Dropout + else: + print("Validing...") + model.train(False) # 不启用 BatchNormalization 和 Dropout + + running_loss = 0.0 + running_corrects = 0 + + for batch, data in enumerate(dataloader[phase], 1): + X, y = data + if Use_gpu: + X, y = Variable(X.cuda()), Variable(y.cuda()) + else: + X, y = Variable(X), Variable(y) + + y_pred = model(X) + _, pred = torch.max(y_pred.data, 1) + + optimizer.zero_grad() + loss = loss_f(y_pred, y) + + if phase == "train": + loss.backward() + optimizer.step() + + running_loss += loss.item() + running_corrects += torch.sum(pred == y.data) + + if batch % 50 == 0 and phase == "train": + print("Batch {}, Train Loss:{:.4f},Train ACC:{:.4f}%".format( + batch, running_loss / batch, 100.0 * running_corrects / (16 * batch) + ) + ) + + epoch_loss = running_loss * 16 / len(image_datasets[phase]) + epoch_acc = 100.0 * running_corrects / len(image_datasets[phase]) + + print("{} Loss:{:.4f} Acc:{:.4f}%".format(phase, epoch_loss, epoch_acc)) + + # 保存模型参数 + best_acc = 0.0 + if phase == 'valid' and epoch_acc > best_acc: + best_acc = epoch_acc + torch.save(model.state_dict(), model_params_path) + +time_end = time.time() - time_open +print("程序运行时间:{}分钟...".format(int(time_end / 60))) + +# 读取训练集中的一个batch +X_example, Y_example = next(iter(dataloader['train'])) + +if Use_gpu: + X_example, Y_example = Variable(X_example.cuda()), Variable(Y_example.cuda()) +else: + X_example, Y_example = Variable(X_example), Variable(Y_example) + +y_pred = model(X_example) + +index_classes = image_datasets['train'].class_to_idx # 显示类别对应的索引 +example_classes = image_datasets['train'].classes # 将原始图像的类别保存起来 + +img = torchvision.utils.make_grid(X_example) +img = img.cpu().numpy().transpose([1,2,0]) + +print("实际:",[example_classes[i] for i in Y_example]) +_, y_pred = torch.max(y_pred.data,1) +print("预测:",[example_classes[i] for i in y_pred]) + +plt.imshow(img) +plt.show()