{"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":"# TFRecord の読み込み について\n\n## 検証の流れ\n\n* TFRecord の読み込みをとりあえず読み込む\n* モデルの学習に使えるように読み込む\n\n## TFRecord とは\n\nTFRecord は、TensorFlow のサポートするデータ形式です。\n\nhttps://www.tensorflow.org/tutorials/load_data/tfrecord?hl=ja#tfrecords_%E5%BD%A2%E5%BC%8F%E3%81%AE%E8%A9%B3%E7%B4%B0","metadata":{}},{"cell_type":"markdown","source":"# TFRecord の読み込みをとりあえず読み込む\n\nまずは、TFRecord を読み込む方法です。\n\n## ファイルリストの取得\n\nTFRecord のデータである拡張子が tfrec であるファイルのリストを作成します","metadata":{}},{"cell_type":"code","source":"import glob\n\ntraining_files = glob.glob('/kaggle/input/tpu-getting-started/tfrecords-jpeg-512x512/train/**/*.tfrec', recursive=True)\ntraining_files","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.616338Z","iopub.execute_input":"2023-02-01T09:26:14.616697Z","iopub.status.idle":"2023-02-01T09:26:14.626008Z","shell.execute_reply.started":"2023-02-01T09:26:14.616672Z","shell.execute_reply":"2023-02-01T09:26:14.625258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## バイナリデータとして読み込む\n\nTFRecord ファイルのリストを`TFRecordDataset`に入力することでバイナリデータとして読み込むことができます","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\nraw_dataset = tf.data.TFRecordDataset(training_files)\nraw_dataset","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.630145Z","iopub.execute_input":"2023-02-01T09:26:14.631089Z","iopub.status.idle":"2023-02-01T09:26:14.644363Z","shell.execute_reply.started":"2023-02-01T09:26:14.631058Z","shell.execute_reply":"2023-02-01T09:26:14.643544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"読み込まれた内容を確認すると、バイナリデータが確認できます。 　  \n※量が多いので初めの５００文字だけ表示しています","metadata":{}},{"cell_type":"code","source":"for example in raw_dataset:\n    print(str(example)[:500])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.648728Z","iopub.execute_input":"2023-02-01T09:26:14.649216Z","iopub.status.idle":"2023-02-01T09:26:14.669990Z","shell.execute_reply.started":"2023-02-01T09:26:14.649183Z","shell.execute_reply":"2023-02-01T09:26:14.668149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## バイナリデータをパース\n\n読み込んだ TFRecord ファイルのバイナリデータはいくつかの項目が含まれています。  \nまずは、その項目をパースすることによって分解します。  \n\nどんな項目があるかを辞書形式で指定して、`parse_single_example`でパースします","metadata":{}},{"cell_type":"code","source":"def _parse_function(example_proto):\n    feature_description = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    \n    return tf.io.parse_single_example(example_proto, feature_description)\n\nparsed_dataset = raw_dataset.map(_parse_function)\nparsed_dataset","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.673017Z","iopub.execute_input":"2023-02-01T09:26:14.673465Z","iopub.status.idle":"2023-02-01T09:26:14.697338Z","shell.execute_reply.started":"2023-02-01T09:26:14.673431Z","shell.execute_reply":"2023-02-01T09:26:14.695662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"パースされた結果を確認します。\n\n`class`は`int64`方で取得されています。  \n今回のデータでは画像のクラス番号（ラベル）です。\n\n`image`は画像データですが、パースの段階でいきなり画像にすることが難しいので、まずは文字列として扱われています。","metadata":{}},{"cell_type":"code","source":"for example in parsed_dataset:\n    print(example.keys())\n    print(example['class'])\n    print(str(example['image'])[:500])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.699153Z","iopub.execute_input":"2023-02-01T09:26:14.699606Z","iopub.status.idle":"2023-02-01T09:26:14.728303Z","shell.execute_reply.started":"2023-02-01T09:26:14.699560Z","shell.execute_reply":"2023-02-01T09:26:14.727307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## データを取得\n\nパースされた情報をもとに、データを使いたい形式に変換します。\n\nTensorFlowに準備された関数を利用することで、文字列形式の画像を画像形式に変換できます。","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = [512, 512]\n\ndef decode_image(image_data):\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.731283Z","iopub.execute_input":"2023-02-01T09:26:14.731864Z","iopub.status.idle":"2023-02-01T09:26:14.739320Z","shell.execute_reply.started":"2023-02-01T09:26:14.731826Z","shell.execute_reply":"2023-02-01T09:26:14.737423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"結果を確認します。","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nfor example in parsed_dataset:\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    \n    print('label :', label)\n    plt.imshow(image)\n    \n    break","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:14.740569Z","iopub.execute_input":"2023-02-01T09:26:14.741580Z","iopub.status.idle":"2023-02-01T09:26:15.006611Z","shell.execute_reply.started":"2023-02-01T09:26:14.741545Z","shell.execute_reply":"2023-02-01T09:26:15.005750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# モデルの学習に使えるように読み込む\n\n上記では、単純に画像ファイルを取得しました。  \n今回は、モデルの学習に使えるように読み込みます。  \n具体的には、`model.fit`の入力として使える様にします。\n\n## 学習用形式でのパース\n\n「TFRecord の読み込みをとりあえず読み込む」では確認のために、パースやフォーマットの変換を順番に行なっていましたが、  \n学習に利用する際には１つの関数で実施するとうまくいきます。","metadata":{}},{"cell_type":"code","source":"def parse_function(example_proto):\n    feature_description = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    \n    example = tf.io.parse_single_example(example_proto, feature_description)\n    \n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    \n    return image, label","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:15.008096Z","iopub.execute_input":"2023-02-01T09:26:15.008605Z","iopub.status.idle":"2023-02-01T09:26:15.016713Z","shell.execute_reply.started":"2023-02-01T09:26:15.008573Z","shell.execute_reply":"2023-02-01T09:26:15.015157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"パース用関数の使い方は、「TFRecord の読み込みをとりあえず読み込む」の時と同じです。 　  \n\n学習時に便利なオプションが準備されています\n* `batch`: 学習時のバッチサイズ\n* `repeat`: 学習時に指定したエポック数繰り返すために、繰り返しデータを利用\n* `shuffle`: データの順番を入れ替える\n* `prefetch`: 学習中に次のバッチをあらかじめ読み込んでおく","metadata":{}},{"cell_type":"code","source":"train_data = raw_dataset.map(parse_function)\ntrain_data = train_data.batch(10)\ntrain_data = train_data.repeat()\ntrain_data = train_data.shuffle(128)\n\nAUTO = tf.data.experimental.AUTOTUNE\ntrain_data = train_data.prefetch(AUTO)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:15.019522Z","iopub.execute_input":"2023-02-01T09:26:15.020257Z","iopub.status.idle":"2023-02-01T09:26:15.108924Z","shell.execute_reply.started":"2023-02-01T09:26:15.020207Z","shell.execute_reply":"2023-02-01T09:26:15.107030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## モデルに学習させてテスト\n\nモデルを適当に作成します。","metadata":{}},{"cell_type":"code","source":"class_num = 104\n\nfrom tensorflow.keras.applications.inception_resnet_v2 import InceptionResNetV2\n\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Activation\nfrom tensorflow.keras.layers import Flatten\nfrom tensorflow.keras.layers import Dense\nfrom tensorflow.keras.layers import Dropout\nfrom tensorflow.keras.layers import Input\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.optimizers import SGD\n\ninput_tensor = Input(shape=(*IMAGE_SIZE, 3))\nbase_model = InceptionResNetV2(include_top=False, weights='imagenet', input_tensor=input_tensor)\n\nbase_model.trainable = False\n\nmodel = Sequential()\nmodel.add(base_model)\nmodel.add(Flatten())\nmodel.add(Dense(128, activation='relu'))\nmodel.add(Dropout(0.5))\nmodel.add(Dense(class_num, activation='sigmoid'))\n\nsdg = SGD(learning_rate=1e-3, momentum=0.9)\nmodel.compile(loss='sparse_categorical_crossentropy', optimizer=sdg, metrics=['accuracy'])","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:15.111701Z","iopub.execute_input":"2023-02-01T09:26:15.111999Z","iopub.status.idle":"2023-02-01T09:26:22.091713Z","shell.execute_reply.started":"2023-02-01T09:26:15.111974Z","shell.execute_reply":"2023-02-01T09:26:22.090675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"学習せてみる。","metadata":{}},{"cell_type":"code","source":"history = model.fit(train_data, epochs=5, steps_per_epoch=10)","metadata":{"execution":{"iopub.status.busy":"2023-02-01T09:26:22.092851Z","iopub.execute_input":"2023-02-01T09:26:22.094013Z"},"trusted":true},"execution_count":null,"outputs":[]}]}