{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"%matplotlib inline\n\nfrom kaggle_datasets import KaggleDatasets\nimport matplotlib.gridspec as gridspec\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nimport tensorflow as tf","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\nDATASET_DIR = Path('/kaggle/input/flowers-tta')\nPATH = Path('/kaggle/input/flower-classification-with-tpus')\nSIZES = {s: f'{s}x{s}' for s in [192, 224, 331, 512]}\nTFRECORD_DIR = KaggleDatasets().get_gcs_path(PATH.parts[-1])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"classes_filename = DATASET_DIR/'classes.csv' \nCLASSES = tf.constant(pd.read_csv(classes_filename).values.squeeze(), tf.string)\n\ndef get_parse_fn(split):\n    def parse_fn(example):\n        features = {\"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n                    \"id\": tf.io.FixedLenFeature([], tf.string),\n                    \"class\": tf.io.FixedLenFeature([], tf.int64)}\n        \n        if split == 'test':\n            del features['class']\n            \n        example = tf.io.parse_single_example(example, features)\n            \n        example['image'] = tf.image.decode_jpeg(example['image'])\n        example['label'] = tf.cast(example['class'], tf.int32)\n        example['class'] = CLASSES[example['label']]\n        \n        return example\n\n    return parse_fn\n\ndef get_ds(split, img_size=224, batch_size=128, shuffle=False):\n    file_pat = f'{TFRECORD_DIR}/tfrecords-jpeg-{SIZES[img_size]}/{split}/*.tfrec'\n    \n    options = tf.data.Options()\n    options.experimental_deterministic = not shuffle\n    \n    ds = (tf.data.Dataset.list_files(file_pat, shuffle=shuffle)\n          .with_options(options)\n          .interleave(tf.data.TFRecordDataset, num_parallel_calls=AUTO)\n          .map(get_parse_fn(split), num_parallel_calls=AUTO)\n         )\n    \n    if shuffle:\n        ds = ds.shuffle(2048)\n            \n    return ds.repeat().batch(batch_size).prefetch(AUTO)\n\ndef show_images(imgs, titles=None, hw=(3,3), rc=(4,4)):\n    \"\"\"Show list of images with optional list of titles.\"\"\"\n    h, w = hw\n    r, c = rc\n    fig=plt.figure(figsize=(w*c, h*r))\n    gs1 = gridspec.GridSpec(r, c, fig, hspace=0.2, wspace=0.05)\n    for i in range(r*c):\n        img = imgs[i].squeeze()\n        ax = fig.add_subplot(gs1[i])\n        if titles != None:\n            ax.set_title(titles[i], {'fontsize': 10})\n        plt.imshow(img)\n        plt.axis('off')\n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Create Datasets"},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_val = get_ds('val', shuffle=False)\nds_val_iter = ds_val.unbatch().batch(16).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"b_val = next(ds_val_iter)\nshow_images(b_val['image'], b_val['class'].tolist(), hw=(2,2), rc=(2,8))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_trn = get_ds('train', shuffle=False)\nds_trn_iter = ds_trn.unbatch().batch(16).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"b_trn = next(ds_trn_iter)\nshow_images(b_trn['image'], b_trn['class'].tolist(), hw=(2,2), rc=(2,8))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Sample"},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_sample = tf.data.experimental.sample_from_datasets([ds_trn.unbatch(), ds_val.unbatch()], [1., 1.])\nds_sample_iter = ds_sample.batch(16).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"b_smp = next(ds_sample_iter)\nshow_images(b_smp['image'], b_smp['class'].tolist(), hw=(2,2), rc=(2,8))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Choose"},{"metadata":{"trusted":true},"cell_type":"code","source":"choices = tf.data.Dataset.range(2).repeat()\nds_choose = tf.data.experimental.choose_from_datasets([ds_trn.unbatch(), ds_val.unbatch()], choices)\nds_choose_iter = ds_choose.batch(16).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"b_ch = next(ds_choose_iter)\nshow_images(b_ch['image'], b_ch['class'].tolist(), hw=(2,2), rc=(2,8))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Zip "},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_zip = tf.data.Dataset.zip((ds_val.unbatch(), ds_trn.unbatch()))\nds_zip_iter = ds_zip.batch(8).as_numpy_iterator()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"b_zip = next(ds_zip_iter)\nshow_images(b_zip[0]['image'], b_zip[0]['class'].tolist(), hw=(2,2), rc=(1,8))\nshow_images(b_zip[1]['image'], b_zip[1]['class'].tolist(), hw=(2,2), rc=(1,8))","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}