{"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","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_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='rgb', classes=classes, batch_size=batch_size,\n    target_size=target_size)\nval_data = generator.flow_from_directory(\n    val_dir, color_mode='rgb', 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-01-05T00:19:56.743401Z","iopub.execute_input":"2023-01-05T00:19:56.743878Z","iopub.status.idle":"2023-01-05T00:19:58.801344Z","shell.execute_reply.started":"2023-01-05T00:19:56.743829Z","shell.execute_reply":"2023-01-05T00:19:58.799672Z"},"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_count":null,"outputs":[]},{"cell_type":"markdown","source":"# モデルの定義\n\n`keras.applications`の学習済みモデルを利用する","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\ninput_tensor = Input(shape=(*target_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='categorical_crossentropy',\n          optimizer=sdg, metrics=['accuracy'])","metadata":{},"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_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予測対象データに対して予測を実施","metadata":{}},{"cell_type":"code","source":"import os\nimport csv\nimport glob\nfrom PIL import Image\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 = Image.open(file_name)\n    test_img = test_img.resize(target_size)\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_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":[]}]}