{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import os\nCFG = {\n    'model_arch': 'tf_efficientnet_b4_ns',\n    'img_size': 512,\n    'num_classes': 5,\n    'tta': 5,\n    'savecsv': \"tmpfilefold/pseudo_submission.csv\",\n    'srcImgPath':'../input/cassava-leaf-disease-classification/test_images',\n    'weights': [1],\n    'models':['EB4.pth.tar'],\n    'baseModelPath':'../input/pseudotraining',\n    'base_save_path':'./tmpfilefold/cassava-pseudo',\n    'saveScores':0.4,\n    \n}","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"if os.path.exists('./tmpfilefold'):\n    !rm -r ./tmpfilefold\nif not os.path.exists('./tmpfilefold'):\n    os.mkdir('./tmpfilefold')\nif not os.path.exists(CFG['base_save_path']):\n    os.mkdir(CFG['base_save_path'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\ndef inference_one_epoch(model, data_loader):\n    model.eval()\n    image_preds_all = []\n    pbar = tqdm(enumerate(data_loader), total=len(data_loader))\n    for i,(img,imgpath) in pbar:\n        img = img.cuda()\n        with torch.no_grad():\n            image_preds = model(img)\n            image_preds_all += [torch.softmax(image_preds, 1).detach().cpu().numpy()]\n    image_preds_all = np.concatenate(image_preds_all, axis=0)\n    return image_preds_all","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import sys\nimport cv2\nimport argparse\nimport pandas as pd\nimport os\nfrom torch.utils.data import Dataset,DataLoader\nimport albumentations as A\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nfrom albumentations.pytorch.transforms import ToTensor\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)\nfrom albumentations.pytorch import ToTensorV2\nimport torch\ndef test_transform(width = 512,height = 512):\n    return A.Compose([A.CLAHE(),\n                       A.Resize(height = height,width = width,p = 1.0),\n                       ToTensor(sigmoid=False,normalize={'mean':[0.485, 0.456, 0.406],'std':[0.229, 0.224, 0.225]})]\n                       )\n\ndef test_transform_v2(width = 512,height = 512):\n    return Compose([\n#             CenterCrop(width, height, p=1.),\n            Resize(width, height,p=1.0),\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        ], p=1.)\n\ndef test_transform_v3():\n    return Compose([\n            RandomResizedCrop(CFG['img_size'], CFG['img_size']),\n            Transpose(p=0.5),\n            HorizontalFlip(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        ], p=1.)\n\ndef test_transform_v4():\n    return Compose([\n            Resize(CFG['img_size'], CFG['img_size']),\n            Transpose(p=0.5),\n            HorizontalFlip(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        ], p=1.)\n\n\nclass LoadImagesAndLabels(Dataset):  # for training/testing\n    def __init__(self,baseImgPath = '', trans = None):\n        self.image_ids = []\n        self.baseImgPath = baseImgPath\n        for imgname in os.listdir(self.baseImgPath):\n            self.image_ids.append(os.path.join(self.baseImgPath,imgname))\n        self.trans = trans\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, index):\n        imgpath = self.image_ids[index]\n        img = cv2.imread(imgpath)\n        assert img is not None,imgpath\n        img = img[:, :, ::-1]\n        img = self.trans(image = img)['image']\n        imgname = imgpath[imgpath.rfind('/')+1:]\n        return img,imgname","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import sys\nsys.path.insert(0,'../input/pseudotraining/pytorch-image-models')\nfrom timm import create_model\nimport torch\nfrom itertools import compress\n\ndef makePseudolabel():\n    args = CFG\n    test = pd.DataFrame()\n    model = create_model(model_name = args['model_arch'],num_classes=args['num_classes'])\n    trans = test_transform_v4()\n    val_dataset = LoadImagesAndLabels(args['srcImgPath'],trans)\n    val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=60, shuffle=False,num_workers=12, pin_memory=True)\n    tst_preds = []\n    with torch.no_grad():\n        for i,model_name in enumerate(CFG['models']):\n            model_path = os.path.join(args['baseModelPath'],model_name)\n            print('model_path: ',model_path)\n            state_dict = torch.load(model_path,map_location='cpu')\n            model.load_state_dict(state_dict[\"state_dict\"], strict=True)\n            model.cuda()\n            for _ in range(CFG['tta']):\n                tst_preds += [CFG['weights'][i]/sum(CFG['weights'])/CFG['tta']*inference_one_epoch(model, val_loader)]\n    tst_preds = np.sum(tst_preds, axis=0)\n    tst_preds_max = np.max(tst_preds,axis = -1)\n    index = tst_preds_max > args['saveScores']\n    test['image_id'] = list(compress(list(os.listdir(args['srcImgPath'])), list(index)))\n    test['label'] = np.argmax(tst_preds, axis=1)[index]\n    test.head()\n    test.to_csv(args['savecsv'], index=False)\n    print('makePseudolabel Done!')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def coptfile():\n    import csv\n    import os\n    import shutil\n    args = CFG\n    test = pd.DataFrame()\n    csv_reader = csv.reader(open(args['savecsv']))\n    label_dict = {'0':'CBB', '1':'CBSD', '2':\"CGM\", '3':\"CMD\",'4':\"Healthy\"}\n    count = {'CBB':0,'CBSD':0,\"CGM\":0,\"CMD\":0,\"Healthy\":0}\n    for line in csv_reader:\n        if line[0].find('.jpg') == -1:\n            continue\n        img_name = line[0]\n        label = line[1]\n        img_short_dir = label_dict[label]\n        dst_path = os.path.join(args['base_save_path'],img_name) \n        src_path = os.path.join(args['srcImgPath'],img_name)\n        count[img_short_dir] += 1\n        shutil.copyfile(src_path,dst_path)\n    print('coptfile Done!')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import csv\nimport cv2\nimport sys\nimport numpy as np\nimport torch\nsys.path.insert(0,'../input/pseudotraining')\nargs = CFG\nnum_imgs = len(list(os.listdir(args['srcImgPath'])))\nif num_imgs >= 10:\n    makePseudolabel()\n    coptfile()\n    num_saveImgs = len(list(os.listdir(args['base_save_path'])))\n    if num_saveImgs >= 10:\n        csv.reader(open('./tmpfilefold/pseudo_submission.csv'))\n        torch.cuda.empty_cache()\n        !python ../input/pseudotraining/pytorch-image-models/train_pseudo.py --data './tmpfilefold/cassava-pseudo' --cvsPath './tmpfilefold/pseudo_submission.csv' --initial-checkpoint '../input/pseudotraining/EB4.pth.tar' --output 'tmpfilefold' --batch-size 10\nelse:\n    print('NO PSEUDO!')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.cuda.empty_cache()\nCFG = {\n    'model_arch': 'tf_efficientnet_b4_ns',\n    'img_size': 512,\n    'num_classes': 5,\n    'tta': 5,\n    'srcImgPath':'../input/cassava-leaf-disease-classification/test_images',\n    'weights': [1],\n    'pseudo_model_path': './tmpfilefold',\n    'base_model_path': './tmpfilefold/cassava-pseudo'\n}\nif os.path.exists(os.path.join(CFG['pseudo_model_path'],'model_best.pth.tar')):\n    CFG['models'] = ['model_best.pth.tar','EB4.pth.tar']\n    CFG['base_model_path'] = CFG['pseudo_model_path']\n    CFG['weights'] = [1,1] \nelse:\n    CFG['models'] = ['EB4.pth.tar']\n\nprint(CFG['models'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from timm import create_model\nimport torch\n\nargs = CFG\ntest = pd.DataFrame()\ntest['image_id'] = list(os.listdir(args['srcImgPath']))\nmodel = create_model(model_name = args['model_arch'],num_classes=args['num_classes'])\ntrans = test_transform_v4()\nval_dataset = LoadImagesAndLabels(args['srcImgPath'],trans)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size=60, shuffle=False,num_workers=16, pin_memory=True)\ntst_preds = []\nwith torch.no_grad():\n    for i,model_name in enumerate(CFG['models']):\n        if model_name.find('model_best') != -1:\n            model_path = os.path.join(CFG['base_model_path'],model_name)\n        else:\n            model_path = os.path.join('../input/pseudotraining',model_name)\n        print('model_path: ',model_path)\n        state_dict = torch.load(model_path,map_location='cpu')\n        model.load_state_dict(state_dict[\"state_dict\"], strict=True)\n        model.cuda()\n        for _ in range(CFG['tta']):\n            tst_preds += [CFG['weights'][i]/sum(CFG['weights'])/CFG['tta']*inference_one_epoch(model, val_loader)]\ntst_preds = np.sum(tst_preds, axis=0)\ntest['label'] = np.argmax(tst_preds, axis=1)\ntest.head()\ntest.to_csv('submission.csv', index=False)\nprint('Done!')\nif os.path.exists('./tmpfilefold'):\n    !rm -r ./tmpfilefold","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}