{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport tensorflow as tf\nimport tensorflow_hub as hub\nimport numpy as np\nimport pandas as pd","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp -R ../input/cassava-layer/ /kaggle/working/cassava-layer/","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.environ[\"TFHUB_CACHE_DIR\"] = \"/kaggle/working/cassava-layer/\"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"cassava = hub.KerasLayer('https://tfhub.dev/google/cropnet/classifier/cassava_disease_V1/2')\nmodel = tf.keras.Sequential([tf.keras.Input(shape=(224,224,3)),\n                             cassava])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.load_weights(\"../input/cassava-model/cassava_model.h5\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\nBATCH_SIZE = 32\nIMAGE_SIZE = [512, 512]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def 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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def build_valid_ds(valid_fnames):\n    ds = load_dataset(valid_fnames)\n    ds = ds.batch(BATCH_SIZE).prefetch(AUTO)\n    return ds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"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)]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_ds = build_valid_ds(valid_fnames)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"preds = model.predict(test_ds)\nlabels = tf.argmax(preds, axis=-1)\nlabels = labels.numpy()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_ds = build_valid_ds(valid_fnames)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"names = []\nfor item in test_ds:\n    names.append(item[2].numpy())\nnames = np.concatenate(names)\nnames = [name.decode() for name in names]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df = pd.DataFrame({'image_id':names, 'label':labels})\nsubmission_df.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","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}