dandan4272/hand_gesture_rec
1
1# python imports2# external imports3import torch4import torchvision5from ignite.metrics import Accuracy, Recall, Precision6from matplotlib import pyplot as plt7from torch.utils.data import DataLoader8import torch.nn as nn9 10# model imports11from Mydataset import MyDataset212from torch.nn import functional as F13 14from model.stgcn.Models import TwoStreamSpatialTemporalGraph15 16 17def eval1(net,testloader,use_cuda):18 # 假设您已经定义了模型和测试数据加载器19 net.eval()20 confusion_matrix = torch.zeros(3, 3)21 with torch.no_grad():22 for data, target in testloader:23 if use_cuda:24 data, target = data.cuda(), target.cuda()25 output = model(data)26 predict = output.argmax(dim=1)27 for t, p in zip(predict.view(-1), target.view(-1)):28 confusion_matrix[t.long(), p.long()] += 129 30 TP = confusion_matrix[0, 0]31 FP = confusion_matrix[0, 1]32 FN = confusion_matrix[1, 0]33 TN = confusion_matrix[1, 1]34 35 accuracy = (TP + TN) / (TP + FP + FN + TN)36 precision = TP / (TP + FP)37 recall = TP / (TP + FN)38 f1_score = 2 * precision * recall / (precision + recall)39 40 print(f'Accuracy: {accuracy}')41 print(f'Precision: {precision}')42 print(f'Recall: {recall}')43 print(f'F1 Score: {f1_score}')44 45def eval(net,testloader,use_cuda):46 net.eval()47 test_loss = 048 correct = 049 total = 050 classnum = 351 target_num = torch.zeros((1,classnum))52 predict_num = torch.zeros((1,classnum))53 acc_num = torch.zeros((1,classnum))54 for batch_idx, (inputs, targets) in enumerate(testloader):55 if use_cuda:56 inputs, targets = inputs.cuda(), targets.cuda()57 # inputs, targets = Variable(inputs, volatile=True), Variable(targets)58 outputs = net(inputs)59 loss = F.cross_entropy(outputs, targets)60 # loss = criterion(outputs, targets)61 # loss is variable , if add it(+=loss) directly, there will be a bigger ang bigger graph.62 test_loss += loss.data63 _, predicted = torch.max(outputs.data, 1)64 total += targets.size(0)65 correct += predicted.eq(targets.data).cpu().sum()66 pre_mask = torch.zeros(outputs.size()).scatter_(1, predicted.cpu().view(-1, 1), 1.)67 predict_num += pre_mask.sum(0)68 tar_mask = torch.zeros(outputs.size()).scatter_(1, targets.data.cpu().view(-1, 1), 1.)69 target_num += tar_mask.sum(0)70 acc_mask = pre_mask*tar_mask71 acc_num += acc_mask.sum(0)72 recall = acc_num/target_num73 precision = acc_num/predict_num74 F1 = 2*recall*precision/(recall+precision)75 accuracy = acc_num.sum(1)/target_num.sum(1)76#精度调整77 recall = (recall.numpy()[0]*100).round(3)78 precision = (precision.numpy()[0]*100).round(3)79 F1 = (F1.numpy()[0]*100).round(3)80 accuracy = (accuracy.numpy()[0]*100).round(3)81# 打印格式方便复制82 print('recall'," ".join('%s' % id for id in recall))83 print('precision'," ".join('%s' % id for id in precision))84 print('F1'," ".join('%s' % id for id in F1))85 print('accuracy',accuracy)86 87 88def test_loop(dataloader, model, loss_fn, num_class):89 # 实例化相关metrics的计算对象90 test_acc = Accuracy()91 test_recall = Recall()92 test_precision = Precision()93 94 size = len(dataloader.dataset)95 num_batches = len(dataloader)96 test_loss, correct = 0, 097 98 with torch.no_grad():99 for X, z in dataloader:100 X = X.cuda()101 z = z.cuda()102 pred = model(X)103 test_loss += loss_fn(pred, z).item()104 correct += (pred.argmax(1) == z).type(torch.float).sum().item()105 # 一个batch进行计算迭代106 test_acc.update((pred, z))107 test_recall.update((pred, z))108 test_precision.update((pred, z))109 110 test_loss /= num_batches111 correct /= size112 113 # 计算一个epoch的accuray、recall、precision114 total_acc = test_acc.compute()115 total_recall = test_recall.compute()116 total_precision = test_precision.compute()117 AV_precision = torch.sum(total_precision)/num_class118 print(f"Test Error: \n Accuracy: {(100 * correct):>0.1f}%, "119 f"Avg loss: {test_loss:>8f}, "120 f"ignite acc: {(100 * total_acc):>0.1f}%\n")121 print("recall of every test dataset class: ", total_recall)122 print("precision of every test dataset class: ", total_precision)123 124 # 清空计算对象125 test_precision.reset()126 test_acc.reset()127 test_recall.reset()128 return test_loss,correct,AV_precision129 130 131#%% Visualizing the STN results132def convert_image_np(inp):133 """Convert a Tensor to numpy image."""134 # inp = torch.squeeze(inp)135 inp = inp.numpy().transpose((1, 2, 3, 0))136 137 # for i in inp:138 # inp = inp.numpy().transpose((1, 2,3, 0))139 # inp = inp[5,:,:,:]140 # mean = np.array([0.485, 0.456, 0.406])141 # std = np.array([0.229, 0.224, 0.225])142 # inp = std * inp + mean143 # inp = np.clip(inp, 0, 1)144 return inp145 146def visualize_stn(test_loader):147 with torch.no_grad():148 # Get a batch of training data149 data = next(iter(test_loader))[0].cuda()150 151 input_tensor = data.cpu()152 transformed_input_tensor = model.stn(data).cpu()153 in_grid = convert_image_np(154 torchvision.utils.make_grid(input_tensor))155 156 out_grid = convert_image_np(157 torchvision.utils.make_grid(transformed_input_tensor))158 159 # x = (in_grid,out_grid)160 # Plot the results side-by-side161 for g in range(len(in_grid)):162 f, axarr = plt.subplots(1, 2)163 axarr[0].imshow(in_grid[g])164 axarr[0].set_title('Dataset Images')165 166 axarr[1].imshow(out_grid[g])167 axarr[1].set_title('Transformed Images')168 169 plt.show()170if __name__ == "__main__":171 actions = ['shake_hand', 'palm', 'fist', 'clock_wise', 'anti_clockwise', 'ok', 'thumb', 'v', 'heart','no_gesture']172 # test()173 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")174 epoch = 100 # 加载之前训练的模型(指定迭代次数)175 class_names = ['shake_hand', 'palm', 'fist', 'clock_wise', 'anti_clockwise', 'ok', 'thumb', 'v', 'heart','no_gesture']176 num_class = len(class_names)177 graph_args = {'strategy': 'spatial'}178 model = TwoStreamSpatialTemporalGraph(graph_args, num_class).to(device)179 model.eval()180 los = nn.CrossEntropyLoss()181 # 加载权重182 if(epoch!=0):183 # pre = torch.load(os.path.join('weights/5', '_hand_stgcn_fps15_%d.pth' % epoch))184 pre = torch.load('weights/5/_hand_stgcn_fps15_21node_sgd_200.pth')185 186 model.load_state_dict(pre)187 # model = STGCN(weight_file='weights/_hand_stgcn_100.pth', device=device)188 189 # model.cuda()190 # 加载训练数据191 Data_test = 'datasets/hand/dataset/test'192 dataset = MyDataset2(Data_test)193 dataloader = DataLoader(dataset, batch_size=8, shuffle=True, num_workers=0, drop_last=True)194 195 # eval1(model,dataloader,use_cuda=True)196 # eval(model,dataloader,use_cuda=True)197 test_loop(dataloader,model,los,num_class)198 # visualize_stn(test_loader=dataloader)199# Test200# Error:201# Accuracy: 97.0 %, Avg202# loss: 0.014478, ignite203# acc: 99.4 %204#205# recall206# of207# every208# test209# dataset210#211#212# class: tensor([1.0000, 1.0000, 0.8889, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000],213# dtype=torch.float64)214#215#216# precision217# of218# every219# test220# dataset221#222#223# class: tensor([1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 0.9091, 1.0000, 1.0000],224# dtype=torch.float64)