{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/working/weights'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-27T04:51:27.970408Z","iopub.execute_input":"2022-05-27T04:51:27.970739Z","iopub.status.idle":"2022-05-27T04:51:28.003377Z","shell.execute_reply.started":"2022-05-27T04:51:27.970656Z","shell.execute_reply":"2022-05-27T04:51:28.00267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modules","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"..\")","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:29.336337Z","iopub.execute_input":"2022-05-27T04:51:29.336875Z","iopub.status.idle":"2022-05-27T04:51:29.340973Z","shell.execute_reply.started":"2022-05-27T04:51:29.336836Z","shell.execute_reply":"2022-05-27T04:51:29.340173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install Augmentor","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:29.71465Z","iopub.execute_input":"2022-05-27T04:51:29.715135Z","iopub.status.idle":"2022-05-27T04:51:39.150669Z","shell.execute_reply.started":"2022-05-27T04:51:29.715099Z","shell.execute_reply":"2022-05-27T04:51:39.149826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport copy\nimport random\nfrom functools import partial\nfrom collections import OrderedDict\nfrom typing import Optional, Callable\n\nimport torch\nimport torch.nn as nn\nfrom torch import Tensor\nfrom torch.nn import functional as F\n\nimport os\nimport argparse\n\nimport torch.optim as optim\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torchvision import transforms\nimport torch.optim.lr_scheduler as lr_scheduler\n\nfrom input.efficientnetmodel.model import efficientnet_b3 as create_model\nfrom input.efficientnetmodel.my_dataset import MyDataSet\nfrom input.efficientnetmodel.utils import read_split_data, train_one_epoch, evaluate\n\nimport matplotlib.pyplot as plt\nimport cv2\nimport shutil\nimport Augmentor","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:39.154183Z","iopub.execute_input":"2022-05-27T04:51:39.154405Z","iopub.status.idle":"2022-05-27T04:51:41.732392Z","shell.execute_reply.started":"2022-05-27T04:51:39.15438Z","shell.execute_reply":"2022-05-27T04:51:41.731667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Transform","metadata":{}},{"cell_type":"code","source":"# split1 = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/train.csv')\n# split1.head()","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.733774Z","iopub.execute_input":"2022-05-27T04:51:41.734019Z","iopub.status.idle":"2022-05-27T04:51:41.740508Z","shell.execute_reply.started":"2022-05-27T04:51:41.733986Z","shell.execute_reply":"2022-05-27T04:51:41.739496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if os.path.exists('/kaggle/cassava_train/'):\n#     shutil.rmtree('/kaggle/cassava_train/')\n# for i in range(5):\n#     if not os.path.exists('/kaggle/cassava_train/' + str(i)):\n#         os.makedirs('/kaggle/cassava_train/' + str(i))\n#         print('/kaggle/cassava_train/' + str(i))","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.742712Z","iopub.execute_input":"2022-05-27T04:51:41.74327Z","iopub.status.idle":"2022-05-27T04:51:41.751156Z","shell.execute_reply.started":"2022-05-27T04:51:41.743231Z","shell.execute_reply":"2022-05-27T04:51:41.750395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# for i in range(21397):\n#     if i % 1000 == 0:\n#         print(i/1000)\n#     img = cv2.imread('/kaggle/input/cassava-leaf-disease-classification/train_images/' \n#                      + split1['image_id'][i])\n#     cv2.imwrite('/kaggle/cassava_train/' + str(split1['label'][i]) + '/' + split1['image_id'][i], img)","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.753683Z","iopub.execute_input":"2022-05-27T04:51:41.754764Z","iopub.status.idle":"2022-05-27T04:51:41.763871Z","shell.execute_reply.started":"2022-05-27T04:51:41.754735Z","shell.execute_reply":"2022-05-27T04:51:41.763116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# image enhancement","metadata":{}},{"cell_type":"code","source":"# path_out = '/kaggle/working/train/'\n# if os.path.exists(path_out):\n#     shutil.rmtree(path_out)\n# os.makedirs(path_out + '0')\n# os.makedirs(path_out + '1')\n# os.makedirs(path_out + '2')\n# os.makedirs(path_out + '3')\n# os.makedirs(path_out + '4')\n\n# for dirname, _, filenames in os.walk('/kaggle/input/cassava-train/train/train'):\n#     for filename in filenames:\n#         path1 = os.path.join(dirname, filename)\n#         img1 = cv2.imread(path1)\n#         print(path_out + str(dirname[-1]) + '/' + filename)\n#         cv2.imwrite(path_out + str(dirname[-1]) + '/' + filename, img1)\n        \n#         count1 = 0\n#         count2 = 0\n#         count3 = 0\n#         count4 = 0\n        \n#         a = random.random()\n#         if a >= 0.5:\n#             count1 += 1\n#             img2 = img1.transpose((1, 0, 2))\n#             path2 = path_out + str(dirname[-1]) + '/' + '2_' + filename\n#             cv2.imwrite(path2, img2)\n#         b = random.random()\n#         if b >= 0.5:\n#             count2 += 1","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.765334Z","iopub.execute_input":"2022-05-27T04:51:41.765795Z","iopub.status.idle":"2022-05-27T04:51:41.773072Z","shell.execute_reply.started":"2022-05-27T04:51:41.765762Z","shell.execute_reply":"2022-05-27T04:51:41.772341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train","metadata":{}},{"cell_type":"code","source":"def main(args):\n    device = torch.device(args.device if torch.cuda.is_available() else \"cpu\")\n\n    print(args)\n    print('Start Tensorboard with \"tensorboard --logdir=runs\", view at http://localhost:6006/')\n    tb_writer = SummaryWriter()\n    if os.path.exists(\"./weights\") is False:\n        os.makedirs(\"./weights\")\n\n    train_images_path, train_images_label, val_images_path, val_images_label = read_split_data(args.data_path)\n\n    img_size = {\"B0\": 224,\n                \"B1\": 240,\n                \"B2\": 260,\n                \"B3\": 300,\n                \"B4\": 380,\n                \"B5\": 456,\n                \"B6\": 528,\n                \"B7\": 600}\n    num_model = \"B3\"\n\n    data_transform = {\n        \"train\": transforms.Compose([transforms.RandomResizedCrop(img_size[num_model]),\n                                     transforms.RandomHorizontalFlip(),\n                                     transforms.RandomVerticalFlip(),\n                                     transforms.ToTensor(),\n                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]),\n        \"val\": transforms.Compose([transforms.RandomResizedCrop(img_size[num_model]),\n                                     transforms.RandomHorizontalFlip(),\n                                     transforms.RandomVerticalFlip(),\n                                     transforms.ToTensor(),\n                                     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])}\n\n    # 实例化训练数据集\n    train_dataset = MyDataSet(images_path=train_images_path,\n                              images_class=train_images_label,\n                              transform=data_transform[\"train\"])\n\n    # 实例化验证数据集\n    val_dataset = MyDataSet(images_path=val_images_path,\n                            images_class=val_images_label,\n                            transform=data_transform[\"val\"])\n\n    batch_size = args.batch_size\n    nw = min([os.cpu_count(), batch_size if batch_size > 1 else 0, 8])  # number of workers\n    print('Using {} dataloader workers every process'.format(nw))\n    train_loader = torch.utils.data.DataLoader(train_dataset,\n                                               batch_size=batch_size,\n                                               shuffle=True,\n                                               pin_memory=True,\n                                               num_workers=nw,\n                                               collate_fn=train_dataset.collate_fn)\n\n    val_loader = torch.utils.data.DataLoader(val_dataset,\n                                             batch_size=batch_size,\n                                             shuffle=False,\n                                             pin_memory=True,\n                                             num_workers=nw,\n                                             collate_fn=val_dataset.collate_fn)\n\n    # 如果存在预训练权重则载入\n    model = create_model(num_classes=args.num_classes).to(device)\n    if args.weights != \"\":\n        if os.path.exists(args.weights):\n            weights_dict = torch.load(args.weights, map_location=device)\n            load_weights_dict = {k: v for k, v in weights_dict.items()\n                                 if model.state_dict()[k].numel() == v.numel()}\n            print(model.load_state_dict(load_weights_dict, strict=False))\n        else:\n            raise FileNotFoundError(\"not found weights file: {}\".format(args.weights))\n\n    # 是否冻结权重\n    if args.freeze_layers:\n        for name, para in model.named_parameters():\n            # 除最后一个卷积层和全连接层外，其他权重全部冻结\n            if (\"features.top\" not in name) and (\"classifier\" not in name):\n                para.requires_grad_(False)\n            else:\n                print(\"training {}\".format(name))\n\n    pg = [p for p in model.parameters() if p.requires_grad]\n    optimizer = optim.SGD(pg, lr=args.lr, momentum=0.9, weight_decay=1E-4)\n    # Scheduler https://arxiv.org/pdf/1812.01187.pdf\n    lf = lambda x: ((1 + math.cos(x * math.pi / args.epochs)) / 2) * (1 - args.lrf) + args.lrf  # cosine\n    scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lf)\n    \n    for epoch in range(args.epochs):\n        # train\n        mean_loss = train_one_epoch(model=model,\n                                    optimizer=optimizer,\n                                    data_loader=train_loader,\n                                    device=device,\n                                    epoch=epoch)\n\n        scheduler.step()\n\n        # validate\n        acc = evaluate(model=model,\n                       data_loader=val_loader,\n                       device=device)\n        print(\"[epoch {}] accuracy: {}\".format(epoch, round(acc, 4)))\n        tags = [\"loss\", \"accuracy\", \"learning_rate\"]\n        tb_writer.add_scalar(tags[0], mean_loss, epoch)\n        tb_writer.add_scalar(tags[1], acc, epoch)\n        tb_writer.add_scalar(tags[2], optimizer.param_groups[0][\"lr\"], epoch)\n\n        torch.save(model.state_dict(), \"./weights/model-{}.pth\".format(epoch))\n","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.774546Z","iopub.execute_input":"2022-05-27T04:51:41.775101Z","iopub.status.idle":"2022-05-27T04:51:41.798883Z","shell.execute_reply.started":"2022-05-27T04:51:41.775062Z","shell.execute_reply":"2022-05-27T04:51:41.798114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nparser = argparse.ArgumentParser()\nparser.add_argument('--num_classes', type=int, default=5)\nparser.add_argument('--epochs', type=int, default=30)\nparser.add_argument('--batch-size', type=int, default=16)\nparser.add_argument('--lr', type=float, default=0.01)\nparser.add_argument('--lrf', type=float, default=0.01)\n\n# 数据集所在根目录\n# https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz\nparser.add_argument('--data-path', type=str,\n                    default=\"/kaggle/input/cassava-train/train/train\")\n\n# download model weights\n# 链接: https://pan.baidu.com/s/1ouX0UmjCsmSx3ZrqXbowjw  密码: 090i\nparser.add_argument('--weights', type=str, default='/kaggle/input/efficientnet/efficientnetb3.pth',\n                    help='initial weights path')\nparser.add_argument('--freeze-layers', type=bool, default=False)\nparser.add_argument('--device', default='cuda:0', help='device id (i.e. 0 or 0,1 or cpu)')\n\nopt = parser.parse_known_args()[0]\n\nmain(opt)","metadata":{"execution":{"iopub.status.busy":"2022-05-27T04:51:41.799889Z","iopub.execute_input":"2022-05-27T04:51:41.801806Z","iopub.status.idle":"2022-05-27T07:08:46.611214Z","shell.execute_reply.started":"2022-05-27T04:51:41.80177Z","shell.execute_reply":"2022-05-27T07:08:46.610319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# test","metadata":{}},{"cell_type":"code","source":"import json\nimport torch\nfrom PIL import Image\nfrom torchvision import transforms\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-05-27T12:41:41.35986Z","iopub.execute_input":"2022-05-27T12:41:41.360596Z","iopub.status.idle":"2022-05-27T12:41:43.106578Z","shell.execute_reply.started":"2022-05-27T12:41:41.360505Z","shell.execute_reply":"2022-05-27T12:41:43.10587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main1(num_model):\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    img_size = {\"B0\": 224,\n                \"B1\": 240,\n                \"B2\": 260,\n                \"B3\": 300,\n                \"B4\": 380,\n                \"B5\": 456,\n                \"B6\": 528,\n                \"B7\": 600}\n    \n    data_transform = transforms.Compose(\n        [transforms.Resize(img_size[num_model]),\n         transforms.CenterCrop(img_size[num_model]),\n         transforms.ToTensor(),\n         transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])\n\n    # load image\n    img_path = \"/kaggle/input/cassava-leaf-disease-classification/test_images/2216849948.jpg\"\n    assert os.path.exists(img_path), \"file: '{}' dose not exist.\".format(img_path)\n    img = Image.open(img_path)\n    plt.imshow(img)\n    # [N, C, H, W]\n    img = data_transform(img)\n    # expand batch dimension\n    img = torch.unsqueeze(img, dim=0)\n\n    # read class_indict\n    json_path = './class_indices.json'\n    assert os.path.exists(json_path), \"file: '{}' dose not exist.\".format(json_path)\n\n    with open(json_path, \"r\") as f:\n        class_indict = json.load(f)\n\n    # create model\n    model = create_model(num_classes=5).to(device)\n    # load model weights\n    model_weight_path = \"./weights/model-29.pth\"\n    model.load_state_dict(torch.load(model_weight_path, map_location=device))\n    model.eval()\n    with torch.no_grad():\n        # predict class\n        output = torch.squeeze(model(img.to(device))).cpu()\n        predict = torch.softmax(output, dim=0)\n        predict_cla = torch.argmax(predict).numpy()\n\n    print_res = \"class: {}   prob: {:.3}\".format(class_indict[str(predict_cla)],\n                                                 predict[predict_cla].numpy())\n    plt.title(print_res)\n    for i in range(len(predict)):\n        print(\"class: {:10}   prob: {:.3}\".format(class_indict[str(i)],\n                                                  predict[i].numpy()))\n    plt.show()\nmain1('B3')\n# b2","metadata":{"execution":{"iopub.status.busy":"2022-05-27T12:41:45.880593Z","iopub.execute_input":"2022-05-27T12:41:45.881128Z","iopub.status.idle":"2022-05-27T12:41:46.049741Z","shell.execute_reply.started":"2022-05-27T12:41:45.881088Z","shell.execute_reply":"2022-05-27T12:41:46.048352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}