{"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を利用して「Petals to the Metal - Flower Classification on TPU」を実施\n\n* ポイント\n    * データの読み込み、学習、予測の一連の流れを実施\n    * Keras の `ImageDataGenerator` を利用してデータ拡張を実施\n    * 白黒画像に対して実施\n","metadata":{}},{"cell_type":"markdown","source":"# 設定値の決定\n\n* target_size: 読み込んだ画像をこのサイズに変換して利用する\n* batch_size: 学習時のデータに対するbatch_size\n* learning_rate: 学習を実施する際の学習率\n* epochs: 学習を実施する際のエポック数\n","metadata":{}},{"cell_type":"code","source":"target_size = [192, 192]\nbatch_size = 50\nlearning_rate = 5e-2\nepochs = 100","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:13:00.754605Z","iopub.execute_input":"2023-02-14T05:13:00.755003Z","iopub.status.idle":"2023-02-14T05:13:00.760613Z","shell.execute_reply.started":"2023-02-14T05:13:00.754970Z","shell.execute_reply":"2023-02-14T05:13:00.759297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# データの読み込み\n\n「104-flowers-garden-of-eden」という外部データを利用する。  \nこちらのデータは「Petals to the Metal - Flower Classification on TPU」の画像をJPEG形式に変換して、クラス毎のディレクトリに保存してくれている。\n`flow_from_directory`で画像を読み込む際の形式と一致している。\n","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras.preprocessing.image import ImageDataGenerator\n\ntrain_dir = '/kaggle/input/104-flowers-garden-of-eden/jpeg-192x192/train'\nval_dir = '/kaggle/input/104-flowers-garden-of-eden/jpeg-192x192/val'\n\nclasses = ['pink primrose',    'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea',     'wild geranium',     'tiger lily',           'moon orchid',              'bird of paradise', 'monkshood',        'globe thistle',         # 00 - 09\n           'snapdragon',       \"colt's foot\",               'king protea',      'spear thistle', 'yellow iris',       'globe-flower',         'purple coneflower',        'peruvian lily',    'balloon flower',   'giant white arum lily', # 10 - 19\n           'fire lily',        'pincushion flower',         'fritillary',       'red ginger',    'grape hyacinth',    'corn poppy',           'prince of wales feathers', 'stemless gentian', 'artichoke',        'sweet william',         # 20 - 29\n           'carnation',        'garden phlox',              'love in the mist', 'cosmos',        'alpine sea holly',  'ruby-lipped cattleya', 'cape flower',              'great masterwort', 'siam tulip',       'lenten rose',           # 30 - 39\n           'barberton daisy',  'daffodil',                  'sword lily',       'poinsettia',    'bolero deep blue',  'wallflower',           'marigold',                 'buttercup',        'daisy',            'common dandelion',      # 40 - 49\n           'petunia',          'wild pansy',                'primula',          'sunflower',     'lilac hibiscus',    'bishop of llandaff',   'gaura',                    'geranium',         'orange dahlia',    'pink-yellow dahlia',    # 50 - 59\n           'cautleya spicata', 'japanese anemone',          'black-eyed susan', 'silverbush',    'californian poppy', 'osteospermum',         'spring crocus',            'iris',             'windflower',       'tree poppy',            # 60 - 69\n           'gazania',          'azalea',                    'water lily',       'rose',          'thorn apple',       'morning glory',        'passion flower',           'lotus',            'toad lily',        'anthurium',             # 70 - 79\n           'frangipani',       'clematis',                  'hibiscus',         'columbine',     'desert-rose',       'tree mallow',          'magnolia',                 'cyclamen ',        'watercress',       'canna lily',            # 80 - 89\n           'hippeastrum ',     'bee balm',                  'pink quill',       'foxglove',      'bougainvillea',     'camellia',             'mallow',                   'mexican petunia',  'bromelia',         'blanket flower',        # 90 - 99\n           'trumpet creeper',  'blackberry lily',           'common tulip',     'wild rose']                                                                                                                                               # 100 - 102\nclass_num = len(classes)\n\ngenerator = ImageDataGenerator(\n    rescale=1./255,\n    horizontal_flip=True)\ntrain_data = generator.flow_from_directory(\n    train_dir, color_mode='grayscale', classes=classes, batch_size=batch_size,\n    target_size=target_size)\nval_data = generator.flow_from_directory(\n    val_dir, color_mode='grayscale', classes=classes, batch_size=batch_size,\n    target_size=target_size)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-14T05:13:01.240672Z","iopub.execute_input":"2023-02-14T05:13:01.241054Z","iopub.status.idle":"2023-02-14T05:13:03.498841Z","shell.execute_reply.started":"2023-02-14T05:13:01.241024Z","shell.execute_reply":"2023-02-14T05:13:03.497685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"試しに表示してみる","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n%matplotlib inline\n\nimage_num = 0\nplt.imshow(train_data[0][0][image_num])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:13:03.500868Z","iopub.execute_input":"2023-02-14T05:13:03.501842Z","iopub.status.idle":"2023-02-14T05:13:04.026028Z","shell.execute_reply.started":"2023-02-14T05:13:03.501810Z","shell.execute_reply":"2023-02-14T05:13:04.025021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# モデルの定義\n\n`keras.applications`の学習済みモデルを利用する","metadata":{}},{"cell_type":"markdown","source":"`keras.applications`の学習済みモデルはカラー画像に対して学習してるので、  \n入力するデータは（縦サイズ,横サイズ,3）の形式である必要がある。\n\n```\ninput_tensor = Input(shape=(*target_size, 1))\nbase_model = InceptionResNetV2(include_top=False, weights='imagenet', input_tensor=input_tensor)\n```\n\nのままだと、入力サイズが（縦サイズ,横サイズ,1）なので、次のようなエラーが出る\n\n```\n---------------------------------------------------------------------------\nValueError                                Traceback (most recent call last)\n/tmp/ipykernel_27/1638831701.py in <module>\n     12 \n     13 input_tensor = Input(shape=(*target_size, 1))\n---> 14 base_model = InceptionResNetV2(include_top=False, weights='imagenet', input_tensor=input_tensor)\n     15 \n     16 base_model.trainable = False\n\n...\n\nValueError: Cannot assign to variable conv2d/kernel:0 due to variable shape (3, 3, 1, 32) and value shape (32, 3, 3, 3) are incompatible\n```\n\n次のように`Lambda`層を付け加えることで動作可能となる。","metadata":{}},{"cell_type":"code","source":"# from tensorflow.keras.applications.efficientnet import EfficientNetB1\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\nimport tensorflow as tf\nfrom tensorflow.keras.layers import Lambda\n\nbase_model = InceptionResNetV2(include_top=False, weights='imagenet')\nbase_model.trainable = False\n\nmodel = Sequential()\nmodel.add(Input(shape=(*target_size, 1)))\nmodel.add(Lambda(lambda x: tf.tile(x, [1,1,1,3])))\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='categorical_crossentropy',\n          optimizer=sdg, metrics=['accuracy'])","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:17:00.187153Z","iopub.execute_input":"2023-02-14T05:17:00.187710Z","iopub.status.idle":"2023-02-14T05:17:07.006456Z","shell.execute_reply.started":"2023-02-14T05:17:00.187655Z","shell.execute_reply":"2023-02-14T05:17:07.005392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"`Lambda`層の動作を確認する。\n\n`Lambda`層では任意の関数を適用することができる。  \n今回は`tf.tile`を利用し層の複製をおこなっている。\n\n（縦サイズ,横サイズ,1）が（縦サイズ,横サイズ,3）に変換されていることがわかる\n\n白黒の画像を複製しているので、どのチャンネルを取り出しても同じ画像になる。","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.layers import Lambda\n\ninput_x = train_data[0][0]\nprint(input_x.shape)\noutput_x = Lambda(lambda x: tf.tile(x, [1,1,1,3]))(input_x)\nprint(output_x.shape)\n\nplt.imshow(output_x[0,:,:,0])\nplt.show()\n\n\nplt.imshow(output_x[0,:,:,1])\nplt.show()\n\n\nplt.imshow(output_x[0,:,:,2])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:20:00.906829Z","iopub.execute_input":"2023-02-14T05:20:00.907211Z","iopub.status.idle":"2023-02-14T05:20:01.530463Z","shell.execute_reply.started":"2023-02-14T05:20:00.907179Z","shell.execute_reply":"2023-02-14T05:20:01.529297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 学習の実施","metadata":{}},{"cell_type":"code","source":"history = model.fit(\n    train_data, epochs=epochs, validation_data=val_data)","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:17:11.085795Z","iopub.execute_input":"2023-02-14T05:17:11.086243Z","iopub.status.idle":"2023-02-14T05:17:35.235091Z","shell.execute_reply.started":"2023-02-14T05:17:11.086190Z","shell.execute_reply":"2023-02-14T05:17:35.233678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"学習曲線の表示","metadata":{}},{"cell_type":"code","source":"plt.plot(history.history['loss'])\nplt.plot(history.history['val_loss'])\nplt.title('Model loss')\nplt.ylabel('Loss')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Validation'], loc='best')\nplt.show()\n\nplt.plot(history.history['accuracy'])\nplt.plot(history.history['val_accuracy'])\nplt.title('Model accuracy')\nplt.ylabel('Accuracy')\nplt.xlabel('Epoch')\nplt.legend(['Train', 'Validation'], loc='best')\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 推論する際にメモリ不足になる可能性があるので、メモリ処理（不要かも）","metadata":{}},{"cell_type":"code","source":"import gc\n\ndel train_data, val_data\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport pandas as pd\n\nprint(pd.DataFrame([[val for val in dir()], [sys.getsizeof(eval(val)) for val in dir()]],\n                   index=['name','size']).T.sort_values('size', ascending=False).reset_index(drop=True))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 予測の実施\n\n予測対象データに対して予測を実施\n\n学習画像を白黒でおこなっているので、テスト画像も白黒に返還する","metadata":{}},{"cell_type":"code","source":"import cv2\nimport os\nimport csv\nimport glob\n\nimport numpy as np\n\ntest_dir = '/kaggle/input/104-flowers-garden-of-eden/jpeg-192x192/test'\n\ntest_flist = glob.glob(test_dir+'/*.jpeg', recursive=True)\n\nid_list = []\npredict_list = []\nfor file_name in test_flist:\n    test_img = cv2.imread(file_name)\n    test_img = cv2.cvtColor(test_img, cv2.COLOR_BGR2GRAY)\n    test_img = cv2.resize(test_img, target_size)\n    test_img = np.expand_dims(test_img, axis=2)\n    x = np.expand_dims(test_img, axis=0)/255.\n\n    id_list.append(os.path.splitext(os.path.basename(file_name))[0])\n    test_predict = model.predict(x)\n\n    predict_list.append(test_predict)\n        \npredict_label_list = np.concatenate(predict_list).argmax(axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-02-14T05:32:07.564324Z","iopub.execute_input":"2023-02-14T05:32:07.564686Z","iopub.status.idle":"2023-02-14T05:37:18.775315Z","shell.execute_reply.started":"2023-02-14T05:32:07.564656Z","shell.execute_reply":"2023-02-14T05:37:18.773860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 提出ファイルの作成","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nsubmission = pd.DataFrame(data={'id':id_list, 'label':predict_label_list.tolist()})\nsubmission","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"./submission.csv\", index=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}