{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"#训练好的模型\n#../input/effadamplabelsmooth\n#../input/cassava-leaf-disease-classification\n#../input/timm-pytorch-image-models\n\n\n\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm\n\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom torch.autograd import Variable\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport torchvision\nfrom tqdm import tqdm\nimport cv2\nfrom skimage import io\nimport time\n\nfrom albumentations import (\n    HorizontalFlip, VerticalFlip, IAAPerspective, ShiftScaleRotate, CLAHE, RandomRotate90,\n    Transpose, ShiftScaleRotate, Blur, OpticalDistortion, GridDistortion, HueSaturationValue,\n    IAAAdditiveGaussianNoise, GaussNoise, MotionBlur, MedianBlur, IAAPiecewiseAffine, RandomResizedCrop,\n    IAASharpen, IAAEmboss, RandomBrightnessContrast, Flip, OneOf, Compose, Normalize, Cutout, CoarseDropout, ShiftScaleRotate, CenterCrop, Resize\n)\n\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.metrics import accuracy_score\n\nfrom sklearn.model_selection import GroupKFold, StratifiedKFold\nimport tqdm.notebook as tq\nfrom sklearn.model_selection import train_test_split\nfrom scipy.special import softmax\n\n\nDATA_DIR = '../input/cassava-leaf-disease-classification'\nCFG = {\n    'img_size': 512,\n    'tta': 3,\n    'valid_bs': 32,\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'effnet_models': ['tf_efficientnet_b4_ns_fold_0_5.pt', 'tf_efficientnet_b4_ns_fold_0_6.pt', 'tf_efficientnet_b4_ns_fold_0_7.pt', 'tf_efficientnet_b4_ns_fold_0_8.pt'],\n}\n\n\nclass DiseaseDatasetInference(torch.utils.data.Dataset):\n\n    def __init__ (self, df, transform=None, opt_label=True):\n        self.df = df.reset_index(drop=True).copy()\n        self.transform = transform\n        self.opt_label = opt_label\n\n        if self.opt_label:\n            self.data = [(row['image_id'], row['label']) for _, row in self.df.iterrows()]\n\n        else:\n            self.data = [(row['image_id']) for _, row in self.df.iterrows()]\n\n        self.data = np.asarray(self.data)\n  \n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__ (self, index):\n            # np.random.shuffle(self.data)\n        if self.opt_label:\n            image_path, label = self.data[index]    \n        else:\n            image_path = self.data[index]\n\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        if self.transform is not None:\n            image = self.transform(image=image)['image']\n\n        if self.opt_label == True:\n            return (image, int(label))\n\n        else:\n            return image\n\ndef get_inference_transforms():\n    return Compose([\n            CenterCrop(CFG['img_size'], CFG['img_size'], p=0.5),\n            Resize(CFG['img_size'], CFG['img_size']),\n            Transpose(p=0.5),\n            ShiftScaleRotate(p=0.5),\n            HorizontalFlip(p=0.5),\n            VerticalFlip(p=0.5),\n            HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n            RandomBrightnessContrast(brightness_limit=(-0.1,0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n            ToTensorV2(p=1.0),\n            \n        ], p=1.)\n\n\n\ndf = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv')\nPATH = '/kaggle/input/cassava-leaf-disease-classification/test_images/'\n#train = pd.read_csv(f'{DATA_DIR}/train.csv')\n\n\ntest_csv = df.copy()\ntest_csv['image_id'] = PATH + test_csv['image_id']\n\ntest_ds = DiseaseDatasetInference(test_csv, transform=get_inference_transforms(), opt_label=False)\n\ntest_loader = torch.utils.data.DataLoader(test_ds, batch_size=CFG['valid_bs'], shuffle=False, pin_memory=False) \n\n\ndef inference (model, data_loader, device):\n    preds = []\n    model.eval()\n    test_tqdm = tq.tqdm(data_loader, total=len(data_loader), desc=\"Testing\", position=0, leave=True)\n    for images in test_tqdm:\n        images = images.to(device)\n        preds.extend(model(images).detach().cpu().numpy())\n    return preds\n\n# effiecien  res 修改这里\nclass CassvaImgClassifier(nn.Module):\n    def __init__ (self, efnet_arch, n_class, pretrained=False):\n        super().__init__ ()\n        self.efnet_model = timm.create_model(efnet_arch, pretrained=pretrained)\n        efnet_features = self.efnet_model.classifier.in_features\n        self.efnet_model.classifier = nn.Linear(efnet_features, n_class)\n        #全连接层 5类\n        # self.resnet_model = timm.create_model(resnet_arch, pretrained=pretrained)\n        # resnet_features = self.resnet_model.fc.in_features\n        # self.resnet_model.classifier = nn.Linear(resnet_features, n_class)\n\n    def forward (self, x):\n        efnet_opt = self.efnet_model(x)\n        # resnet_opt = self.resnet_model(x)\n\n        # combined = efnet_opt*0.6 + resnet_opt*0.4\n        return efnet_opt\n\n\neffnet_preds = []\nfor effnet_model_name in CFG['effnet_models']:\n    print(\"Model: \", effnet_model_name)\n    effnet_model = torch.load('../input/hym1216/'+effnet_model_name, map_location=torch.device(CFG['device']))\n   #加载进来\n    with torch.no_grad():\n        for i in range(CFG['tta']):\n            effnet_preds += [inference(effnet_model, test_loader, CFG['device'])]\neffnet_preds = np.mean(effnet_preds, axis=0)\n\n\neffnet_outcomes = pd.concat([df['image_id'], pd.DataFrame(effnet_preds)], axis=1).sort_values(['image_id'])\n\n\nfinal_preds = (effnet_outcomes.drop('image_id', axis=1)).to_numpy()\nfinal_preds = softmax(final_preds).argmax(1)\n\naccuracy_score(final_preds, df['label'].values)\n\nsubmit = pd.DataFrame({'image_id': df['image_id'].values, 'label': final_preds})\nsubmit.to_csv('submission.csv', index=False)\n\n#输出提交","execution_count":null,"outputs":[]}],"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"}},"nbformat":4,"nbformat_minor":4}