CoolFace
Apppublic

dandan4272/hand_gesture_rec

sourceHugging Facemitupdated 3y agoView on Hugging Face
1likes
test.py224 linesDownload Raw Back to root
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)