{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ['KMP_DUPLICATE_LIB_OK']='True'\nimport torch\nimport numpy as np\nimport albumentations as A\n\nfrom PIL import Image\nfrom tqdm import tqdm\nimport glob\nfrom sklearn.metrics import roc_curve, roc_auc_score\nimport matplotlib.pyplot as plt\n\ndef load_model():\n    model = torch.load(model_path)\n    model.to(device)\n    model.eval()\n    return model\n\nvalid_aug = A.Compose([\n            # A.Resize(*[896,896], interpolation=cv2.INTER_NEAREST, p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], p=1.0),\n            ], p=1.0)\n\n\ndef pading(size, img):\n    target_size = size\n    padding_v = tuple([125, 125, 125])\n    interpolation = Image.BILINEAR\n    w, h = img.size\n    if w > h:\n        img = img.resize((int(target_size), int(h * target_size * 1.0 / w)), interpolation)\n    else:\n        img = img.resize((int(w * target_size * 1.0 / h), int(target_size)), interpolation)\n\n    ret_img = Image.new(\"RGB\", (target_size, target_size), padding_v)\n    w, h = img.size\n    st_w = int((ret_img.size[0] - w) / 2.0)\n    st_h = int((ret_img.size[1] - h) / 2.0)\n    ret_img.paste(img, (st_w, st_h))\n    # ret_img = np.array(ret_img)\n    return ret_img\n\ndef cal_roc_auc():\n    total = 0\n    correct = 0\n    y_true_list = list()\n    y_score_list = list()\n    with torch.no_grad():\n\n        for img_path in tqdm(glob.glob(input_dir + \"/*/*\")):\n            # img_name = os.path.split(img_path)[-1]\n            # print(img_path)\n            label = int(img_path.split('\\\\')[-2])\n\n            image = Image.open(img_path).convert(\"RGB\")\n            img = pading(448, image)\n            input_image = valid_aug(image=np.array(img))['image']\n            input_image = np.transpose(input_image, (2, 0, 1))  # [c, h, w]\n\n            input_image = input_image[np.newaxis, :]  # 增加一个维度\n            input_image = torch.tensor(input_image).cuda()\n            res1 = model(input_image)\n            y_score = res1.softmax(dim=1).detach().cpu().numpy()[0][1]\n            y_score_list.append(y_score)\n            y_true_list.append(label)\n            _, predicted = torch.max(res1.data, 1)\n            # print(res1)\n            # prob1 = torch.nn.functional.softmax(res1,dim=1)\n            # result = prob1[:,1].detach().cpu()\n            result0 = predicted.cpu().numpy()[0]\n            if label == result0:\n                correct += 1\n            total += 1\n\n    print(\"acc:\",correct / total)\n    fpr, tpr, _ = roc_curve(y_true_list, y_score_list)\n\n    # 计算 AUC 值\n    auc_score = roc_auc_score(y_true_list, y_score_list)\n    print(\"auc_score:\",auc_score)\n\n    # 绘制 ROC 曲线\n    plt.figure()\n    plt.plot(fpr, tpr, color='darkorange', lw=2, label='ROC curve (area = %0.2f)' % auc_score)\n    plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('Receiver Operating Characteristic (ROC) Curve')\n    plt.legend(loc=\"lower right\")\n    plt.show()\n\n\n\nif __name__ == '__main__':\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    input_dir = r\"C:\\Users\\Administrator\\Desktop\\Data\\VCF2\"\n    model_path = r\"D:\\lwh\\python_project\\Dark_circle\\model_2\\efficientnet_b4_2\\best_fold.pth\"\n    model = load_model()\n    cal_roc_auc()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]}]}