{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport random\nimport cv2\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport os\nimport torchvision.models as models\nimport torchvision.transforms as T\nfrom PIL import Image","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        \n        check_point = models.resnext50_32x4d().state_dict()\n        del check_point[\"fc.weight\"]\n        del check_point[\"fc.bias\"]\n        \n        self.backbone = models.resnext50_32x4d(num_classes=num_classes)\n        self.backbone.load_state_dict(check_point, strict=False)\n        self.loss = nn.BCEWithLogitsLoss()\n        \n    def forward(self, images, labels=None):\n        if self.training:\n            y = self.backbone(images)\n            loss = self.loss(y, labels)\n            return loss\n        else:\n            pred = self.backbone(images)\n            logits = pred.sigmoid()\n            batched_labels = []\n            for item in logits:\n                batched_labels.append(torch.where(item > 0.5)[0].tolist())\n            return batched_labels","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset:\n    def __init__(self):\n        self.root = \"/kaggle/input/plant-pathology-2021-fgvc8/test_images\"\n        self.label_map = {\n            'complex':[\"疑难复杂\", 0], \n            'rust':[\"锈菌，生锈\", 1], \n            'scab': [\"疮痂病，斑点病\", 2], \n            'frog_eye_leaf_spot': [\"青蛙眼叶斑\", 3], \n            'healthy': [\"健康的\", 4], \n            'powdery_mildew': [\"白粉病\", 5]\n        }\n        self.index_to_name = {\n            self.label_map[key][1]: key for key in self.label_map\n        }\n        self.num_classes = len(self.label_map)\n        self.files = []\n        for dirname, _, filenames in os.walk(self.root):\n            for filename in filenames:\n                self.files.append([os.path.join(dirname, filename), filename])\n                \n        self.trans = T.Compose([\n            #T.RandomHorizontalFlip(),\n            T.Resize((300, 300)),\n            T.ToTensor(),\n            T.Normalize(mean=0.5, std=1.0)\n        ])\n        \n    def __getitem__(self, index):\n        file, name = self.files[index]\n        image = Image.open(file)\n        image = self.trans(image)\n        return image, name\n        \n    def __len__(self):\n        return len(self.files)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\ndevice = \"cuda:0\"\ndataloader = torch.utils.data.DataLoader(Dataset(), batch_size=batch_size, shuffle=True, pin_memory=True, num_workers=3)\nmodel = Model(dataloader.dataset.num_classes)\ncheckpoint = torch.load(\"../input/model1pth/030.pth\", map_location=\"cpu\")\nmodel.load_state_dict(checkpoint)\nmodel.to(device)\n_ = model.eval()\n\nall_predict = []\nfor images, names in dataloader:\n    images = images.to(device)\n    batched_labels = model(images)\n    for name, labels in zip(names, batched_labels):\n        all_predict.append([name, \" \".join([dataloader.dataset.index_to_name[index] for index in labels])])\n\ndata = pd.DataFrame(all_predict, columns=(\"image\", \"labels\"))\ndata.to_csv(\"submission.csv\", index=False)\ndata","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}