{"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":"markdown","source":"# About this notebook  \n\nTBD...\n","metadata":{}},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = \"./\"\nMODEL_DIR = \"../input/cassava-model/\"\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_PATH = \"../input/cassava-leaf-disease-classification/train_images\"\nTEST_PATH = \"../input/cassava-leaf-disease-classification/test_images\"","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:04.998353Z","iopub.execute_input":"2022-04-25T10:41:04.998811Z","iopub.status.idle":"2022-04-25T10:41:05.006866Z","shell.execute_reply.started":"2022-04-25T10:41:04.998777Z","shell.execute_reply":"2022-04-25T10:41:05.006025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    debug = False\n    num_workers = 4\n    models = [\n        # \"tf_efficientnet_b3_ns\",\n        \"tf_efficientnet_b4_ns\",\n        \"vit_base_patch16_384\",\n        # \"deit_base_patch16_384\",\n        \"seresnext50_32x4d\",\n    ]\n    size = {\n        #\"tf_efficientnet_b3_ns\": 512,\n        \"tf_efficientnet_b4_ns\": 512,\n        \"vit_base_patch16_384\": 384,\n        #\"deit_base_patch16_384\": 384,\n        \"seresnext50_32x4d\": 512,\n    }\n    batch_size = 128\n    seed = 4021\n    target_size = 5\n    target_col = \"label\"\n    n_fold = 5\n    trn_fold = {  # [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]\n        \"tf_efficientnet_b3_ns\": {\n            \"best\": [0, 1, 2, 3, 4],\n            \"final\": [],\n        },\n        \"tf_efficientnet_b4_ns\": {\n            \"best\": [0, 1, 2, 3, 4],\n            \"final\": [],\n        },\n        \"vit_base_patch16_384\": {\"best\": [0, 1, 2, 3, 4], \"final\": []},\n        #\"deit_base_patch16_384\": {\"best\": [0, 1, 2, 3, 4], \"final\": []},\n        \"seresnext50_32x4d\": {\"best\": [5, 6, 7, 8, 9], \"final\": []},\n    }\n    data_parallel = {\n        #\"tf_efficientnet_b3_ns\": False,\n        \"tf_efficientnet_b4_ns\": True,  # True,\n        \"vit_base_patch16_384\": False,\n        #\"deit_base_patch16_384\": False,\n        \"seresnext50_32x4d\": False,\n    }\n    transform = {\n        # \"tf_efficientnet_b3_ns\": None,\n        \"tf_efficientnet_b4_ns\": \"rotate\",\n        \"vit_base_patch16_384\": \"rotate\",\n        # \"deit_base_patch16_384\": None,\n        \"seresnext50_32x4d\": \"rotate\",\n    }\n    weight = {\n        # \"tf_efficientnet_b3_ns\": None,\n        \"tf_efficientnet_b4_ns\": 1,\n        \"vit_base_patch16_384\": 1,\n        # \"deit_base_patch16_384\": None,\n        \"seresnext50_32x4d\": 1,\n    }\n    tta = 10  # 1: no TTA, >1: TTA\n    no_tta_weight = tta - 1\n    train = False\n    inference = True\n    \ntta_weight_sum = CFG.no_tta_weight + (CFG.tta - 1)\nweight_sum = sum([CFG.weight[model] for model in CFG.models]) * tta_weight_sum","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:05.00874Z","iopub.execute_input":"2022-04-25T10:41:05.009428Z","iopub.status.idle":"2022-04-25T10:41:05.025057Z","shell.execute_reply.started":"2022-04-25T10:41:05.009358Z","shell.execute_reply":"2022-04-25T10:41:05.024318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\n\nsys.path.append(\"../input/pytorch-image-models/pytorch-image-models-master\")\n\nimport math\nimport os\nimport random\nimport shutil\nimport time\nimport warnings\nfrom collections import Counter, defaultdict\nfrom contextlib import contextmanager\nfrom functools import partial\nfrom pathlib import Path\nimport os\nimport tensorflow as tf\nimport tensorflow_hub as hub\nimport numpy as np\nimport pandas as pd\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport scipy as sp\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom albumentations import (\n    CenterCrop,\n    CoarseDropout,\n    Compose,\n    Cutout,\n    HorizontalFlip,\n    HueSaturationValue,\n    IAAAdditiveGaussianNoise,\n    ImageOnlyTransform,\n    Normalize,\n    OneOf,\n    RandomBrightness,\n    RandomBrightnessContrast,\n    RandomContrast,\n    RandomCrop,\n    RandomResizedCrop,\n    Resize,\n    Rotate,\n    ShiftScaleRotate,\n    Transpose,\n    VerticalFlip,\n)\nfrom albumentations.pytorch import ToTensorV2\nfrom PIL import Image\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.nn.parameter import Parameter\nfrom torch.optim import SGD, Adam\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, CosineAnnealingWarmRestarts, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:05.026985Z","iopub.execute_input":"2022-04-25T10:41:05.028675Z","iopub.status.idle":"2022-04-25T10:41:13.945793Z","shell.execute_reply.started":"2022-04-25T10:41:05.028644Z","shell.execute_reply":"2022-04-25T10:41:13.94486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f\"[{name}] start\")\n    yield\n    LOGGER.info(f\"[{name}] done in {time.time() - t0:.0f} s.\")\n\n\ndef init_logger(log_file=OUTPUT_DIR + \"inference.log\"):\n    from logging import INFO, FileHandler, Formatter, StreamHandler, getLogger\n\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n\nseed_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:13.948381Z","iopub.execute_input":"2022-04-25T10:41:13.948768Z","iopub.status.idle":"2022-04-25T10:41:13.964977Z","shell.execute_reply.started":"2022-04-25T10:41:13.94871Z","shell.execute_reply":"2022-04-25T10:41:13.964279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"test = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:13.96771Z","iopub.execute_input":"2022-04-25T10:41:13.96796Z","iopub.status.idle":"2022-04-25T10:41:14.000372Z","shell.execute_reply.started":"2022-04-25T10:41:13.967936Z","shell.execute_reply":"2022-04-25T10:41:13.999452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df[\"image_id\"].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f\"{TEST_PATH}/{file_name}\"\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented[\"image\"]\n        return image","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:14.001881Z","iopub.execute_input":"2022-04-25T10:41:14.002205Z","iopub.status.idle":"2022-04-25T10:41:14.010971Z","shell.execute_reply.started":"2022-04-25T10:41:14.002171Z","shell.execute_reply":"2022-04-25T10:41:14.01006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data, size):\n\n    if data == \"train\":\n        return Compose(\n            [\n                # Resize(size, size),\n                RandomResizedCrop(size, size),\n                Transpose(p=0.5),\n                HorizontalFlip(p=0.5),\n                VerticalFlip(p=0.5),\n                ShiftScaleRotate(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                CoarseDropout(p=0.5),\n                Cutout(p=0.5),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"valid\":\n        return Compose(\n            [\n                Resize(size, size),\n                CenterCrop(size, size),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"simple\":\n        return Compose(\n            [\n                # Resize(size, size),\n                RandomResizedCrop(size, size),\n                # Transpose(p=0.5),\n                # HorizontalFlip(p=0.5),\n                # VerticalFlip(p=0.5),\n                # ShiftScaleRotate(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                # CoarseDropout(p=0.5),\n                # Cutout(p=0.5),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"rotate\":\n        return Compose(\n            [\n                # Resize(size, size),\n                RandomResizedCrop(size, size),\n                Transpose(p=0.5),\n                HorizontalFlip(p=0.5),\n                VerticalFlip(p=0.5),\n                # ShiftScaleRotate(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                # CoarseDropout(p=0.5),\n                # Cutout(p=0.5),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:14.012554Z","iopub.execute_input":"2022-04-25T10:41:14.013153Z","iopub.status.idle":"2022-04-25T10:41:14.030101Z","shell.execute_reply.started":"2022-04-25T10:41:14.013116Z","shell.execute_reply":"2022-04-25T10:41:14.028841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CassvaImgClassifier(nn.Module):\n    def __init__(self, model_name=\"resnext50_32x4d\", pretrained=False):\n        super().__init__()\n\n        if model_name == \"deit_base_patch16_384\":\n            # self.model = torch.hub.load(\"facebookresearch/deit:main\", model_name, pretrained=pretrained)\n            self.model = torch.hub.load(\"../input/fair-deit\", model_name, pretrained=pretrained, source=\"local\")\n            n_features = self.model.head.in_features\n            self.model.head = nn.Linear(n_features, CFG.target_size)\n\n        else:\n            self.model = timm.create_model(model_name, pretrained=pretrained)\n\n            if \"resnext50_32x4d\" in model_name:\n                n_features = self.model.fc.in_features\n                self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n            elif model_name.startswith(\"tf_efficientnet\"):\n                n_features = self.model.classifier.in_features\n                self.model.classifier = nn.Linear(n_features, CFG.target_size)\n\n            elif model_name.startswith(\"vit_\"):\n                n_features = self.model.head.in_features\n                self.model.head = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:14.031669Z","iopub.execute_input":"2022-04-25T10:41:14.031966Z","iopub.status.idle":"2022-04-25T10:41:14.045117Z","shell.execute_reply.started":"2022-04-25T10:41:14.031937Z","shell.execute_reply":"2022-04-25T10:41:14.044189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\ndef inference(model, states, test_loader, device, data_parallel):\n    model.to(device)\n\n    # Use multi GPU\n    if device == torch.device(\"cuda\") and data_parallel:\n        model = torch.nn.DataParallel(model)  # make parallel\n        # torch.backends.cudnn.benchmark=True\n\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state[\"model\"])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.softmax(1).to(\"cpu\").numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:14.046433Z","iopub.execute_input":"2022-04-25T10:41:14.046894Z","iopub.status.idle":"2022-04-25T10:41:14.058337Z","shell.execute_reply.started":"2022-04-25T10:41:14.046857Z","shell.execute_reply":"2022-04-25T10:41:14.057419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# inference","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\npredictions = None\nfor model_name in CFG.models:\n    for i in range(CFG.tta):\n        model = CassvaImgClassifier(model_name, pretrained=False)\n        states = []\n        for saved_model in [\"best\", \"final\"]:\n            if CFG.trn_fold[model_name][saved_model] != []:\n                LOGGER.info(\n                    f\"========== Model: {model_name}, TTA: {i}, Saved: {saved_model}, Fold: {CFG.trn_fold[model_name][saved_model]} ==========\"\n                )\n                states += [\n                    torch.load(MODEL_DIR + f\"{model_name}_fold{fold}_{saved_model}.pth\")\n                    for fold in CFG.trn_fold[model_name][saved_model]\n                ]\n\n        if i == 0:  # no TTA\n            test_dataset = TestDataset(test, transform=get_transforms(data=\"valid\", size=CFG.size[model_name]))\n            tta_weight = CFG.no_tta_weight\n        else:\n            test_dataset = TestDataset(\n                test, transform=get_transforms(data=CFG.transform[model_name], size=CFG.size[model_name])\n            )\n            tta_weight = 1\n\n        test_loader = DataLoader(\n            test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True\n        )\n\n        inf = inference(model, states, test_loader, device, CFG.data_parallel[model_name])\n        LOGGER.info(f\"Inference example: {inf[0]}\")\n\n        if predictions is None:\n            predictions = inf[np.newaxis] * CFG.weight[model_name] * tta_weight\n        else:\n            predictions = np.append(predictions, inf[np.newaxis] * CFG.weight[model_name] * tta_weight, axis=0)\n\nsub = np.sum(predictions, axis=0) / weight_sum\nLOGGER.info(f\"========== Overall ==========\")\nLOGGER.info(f\"Submission example: {sub[0]}\")","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:41:14.06006Z","iopub.execute_input":"2022-04-25T10:41:14.060317Z","iopub.status.idle":"2022-04-25T10:42:55.553833Z","shell.execute_reply.started":"2022-04-25T10:41:14.060286Z","shell.execute_reply":"2022-04-25T10:42:55.552785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Pipeline 2\n","metadata":{}},{"cell_type":"code","source":"!cp -R ../input/cassavalayer/ /kaggle/working/cassava-layer/","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:42:55.555868Z","iopub.execute_input":"2022-04-25T10:42:55.556289Z","iopub.status.idle":"2022-04-25T10:42:56.654803Z","shell.execute_reply.started":"2022-04-25T10:42:55.556242Z","shell.execute_reply":"2022-04-25T10:42:56.653807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"TFHUB_CACHE_DIR\"] = \"/kaggle/working/cassava-layer/\"","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:42:56.656795Z","iopub.execute_input":"2022-04-25T10:42:56.657365Z","iopub.status.idle":"2022-04-25T10:42:56.662344Z","shell.execute_reply.started":"2022-04-25T10:42:56.657321Z","shell.execute_reply":"2022-04-25T10:42:56.661486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow.compat.v2 as tf\nimport tensorflow_hub as hub\ncassava = hub.KerasLayer(\"../input/cropnet-classifier-cassava-disease-v1-2\")","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:42:56.663959Z","iopub.execute_input":"2022-04-25T10:42:56.664639Z","iopub.status.idle":"2022-04-25T10:43:00.095273Z","shell.execute_reply.started":"2022-04-25T10:42:56.664598Z","shell.execute_reply":"2022-04-25T10:43:00.094488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = tf.keras.Sequential([tf.keras.Input(shape=(224,224,3)), cassava])\nmodel.load_weights(\"../input/cassavamodelg/cassava_model.h5\")","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:00.097082Z","iopub.execute_input":"2022-04-25T10:43:00.097452Z","iopub.status.idle":"2022-04-25T10:43:01.461671Z","shell.execute_reply.started":"2022-04-25T10:43:00.097415Z","shell.execute_reply":"2022-04-25T10:43:01.460644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\nBATCH_SIZE = 128\nIMAGE_SIZE = [512, 512]","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:01.469274Z","iopub.execute_input":"2022-04-25T10:43:01.476652Z","iopub.status.idle":"2022-04-25T10:43:01.482008Z","shell.execute_reply.started":"2022-04-25T10:43:01.476605Z","shell.execute_reply":"2022-04-25T10:43:01.481207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _parse_function(proto):\n    # feature_description needs to be defined since datasets use graph-execution\n    # - its used to build their shape and type signature\n    feature_description = {\n        'image': tf.io.FixedLenFeature([], tf.string, default_value=''),\n        'image_name': tf.io.FixedLenFeature([], tf.string, default_value=''),\n        'target': tf.io.FixedLenFeature([], tf.int64, default_value=-1)\n    }\n\n    parsed_features = tf.io.parse_single_example(proto, feature_description)\n    image = tf.image.decode_jpeg(parsed_features['image'], channels=3)\n    image = tf.cast(image, tf.float32) # :: [0.0, 255.0]\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    target = tf.one_hot(parsed_features['target'], depth=5)\n    image_id = parsed_features['image_name']\n    return image, target, image_id\n\ndef _preprocess_fn(image, label, image_id):\n    image = image / 255.0\n    image = tf.image.resize(image, (224, 224))\n    label = tf.concat([label, [0]], axis=0)\n    return image, label, image_id\n\ndef load_dataset(tfrecords_fnames):\n    raw_ds = tf.data.TFRecordDataset(tfrecords_fnames, num_parallel_reads=AUTO)\n    parsed_ds = raw_ds.map(_parse_function, num_parallel_calls=AUTO)\n    parsed_ds = parsed_ds.map(_preprocess_fn, num_parallel_calls=AUTO)\n    return parsed_ds\n\ndef build_valid_ds(valid_fnames):\n    ds = load_dataset(valid_fnames)\n    ds = ds.batch(BATCH_SIZE).prefetch(AUTO)\n    return ds","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:01.483507Z","iopub.execute_input":"2022-04-25T10:43:01.484168Z","iopub.status.idle":"2022-04-25T10:43:01.52311Z","shell.execute_reply.started":"2022-04-25T10:43:01.484124Z","shell.execute_reply":"2022-04-25T10:43:01.522186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_PATH = '../input/cassava-leaf-disease-classification/test_tfrecords/'\nvalid_fnames = [TEST_PATH + fname for fname in os.listdir(TEST_PATH)]\ntest_ds = build_valid_ds(valid_fnames)\npreds = model.predict(test_ds)\nlabels = tf.argmax(preds, axis=-1)\nlabels = labels.numpy()\ntest_ds = build_valid_ds(valid_fnames)\n\nnames = []\nfor item in test_ds:\n    names.append(item[2].numpy())\nnames = np.concatenate(names)\nnames = [name.decode() for name in names]\ntst_preds = 0.95 * sub + 0.05 * preds[:,:5]","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:01.524549Z","iopub.execute_input":"2022-04-25T10:43:01.525056Z","iopub.status.idle":"2022-04-25T10:43:04.845938Z","shell.execute_reply.started":"2022-04-25T10:43:01.525015Z","shell.execute_reply":"2022-04-25T10:43:04.844965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Saving results**","metadata":{}},{"cell_type":"code","source":"test[\"label\"] = tst_preds.argmax(1)\ntest[[\"image_id\", \"label\"]].to_csv(OUTPUT_DIR + \"submission.csv\", index=False)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:04.936078Z","iopub.execute_input":"2022-04-25T10:43:04.93695Z","iopub.status.idle":"2022-04-25T10:43:04.95262Z","shell.execute_reply.started":"2022-04-25T10:43:04.936905Z","shell.execute_reply":"2022-04-25T10:43:04.951763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test","metadata":{"execution":{"iopub.status.busy":"2022-04-25T10:43:04.953815Z","iopub.execute_input":"2022-04-25T10:43:04.954338Z","iopub.status.idle":"2022-04-25T10:43:04.964873Z","shell.execute_reply.started":"2022-04-25T10:43:04.954288Z","shell.execute_reply":"2022-04-25T10:43:04.963831Z"},"trusted":true},"execution_count":null,"outputs":[]}]}