{"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":"## Keras Baseline(EfficientNetB3)\n* [Baseline](https://www.kaggle.com/code/juhjoo/0-56-tf-keras-efficientnet-rsna-baseline)\n\n## TFRecord Training\n* [Training](https://www.kaggle.com/code/juhjoo/rsna-2022-training)\n\n## TFRecord Data\n* Data: [TFRecord Data](https://www.kaggle.com/code/juhjoo/rsna-2022-tfrecords)","metadata":{}},{"cell_type":"code","source":"!conda install '../input/rsna2022-pydicom-conda-helper/libjpeg-turbo-2.1.0-h7f98852_0.tar.bz2' -c conda-forge -y\n!conda install '../input/rsna2022-pydicom-conda-helper/libgcc-ng-9.3.0-h2828fa1_19.tar.bz2' -c conda-forge -y\n!conda install '../input/rsna2022-pydicom-conda-helper/gdcm-2.8.9-py37h500ead1_1.tar.bz2' -c conda-forge -y\n!conda install '../input/rsna2022-pydicom-conda-helper/conda-4.10.1-py37h89c1867_0.tar.bz2' -c conda-forge -y\n!conda install '../input/rsna2022-pydicom-conda-helper/certifi-2020.12.5-py37h89c1867_1.tar.bz2' -c conda-forge -y\n!conda install '../input/rsna2022-pydicom-conda-helper/openssl-1.1.1k-h7f98852_0.tar.bz2' -c conda-forge -y","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:50:31.684385Z","iopub.execute_input":"2022-10-18T12:50:31.68567Z","iopub.status.idle":"2022-10-18T12:51:35.146499Z","shell.execute_reply.started":"2022-10-18T12:50:31.68547Z","shell.execute_reply":"2022-10-18T12:51:35.145316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qU \"python-gdcm\" pydicom pylibjpeg \"opencv-python-headless\"","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:51:35.15013Z","iopub.execute_input":"2022-10-18T12:51:35.150625Z","iopub.status.idle":"2022-10-18T12:51:54.613609Z","shell.execute_reply.started":"2022-10-18T12:51:35.150573Z","shell.execute_reply":"2022-10-18T12:51:54.612084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Core\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport gc\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom tqdm.notebook import tqdm\nfrom datetime import datetime\nimport json,itertools\nfrom typing import Optional\nfrom glob import glob\nimport warnings\nfrom IPython import display as ipd\nwarnings.filterwarnings(\"ignore\")\nimport matplotlib.gridspec as gridspec\nimport matplotlib.patches as mpatches\nimport matplotlib as mpl\nfrom matplotlib.patches import Rectangle\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\nfrom sklearn.preprocessing import LabelEncoder\nimport random\nfrom joblib import Parallel, delayed\nimport os, shutil\nimport math\nfrom sklearn.utils import check_random_state\nimport glob\nfrom collections import Counter, defaultdict\nimport pydicom as dicom\nfrom sklearn.model_selection import GroupKFold\nfrom glob import glob\nimport re\n\n# Keras\nfrom tensorflow import keras\nimport tensorflow as tf\nimport keras\nfrom keras import backend as K\nfrom keras.models import Model\nfrom keras.layers import Input\nfrom keras.layers.convolutional import Conv2D, Conv2DTranspose\nfrom keras.layers.pooling import MaxPooling2D\nfrom keras.layers.merge import concatenate\nfrom keras.losses import binary_crossentropy\nfrom keras.callbacks import Callback, ModelCheckpoint, EarlyStopping\nfrom keras.models import load_model, save_model","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:51:54.616102Z","iopub.execute_input":"2022-10-18T12:51:54.617096Z","iopub.status.idle":"2022-10-18T12:52:00.840517Z","shell.execute_reply.started":"2022-10-18T12:51:54.617027Z","shell.execute_reply":"2022-10-18T12:52:00.839347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n\n    # Use verbose=0 for silent, 1 for interactive\n    verbose = 0\n    apex = False\n    print_freq = 100\n    display_plot = True\n    debug = False\n\n    # Device for training\n    device = None  # device is automatically selected\n\n    # Model\n    model_name = (\n        \"resnext50_32x4d\"  # resnext50_32x4d, tf_efficientnet_b3_ns, tf_efficientnet_b4_ns, vit_base_patch16_384\n    )\n    rsna_2022_path = '../input/rsna-2022-cervical-spine-fracture-detection'\n    train_image_path = f'{rsna_2022_path}/train_images'\n    test_image_path = f'{rsna_2022_path}/test_images'\n    meta_path = '../input/rsna-2022-dataset'\n    target_size = 7\n\n    # Seeding for reproducibility\n    seed = 101\n\n    # Number of folds\n    folds = 5\n\n    # Which Folds to train\n    selected_folds = [0, 1, 2, 3, 4]\n\n    # Image Size\n    img_size = [510, 510]\n    IMAGE_SIZE = 510\n\n    # Batch Size & Epochs\n    batch_size = 30   # resnext50_32x4d: 14-30, tf_efficientnet_b3_ns:10-22, tf_efficientnet_b4_ns: 8-16\n    num_workers = 4\n\n\n    # Loss & Optimizer\n    optimizer = \"Adam\"\n    loss_weight = 2.\n    n_folds = 5\n    \n    scheduler = \"CosineAnnealingWarmRestarts\"\n    epochs = 10\n    # factor=0.2 # ReduceLROnPlateau\n    # patience=4 # ReduceLROnPlateau\n    # eps=1e-6 # ReduceLROnPlateau\n    # T_max=10 # CosineAnnealingLR\n    T_0 = 10  # CosineAnnealingWarmRestarts\n    lr = 1e-4\n    min_lr = 1e-6\n    weight_decay = 1e-6\n    gradient_accumulation_steps = 1\n    max_grad_norm = 1000\n    \n    n_splits = 5\n    \n    clip = False\n    \n    augment = True\n    \n    \n    # Horizontal & Vertical Flip\n    hflip = 0.5\n    vflip = 0.5\n\n    # Random Bright\n    brightness_limit = 0.2\n    contrast_limit = 0.2\n    Brightp = 0.75\n\n    # Clip values to [0, 1]\n    clip = False\n\n    # ShiftScaleRotate\n    shift_limit=0.125\n    scale_limit=0.1\n    rotate_limit=20\n    rotatep=0.075\n\n    # CutOut\n    drop_prob = 0.5\n    drop_cnt = 10\n    drop_size = 0.05\n    \nBATCH_SIZE = 32\nAUTO = tf.data.experimental.AUTOTUNE  \nSEED = 101","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:00.843224Z","iopub.execute_input":"2022-10-18T12:52:00.843812Z","iopub.status.idle":"2022-10-18T12:52:00.856128Z","shell.execute_reply.started":"2022-10-18T12:52:00.84378Z","shell.execute_reply":"2022-10-18T12:52:00.854497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(SEED):\n    \"\"\"\n    Sets all random seeds for the program (Python, NumPy, and TensorFlow).\n    \"\"\"\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ[\"PYTHONHASHSEED\"] = str(SEED)\n    os.environ[\"TF_CUDNN_DETERMINISTIC\"] = str(SEED)\n    print(\"seeding done\")\n\n\nseeding(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:00.857709Z","iopub.execute_input":"2022-10-18T12:52:00.858062Z","iopub.status.idle":"2022-10-18T12:52:00.877599Z","shell.execute_reply.started":"2022-10-18T12:52:00.858028Z","shell.execute_reply":"2022-10-18T12:52:00.876233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\ntrain_bbox = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv\")\ntest_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\")\nss = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/sample_submission.csv\")\nbase_path = \"../input/rsna-2022-cervical-spine-fracture-detection\"\n\nprint('train shape:', train_df.shape)\nprint('train bbox shape:', train_bbox.shape)\nprint('test shape:', test_df.shape)\nprint('ss shape:', ss.shape)\nprint('')\n\ntrain_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:00.879321Z","iopub.execute_input":"2022-10-18T12:52:00.880173Z","iopub.status.idle":"2022-10-18T12:52:00.958177Z","shell.execute_reply.started":"2022-10-18T12:52:00.88011Z","shell.execute_reply":"2022-10-18T12:52:00.956813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_train = pd.read_csv(\"../input/rsna-2022-dataset/meta_data.csv\")\nmeta_train[\"StudyInstanceUID\"] = meta_train[\"SOPInstanceUID\"].apply(lambda x: \".\".join(x.split(\".\")[:-2]))\nprint('meta_train shape:', meta_train.shape)\nmeta_train_clean = meta_train.drop(['SOPInstanceUID','ImagePositionPatient','ImageOrientationPatient'], axis=1)\nmeta_train_clean.rename(columns={\"Rows\": \"ImageHeight\", \"Columns\": \"ImageWidth\",\"InstanceNumber\": \"Slice\"}, inplace=True)\nmeta_train_clean = meta_train_clean[['StudyInstanceUID','Slice','ImageHeight','ImageWidth','SliceThickness']]\nmeta_train_clean.sort_values(by=['StudyInstanceUID','Slice'], inplace=True)\nmeta_train_clean.reset_index(drop=True, inplace=True)\n\nmeta_train_clean.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:00.959747Z","iopub.execute_input":"2022-10-18T12:52:00.960196Z","iopub.status.idle":"2022-10-18T12:52:04.827145Z","shell.execute_reply.started":"2022-10-18T12:52:00.960161Z","shell.execute_reply":"2022-10-18T12:52:04.825989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = meta_train_clean.set_index('StudyInstanceUID').join(train_df.set_index('StudyInstanceUID')).reset_index().copy()\ntrain_df = train_df.query('StudyInstanceUID != \"1.2.826.0.1.3680043.20574\"').reset_index(drop=True)\ntrain_df.sample(5)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:04.828437Z","iopub.execute_input":"2022-10-18T12:52:04.828803Z","iopub.status.idle":"2022-10-18T12:52:05.522653Z","shell.execute_reply.started":"2022-10-18T12:52:04.82877Z","shell.execute_reply":"2022-10-18T12:52:05.521883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"split = GroupKFold(CFG.folds)\nfor k, (_, test_idx) in enumerate(split.split(train_df, groups=train_df.StudyInstanceUID)):\n    train_df.loc[test_idx, 'split'] = k\ntrain_df.sample(2)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:05.52377Z","iopub.execute_input":"2022-10-18T12:52:05.524671Z","iopub.status.idle":"2022-10-18T12:52:06.238994Z","shell.execute_reply.started":"2022-10-18T12:52:05.524638Z","shell.execute_reply":"2022-10-18T12:52:06.237735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f'{CFG.rsna_2022_path}/test.csv')\ntest_df","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.242608Z","iopub.execute_input":"2022-10-18T12:52:06.242986Z","iopub.status.idle":"2022-10-18T12:52:06.256564Z","shell.execute_reply.started":"2022-10-18T12:52:06.242955Z","shell.execute_reply":"2022-10-18T12:52:06.255606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if test_df.iloc[0].row_id == '1.2.826.0.1.3680043.10197_C1':\n    test_df = pd.DataFrame({\n        \"row_id\": ['1.2.826.0.1.3680043.22327_C1', '1.2.826.0.1.3680043.25399_C1', '1.2.826.0.1.3680043.5876_C1'],\n        \"StudyInstanceUID\": ['1.2.826.0.1.3680043.22327', '1.2.826.0.1.3680043.25399', '1.2.826.0.1.3680043.5876'],\n        \"prediction_type\": [\"C1\", \"C1\", \"patient_overall\"]}\n    )\ntest_df","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.257907Z","iopub.execute_input":"2022-10-18T12:52:06.258705Z","iopub.status.idle":"2022-10-18T12:52:06.274417Z","shell.execute_reply.started":"2022-10-18T12:52:06.258657Z","shell.execute_reply":"2022-10-18T12:52:06.272934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_slices = glob(f'{CFG.test_image_path}/*/*')\ntest_slices = [re.findall(f'{CFG.test_image_path}/(.*)/(.*).dcm', s)[0] for s in test_slices]\ndf_test_slices = pd.DataFrame(data=test_slices, columns=['StudyInstanceUID', 'Slice'])\ndf_test_slices.sample(2)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.27619Z","iopub.execute_input":"2022-10-18T12:52:06.276693Z","iopub.status.idle":"2022-10-18T12:52:06.411612Z","shell.execute_reply.started":"2022-10-18T12:52:06.276645Z","shell.execute_reply":"2022-10-18T12:52:06.410774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test = test_df.set_index('StudyInstanceUID').join(df_test_slices.set_index('StudyInstanceUID')).reset_index()\ndf_test.sample(2)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.412783Z","iopub.execute_input":"2022-10-18T12:52:06.413635Z","iopub.status.idle":"2022-10-18T12:52:06.435812Z","shell.execute_reply.started":"2022-10-18T12:52:06.4136Z","shell.execute_reply":"2022-10-18T12:52:06.434957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path) :\n    img = dicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    data = cv2.resize(data, (CFG.IMAGE_SIZE, CFG.IMAGE_SIZE), interpolation = cv2.INTER_AREA)\n    return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), data\n\ndef listdirs(folder):\n    return [d for d in os.listdir(folder) if os.path.isdir(os.path.join(folder, d))]    \n\ntrain_dir = f'{CFG.rsna_2022_path}/train_images'\ntest_dir = f'{CFG.rsna_2022_path}/test_images'\npatients = sorted(os.listdir(train_dir))\nlen(patients), patients[:5]","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.437315Z","iopub.execute_input":"2022-10-18T12:52:06.438434Z","iopub.status.idle":"2022-10-18T12:52:06.529292Z","shell.execute_reply.started":"2022-10-18T12:52:06.43839Z","shell.execute_reply":"2022-10-18T12:52:06.528035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"im, meta = load_dicom(\n    f'{CFG.train_image_path}/1.2.826.0.1.3680043.10001/1.dcm')\nplt.figure()\nplt.imshow(im)\nplt.title('regular image')","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.530736Z","iopub.execute_input":"2022-10-18T12:52:06.531089Z","iopub.status.idle":"2022-10-18T12:52:06.845706Z","shell.execute_reply.started":"2022-10-18T12:52:06.531058Z","shell.execute_reply":"2022-10-18T12:52:06.844589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.sample(2)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.847176Z","iopub.execute_input":"2022-10-18T12:52:06.848022Z","iopub.status.idle":"2022-10-18T12:52:06.887059Z","shell.execute_reply.started":"2022-10-18T12:52:06.847984Z","shell.execute_reply":"2022-10-18T12:52:06.885869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folds = train_df\nfolds.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.888626Z","iopub.execute_input":"2022-10-18T12:52:06.889065Z","iopub.status.idle":"2022-10-18T12:52:06.907618Z","shell.execute_reply.started":"2022-10-18T12:52:06.88902Z","shell.execute_reply":"2022-10-18T12:52:06.906423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _bytes_feature(value):\n  \"\"\"Returns a bytes_list from a string / byte.\"\"\"\n  if isinstance(value, type(tf.constant(0))):\n    value = value.numpy() # BytesList won't unpack a string from an EagerTensor.\n  return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))\n\ndef _float_feature(value):\n  \"\"\"Returns a float_list from a float / double.\"\"\"\n  return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))\n\ndef _int64_feature(value):\n  \"\"\"Returns an int64_list from a bool / enum / int / uint.\"\"\"\n  return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.909523Z","iopub.execute_input":"2022-10-18T12:52:06.91009Z","iopub.status.idle":"2022-10-18T12:52:06.919014Z","shell.execute_reply.started":"2022-10-18T12:52:06.910044Z","shell.execute_reply":"2022-10-18T12:52:06.917622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def serialize_example(img, pati_oa, c1, c2, c3, c4, c5, c6, c7):\n    feature = {\n      'image': _bytes_feature(img),\n      'patient_overall': _int64_feature(pati_oa),\n      'C1': _int64_feature(c1),\n      'C2': _int64_feature(c2),\n      'C3': _int64_feature(c3),\n      'C4': _int64_feature(c4),\n      'C5': _int64_feature(c5),\n      'C6': _int64_feature(c6),\n      'C7': _int64_feature(c7),\n  }\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n    \n    return example_proto.SerializeToString()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.9205Z","iopub.execute_input":"2022-10-18T12:52:06.920903Z","iopub.status.idle":"2022-10-18T12:52:06.934425Z","shell.execute_reply.started":"2022-10-18T12:52:06.920866Z","shell.execute_reply":"2022-10-18T12:52:06.93324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for f in range(0, CFG.folds):\n#for f in range(0, 2):\nfor f in range(3, 5):\n    ct = (folds['split'] == f).sum()\n    idx = folds[folds['split'] == f].index\n    print(idx)\n    print(ct)\n    print('Writing TFRecord %i of %i...'%(f,ct))\n    flg = 0\n    flg2 = 1\n    with tf.io.TFRecordWriter('train%.2i-%i.tfrec'%(f,ct)) as writer:\n        for k in tqdm(idx):\n            #imgs = glob.glob(BASE+'/train_images/'+folds['StudyInstanceUID'][idx[k]]+'/*')\n            path = os.path.join(CFG.train_image_path, folds.iloc[k].StudyInstanceUID, f'{folds.iloc[k].Slice}.dcm')\n            img = load_dicom(path)[0]\n            if flg<1: \n                print(img.shape)\n                plt.imshow(img),plt.show()\n                flg += 1\n            img = cv2.imencode('.jpg', img, (cv2.IMWRITE_JPEG_QUALITY, 94))[1].tostring()    \n            example = serialize_example(\n                img, \n                folds.iloc[k].patient_overall,\n                folds.iloc[k].C1,\n                folds.iloc[k].C2,\n                folds.iloc[k].C3,\n                folds.iloc[k].C4,\n                folds.iloc[k].C5,\n                folds.iloc[k].C6,\n                folds.iloc[k].C7,\n                )\n            writer.write(example)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:52:06.936398Z","iopub.execute_input":"2022-10-18T12:52:06.936745Z","iopub.status.idle":"2022-10-18T12:54:52.586745Z","shell.execute_reply.started":"2022-10-18T12:52:06.936715Z","shell.execute_reply":"2022-10-18T12:54:52.585126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAINING_FILENAMES = tf.io.gfile.glob('./train00-142247.tfrec')\nTRAINING_FILENAMES","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:54:52.587867Z","iopub.status.idle":"2022-10-18T12:54:52.588298Z","shell.execute_reply.started":"2022-10-18T12:54:52.588099Z","shell.execute_reply":"2022-10-18T12:54:52.588119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re, math\ndef decode_image(data, height, width, target_size=(320, 320)):\n    image = tf.io.decode_jpeg(data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0 \n    image = tf.reshape(image, [*target_size, 3]) # explicit size needed for TPU\n    return image\n\n\ndef read_labeled_tfrecord(example):\n    LABELED_TFREC_FORMAT = {\n        \"image\" : tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'], 320, 320)\n    return image # returns a dataset of (image, label) pairs\n\ndef load_dataset(fileids, labeled=True, ordered=False):\n    # Read from TFRecords. For optimal performance, reading from multiple files at once and\n    # disregarding data order. Order does not matter since we will be shuffling the data anyway.\n\n    ignore_order = tf.data.Options()\n    if not ordered:\n        ignore_order.experimental_deterministic = False # disable order, increase speed\n\n    dataset = tf.data.TFRecordDataset(fileids, num_parallel_reads=AUTO) # automatically interleaves reads from multiple files\n    dataset = dataset.with_options(ignore_order) # uses data as soon as it streams in, rather than in its original order\n    dataset = dataset.map(read_labeled_tfrecord)\n    # returns a dataset of (image, label) pairs if labeled=True or (image, id) pairs if labeled=False\n    return dataset\n\ndef get_training_dataset():\n    dataset = load_dataset(TRAINING_FILENAMES, labeled=True)\n    dataset = dataset.repeat() # the training dataset must repeat for several epochs\n    dataset = dataset.shuffle(20, seed=SEED)\n    dataset = dataset.batch(BATCH_SIZE)\n    dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)\n    return dataset\n\ndef count_data_items(fileids):\n    # the number of data items is written in the id of the .tfrec files, i.e. flowers00-230.tfrec = 230 data items\n    n = [int(re.compile(r\"-([0-9]*)\\.\").search(fileid).group(1)) for fileid in fileids]\n    return np.sum(n)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:54:52.589494Z","iopub.status.idle":"2022-10-18T12:54:52.589916Z","shell.execute_reply.started":"2022-10-18T12:54:52.589694Z","shell.execute_reply":"2022-10-18T12:54:52.589713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_batch(batch, size=5):\n    imgs = batch\n    plt.figure(figsize=(size*10, 10))\n    for img_idx in range(size):\n        plt.subplot(1, size, img_idx+1)\n        plt.imshow(imgs[img_idx, ], cmap='bone')\n        plt.xticks([])\n        plt.yticks([])\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:54:52.591204Z","iopub.status.idle":"2022-10-18T12:54:52.591606Z","shell.execute_reply.started":"2022-10-18T12:54:52.591413Z","shell.execute_reply":"2022-10-18T12:54:52.591432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_dataset = get_training_dataset().take(3439)\ntraining_dataset = training_dataset.unbatch().batch(20)\ntrain_batch = next(iter(training_dataset))\ndisplay_batch(train_batch, 5)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T12:54:52.593468Z","iopub.status.idle":"2022-10-18T12:54:52.593909Z","shell.execute_reply.started":"2022-10-18T12:54:52.593673Z","shell.execute_reply":"2022-10-18T12:54:52.593692Z"},"trusted":true},"execution_count":null,"outputs":[]}]}