{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"DEBUG = False","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# loading data"},{"metadata":{"trusted":true},"cell_type":"code","source":"from imgaug import augmenters as iaa\nimport numpy as np\nfrom torchvision import transforms\nimport torch\n\ntrans = transforms.Compose(\n    [\n        transforms.ToTensor()]\n)\n\n\ndef argument(image):\n    seq = iaa.Sequential(\n        [\n            iaa.Crop(px=(0, 16)),\n            iaa.Fliplr(0.5),\n            iaa.Flipud(.3),\n            iaa.CropAndPad(percent=(-0.25, 0.25)),\n            # 随机裁剪\n            iaa.CoarseDropout(5 * 0.01, size_percent=0.04, per_channel=0.4),\n            # 通道增强器\n            iaa.WithChannels(0, iaa.Add((-100, 100))),\n            iaa.WithChannels(1, iaa.Add((-100, 100))),\n            iaa.WithChannels(2, iaa.Add((-100, 100))),\n            # 高斯模糊\n            iaa.GaussianBlur((0, 1.)),\n            iaa.AverageBlur(k=((5, 11), (1, 3))),\n        ]\n    )\n\n    return seq.augment_image(image)\n\n\ndef rand_bbox(size, lam):\n    W = size[2]\n    H = size[3]\n\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = np.int(W * cut_rat)\n    cut_h = np.int(H * cut_rat)\n\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, W)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, W)\n\n    return bbx1, bby1, bbx2, bby2\n\n\ndef mixup(image1, image2, label1, label2, clz=5):\n    label1 = torch.tensor(label1)\n    label2 = torch.tensor(label2)\n    image = image1/2 + image2/2\n    label1 = torch.nn.functional.one_hot(label1, clz)\n    label2 = torch.nn.functional.one_hot(label2, clz)\n    label = label1/2 + label2/2\n    return image, label\n\n\nclass Cutout(object):\n    \"\"\"Randomly mask out one or more patches from an image.\n    Args:\n        n_holes (int): Number of patches to cut out of each image.\n        length (int): The length (in pixels) of each square patch.\n    \"\"\"\n    def __init__(self, n_holes, length):\n        self.n_holes = n_holes\n        self.length = length\n\n    def __call__(self, img):\n        \"\"\"\n        Args:\n            img (Tensor): Tensor image of size (C, H, W).\n        Returns:\n            Tensor: Image with n_holes of dimension length x length cut out of it.\n        \"\"\"\n        h = img.size(1)\n        w = img.size(2)\n\n        mask = np.ones((h, w), np.float32)\n\n        for n in range(self.n_holes):\n            y = np.random.randint(h)\n            x = np.random.randint(w)\n\n            y1 = np.clip(y - self.length // 2, 0, h)\n            y2 = np.clip(y + self.length // 2, 0, h)\n            x1 = np.clip(x - self.length // 2, 0, w)\n            x2 = np.clip(x + self.length // 2, 0, w)\n\n            mask[y1: y2, x1: x2] = 0.\n\n        mask = torch.from_numpy(mask)\n        mask = mask.expand_as(img)\n        img = img * mask\n\n        return img","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# data_generator"},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport json\nimport cv2\nimport pandas as pd\nfrom torch.utils.data import DataLoader, Dataset\nimport torch\nimport seaborn\n\nBASE_DIR = r'/kaggle/input'\n\n\nclass TrainDataSet(Dataset):\n    \"\"\"\n    获取了图像的csv之后将结果转换为带表的list\n    \"\"\"\n    def __init__(self, root_path, num_class):\n        super(TrainDataSet, self).__init__()\n        with open(os.path.join(root_path, 'label_num_to_disease_map.json')) as file:\n            map_class = json.loads(file.read())\n            map_class = {int(k): v for k, v in map_class.items()}\n\n        self.root_path = root_path\n        self.df = pd.read_csv(os.path.join(root_path, 'train.csv'))[:401 if DEBUG else 20000]\n        self.df['class_name'] = self.df['label'].map(map_class)\n        self.num_class = num_class\n\n    def __getitem__(self, index: int):\n        \"\"\"\n        应用图像转换算法，尝试将mixup等高阶算法放入图像转换算法\n        :param index:选择图像的位置\n        :return:tensor化之后的图像，one-hot之后的label\n        \"\"\"\n        img_name = self.df.iloc[[index], [0]].values[0][0]\n        target = self.df.iloc[[index], [1]].values[0][0]\n        img_path = os.path.join(self.root_path, 'train_images', img_name)\n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        img = argument(img)\n        tensor = trans(img)\n        target = torch.nn.functional.one_hot(torch.tensor(target), self.num_class)\n        return tensor, target\n\n    def __len__(self) -> int:\n        return len(self.df)\n\n\nclass ValueDataset(Dataset):\n    def __init__(self, root_path, num_class):\n        super(ValueDataset, self).__init__()\n        with open(os.path.join(root_path, 'label_num_to_disease_map.json')) as file:\n            map_class = json.loads(file.read())\n            map_class = {int(k): v for k, v in map_class.items()}\n\n        self.root_path = root_path\n        self.df = pd.read_csv(os.path.join(root_path, 'train.csv'))[20000:]\n        self.df['class_name'] = self.df['label'].map(map_class)\n        self.num_class = num_class\n\n    def __getitem__(self, index: int):\n        \"\"\"\n        :param index:选择图像的位置\n        :return:tensor化之后的图像，one-hot之后的label\n        \"\"\"\n        img_name = self.df.iloc[[index], [0]].values[0][0]\n        target = self.df.iloc[[index], [1]].values[0][0]\n        img_path = os.path.join(self.root_path, 'train_images', img_name)\n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        tensor = trans(img)\n        target = torch.nn.functional.one_hot(torch.tensor(target), self.num_class)\n        return tensor, target\n\n    def __len__(self) -> int:\n        return len(self.df)\n\n\ndef get_train_loader(root_path, batch_size, num_class):\n    train_ds = TrainDataSet(root_path=root_path, num_class=num_class)\n    return DataLoader(dataset=train_ds, shuffle=False, num_workers=8, batch_size=batch_size)\n\n\ndef get_value_loader(root_path, batch_size, num_class):\n    value_ds = ValueDataset(root_path=root_path, num_class=num_class)\n    return DataLoader(dataset=value_ds, shuffle=False, num_workers=8, batch_size=batch_size)\n\n\ndef get_loader(root_path, batch_size, num_class):\n    train_loader = get_train_loader(root_path, batch_size, num_class)\n    value_loader = get_value_loader(root_path, batch_size, num_class)\n    loader = {\n        'train': train_loader,\n        'value': value_loader,\n    }\n    return loader\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# network"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom collections import OrderedDict\nfrom torchvision.models.vgg import vgg16_bn \ntorch.cuda.is_available()\n\nfc_layers = OrderedDict([\n    ('fc1', nn.Linear(1000, 100)),\n    ('relu1', nn.ReLU()),\n    ('dropout', nn.Dropout(.5)),\n    ('relu2', nn.ReLU()),\n    ('fc2', nn.Linear(100, 5)),\n    ('softmax', nn.Softmax(dim=1))\n])\n\n\nclass ResNet(nn.Module):\n    def __init__(self):\n        super(ResNet, self).__init__()\n        resnet = resnet34(pretrained=True)\n        self.backbone = resnet\n        self.fc_layers = nn.Sequential(fc_layers)\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.fc_layers(x)\n        return x\n\nclass Vgg16(nn.Module):\n    def __init__(self):\n        super(Vgg16, self).__init__()\n        vgg = vgg16_bn(pretrained=True)\n        self.backbone = vgg\n        self.fc_layers = nn.Sequential(fc_layers)\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.fc_layers(x)\n        return x\n    \ndef get_network():\n    return Vgg16()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# static args"},{"metadata":{"trusted":true},"cell_type":"code","source":"import argparse\nimport os \n\ndef get_args():\n    parser = argparse.ArgumentParser()\n    parser.add_argument('-B', '--base_dir',  type=str, default=r'../input/cassava-leaf-disease-classification')\n    parser.add_argument('-BB', '--backbone_lr',  type=float, default=5e-5)\n    parser.add_argument('-R', '--lr',  type=float, default=5e-3)\n    parser.add_argument('-T', '--is_train',  action='store_false', default=True)\n    parser.add_argument('-BS', '--batch_size',  type=int, default=30)\n    parser.add_argument('-N', '--num_class',  type=int, default=5)\n    parser.add_argument('-E', '--epochs',  type=int, default=200)\n    parser.add_argument('-S', '--save_dir',  type=str, default=os.path.join(os.getcwd(), 'save_dir'))\n    args = parser.parse_known_args()[0]\n    return args\n\n\nstart_epoch = 0\nacc_list = []\nepoch_list = []\nargs =get_args()\nsave_dir = args.save_dir\nmodel = get_network().cuda()\nbase_params = list(map(id, model.backbone.parameters()))\nlogits_params = filter(lambda p: id(p) not in base_params, model.parameters())\nparams = [\n    {\"params\": model.backbone.parameters(), \"lr\": args.backbone_lr},\n    {\"params\": logits_params, \"lr\": args.lr},\n]\nloss_fn = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(params=params)\nloader = get_loader(args.base_dir, args.batch_size, args.num_class)\ntrain_loader = loader['train']\nvalue_loader = loader['value']","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# train"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\n\n\ndef train(loader, optimizer, loss_fn, epoch):\n\n    train_loader = loader\n    correct = 0.\n    total = 0.\n    loss_sum = 0.\n    steps = len(loader)\n    for step, (image, label) in enumerate(train_loader):\n        optimizer.zero_grad()\n        image = image.cuda()\n        label = label.cuda()\n        logic = model(image)\n        label = torch.argmax(label, -1).long()\n        loss = loss_fn(logic, label)\n        logic = torch.argmax(logic, -1)\n        loss.backward()\n        optimizer.step()\n        # print(torch.argmax(label, -1).item(), torch.argmax(logic, -1).item())\n\n        correct += (logic == label).sum().float()\n        total += len(label)\n        loss_sum += loss.item()\n\n        if step % 10 == 0 and step != 0:\n            acc = (correct / total).item()\n            print('[train][epoch: {:5d}][step:{:5d}/{}] ------------acc:{:3.3f}------loss:{:.7f}------------------'.format(epoch, step,steps, acc, loss_sum))\n\n            correct = 0.\n            total = 0.\n            loss_sum = 0.\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install ttach","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# value"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport ttach as tta\nimport matplotlib.pyplot as plt\n\n\ndef value(loader, batch_size, epoch):\n    tta_model = tta.ClassificationTTAWrapper(model,\n                                             tta.aliases.five_crop_transform(crop_height=256, crop_width=256),\n                                             merge_mode='mean')\n    # tta算法将结果显示\n    value_loader = loader\n    correct = 0.\n    total = 0.\n    steps = len(loader)\n    for step, (image, label) in enumerate(value_loader):\n        image = image.cuda()\n        label = label.cuda()\n        logic = tta_model(image)\n        label = torch.argmax(label, -1)\n        logic = torch.argmax(logic, -1)\n        correct += (logic == label).sum().float()\n        total += len(label)\n\n        if step % 10 == 0 and step != 0:\n            print('[test] ----------loading {}/{}---------------'.format(step, steps))\n    acc = (correct / total).item()\n    return acc\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# main"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport os\nimport matplotlib.pyplot as plt\n\nimport os\n\n\ndef load():\n    \"\"\"\n    save items:\n    'ckpt': checkpoint of last epoch,\n    'start_epoch': the latest epoch,\n    'loss_fn': loss items,\n    :return: None\n    \"\"\"\n    global loss_fn, start_epoch, epoch_list, acc_list\n    save_path = os.path.join(save_dir, 'classify_latest_v2.pth')\n    if os.path.exists(save_path):\n        ckpt = torch.load(save_path)\n        start_epoch = ckpt['epoch']\n        loss_fn.load_state_dict(ckpt['loss_fn'])\n        model.load_state_dict(ckpt['model_state_dict'])\n        optimizer.load_state_dict(ckpt['optimizer'])\n        epoch_list = ckpt['epoch_list']\n        acc_list = ckpt['acc_list']\n\n        for state in optimizer.state.values():\n            for k, v in state.items():\n                if torch.is_tensor(v):\n                    state[k] = v.cuda()\n\n\ndef save():\n\n    # saving network\n    ckpt = {\n        'epoch': epoch+1,\n        'model_state_dict': model.state_dict(),\n        'optimizer': optimizer.state_dict(),\n        'loss_fn': loss_fn.state_dict(),\n        'acc_list': acc_list,\n        'epoch_list': epoch_list,\n    }\n    if not os.path.exists(save_dir):\n        os.mkdir(save_dir)\n    save_path = os.path.join(save_dir, 'classify_latest_v2.pth')\n    torch.save(ckpt, save_path)\n\n\nif __name__ == '__main__':\n    \n    load()\n\n    for epoch in range(start_epoch, args.epochs):\n        model.train()\n        train(train_loader, optimizer, loss_fn, epoch) # reset model weight\n        \n        model.eval()\n        accur = value(value_loader, args.batch_size, epoch) # get accury \n        epoch_list.append(epoch)\n        acc_list.append(accur)\n        \n        \n        plt.figure()\n        plt.plot(epoch_list[:20], acc_list[:20], mec='r', mfc='w', ms=10, marker='x')\n        plt.savefig(r'value_acc.png')\n        \n        if len(acc_list) <= 2 or (len(acc_list) > 2 and acc_list[-1] > acc_list[-2]):\n            print(accur)\n            save()\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}