{"cells":[{"metadata":{"_uuid":"58b7f63bc3fdb35dfce4a153022921c05f4b0834"},"cell_type":"markdown","source":"# Multi-Label Classification Example using Keras\n\nThis notebook is an example of a baseline model I trained. In this example, I used the tuning labels as training data, validation data, and test data. However the real baseline model used 1.7M training images from the Open Images Dataset and training labels created from concatenating thr train_human_labels, train_machine_labels, and train_bounding_boxes."},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2 as cv\nimport matplotlib.pyplot as plt\nimport os\nimport pickle\nfrom collections import defaultdict","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"from keras.applications.densenet import DenseNet121\nfrom keras.applications.densenet import preprocess_input\nfrom keras.optimizers import Adam, SGD\nfrom keras.models import Model, load_model\nfrom keras.layers import *\nfrom sklearn.model_selection import train_test_split\nfrom keras.callbacks import *\n\n# from keras.utils import multi_gpu_model","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"cfe503ce1ab7cc7c53e1a027f26dea7de2c1ee52"},"cell_type":"code","source":"DATA_DIR = '../input' #Or wherever your data is\n# Test with training the stage 1 tuning labels (come up with your own labels)\nCHALLENGE_DATA_DIR = DATA_DIR + '/'\nIMG_DIR = CHALLENGE_DATA_DIR + '/stage_1_test_images/stage_1_test_images'","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"204bcac12423d0c0e7fe69d3ce6c1d6a2cf40f14"},"cell_type":"markdown","source":"# Load our (fake) datatset\nThis dataset should actually be all the images inside train_human_labels, train_machine_labels,  and train_bounding_boxes\nBut for demo purposes we'll just use the tuning labels as the dataset"},{"metadata":{"trusted":true,"_uuid":"d8417778531086d34efb32bee42d400dfde7b565"},"cell_type":"code","source":"tuning_labels = pd.read_csv(CHALLENGE_DATA_DIR + '/tuning_labels.csv', \n                            names=['ImageID', 'Caption'],\n                            index_col=['ImageID'])\ntuning_labels.head()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"db5148d24e82d9128757a44df3e6f684cefa0080"},"cell_type":"markdown","source":"# Get list of unique classes in tuning dataset"},{"metadata":{"trusted":true,"_uuid":"07c0506df6312364ff923497db2418cd71fd1f4c"},"cell_type":"code","source":"tuning_labels_freq = defaultdict(int)\n\nfor r in tuning_labels['Caption']:\n    labels = r.split()\n    for l in labels:\n        tuning_labels_freq[l] += 1\n\ntuning_labels_list = list(tuning_labels_freq.keys())\nprint('Unique tuning labels', len(tuning_labels_list))\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d0b0c30a382d047e464c444150451bbba126a537"},"cell_type":"markdown","source":"# Prepare labels"},{"metadata":{"trusted":true,"_uuid":"c68859a392c510c90fd93dc4fe8c7d2406e1b3dc"},"cell_type":"code","source":"label_2_idx = {}\nidx_2_label = {}\nfor i,v in enumerate(tuning_labels_list):\n    label_2_idx[v] = i\n    idx_2_label[i] = v","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"5de2b5fce7c412e83e3483dda8a48fb1bcd12d68"},"cell_type":"code","source":"class_descriptions = pd.read_csv(CHALLENGE_DATA_DIR + '/class-descriptions.csv', index_col='label_code')\n\nclass_descriptions.loc['/m/0104x9kv']['description']","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"6f4f5829df99b28c0a6bbd6b61c31ba208c67664"},"cell_type":"code","source":"all_img_ids = list(tuning_labels.index.unique())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"50a8836d98431f12d80ad2dea3b1f11610537bf2"},"cell_type":"code","source":"train_ids, test_ids = train_test_split(all_img_ids, test_size=0.01, random_state=21)\ntrain_ids, valid_ids = train_test_split(train_ids, test_size=0.1, random_state=21)\n\nprint('Training on {} samples'.format(len(train_ids)))\nprint('Validating on {} samples'.format(len(valid_ids)))\nprint('Testing on {} samples'.format(len(test_ids)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"867df6c0443d8abbcfaefb84f64df6807ab6139a"},"cell_type":"code","source":"N_CLASSES = len(label_2_idx)\nBATCH_SIZE = 8\nINPUT_SIZE = 224\nprint(N_CLASSES)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"75f53214605291e56424f75d9e7ad5b8587d0f59"},"cell_type":"code","source":"def caption_2_one_hot(caption, n_classes=1, lookup_dict=None):\n    y = np.zeros((n_classes))\n    for w in caption.split():\n        idx = lookup_dict[w]\n        y[idx] = 1\n    return y","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b02c78cb8a52fff178a39930e38ea8ac34de9197"},"cell_type":"code","source":"def ImageDataGen(ids, df,\n                 lookup_dict=label_2_idx,\n                 n_classes=N_CLASSES,\n                 img_dir=IMG_DIR, input_size=INPUT_SIZE,\n                 bs=BATCH_SIZE, returnIds=False):\n    while True:\n        for start in range(0, len(ids), bs):\n            x_batch = []\n            y_batch = []\n            end = min(start+bs, len(ids))\n            sample = ids[start:end]\n            for img_id in sample:\n                img = cv.imread('{}/{}.jpg'.format(img_dir, img_id))\n                if img is not None:\n                    img = cv.resize(img, (input_size, input_size))\n                    img = preprocess_input(img.astype(np.float32))\n                    x_batch.append(img)\n                    caption = df.loc[img_id]['Caption']\n                    y = caption_2_one_hot(caption, n_classes=n_classes, lookup_dict=lookup_dict)\n                    y_batch.append(y)\n                    \n            x_batch = np.array(x_batch, np.float32)\n            y_batch = np.array(y_batch, np.float32)\n            \n            if returnIds:\n                yield x_batch, y_batch, sample\n            else:\n                yield x_batch, y_batch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b7cdb26162470ac3f7c382dd63db3de475128ccb"},"cell_type":"code","source":"test_gen = ImageDataGen(test_ids, tuning_labels)\ntest_batch = next(test_gen)\nfig = plt.figure(figsize=(20, 8))\nfor sample_idx in range(BATCH_SIZE):\n    ax = fig.add_subplot(3,3, sample_idx + 1)\n    ax.set_title(','.join([class_descriptions.loc[idx_2_label[i]]['description'] for i in (np.argwhere(test_batch[1][sample_idx]>0)).flatten()]))\n    ax.imshow(test_batch[0][sample_idx])\n    ax.set_axis_off()\nplt.show()\n    ","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"87d3e85be2033243bc0a533488e225d5fc099b4e"},"cell_type":"markdown","source":"# Define Model"},{"metadata":{"trusted":true,"_uuid":"526abf426b3cbaf29ae3bf6e6d0e643704d43396"},"cell_type":"code","source":"def ClsModel(n_classes=1, input_shape=(224,224,3)):\n    base_model = DenseNet121(weights=None, include_top=False, input_shape=input_shape)\n    x = AveragePooling2D(pool_size=(3,3), name='avg_pool')(base_model.output)\n    x = Flatten()(x)\n    x = Dense(1024, activation='relu', name='dense_post_pool')(x)\n    x = Dropout(0.5)(x)\n    output = Dense(n_classes, activation='sigmoid', name='predictions')(x)\n    model = Model(inputs=base_model.input, output=output)\n    return model","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"68aeefb95f560c1f02f799fea8cdd17b806fac09"},"cell_type":"code","source":"model = ClsModel(N_CLASSES)\nmodel.summary()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b5008ee83cf4d8739eab8037057a9c7bb3d4371f"},"cell_type":"code","source":"model.compile(optimizer=Adam(lr=0.001), loss='binary_crossentropy', metrics=['accuracy'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"eb2c21f8b005baa4f7123b7f69ed4d30262cf2ef"},"cell_type":"code","source":"model_checkpoint = ModelCheckpoint(('./densenet.{epoch:02d}.hdf5'),\n                                   monitor='val_loss',\n                                   verbose=1,\n                                   save_best_only=True,\n                                   save_weights_only=True)\n\nreduce_learning_rate = ReduceLROnPlateau(monitor='val_loss', factor=0.1,\n                                         patience=2, verbose=1)\n\ncallbacks = [model_checkpoint, reduce_learning_rate]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"bc9e327210148301715189250da97ee40575216d"},"cell_type":"code","source":"train_gen = ImageDataGen(train_ids, tuning_labels)\nvalid_gen = ImageDataGen(valid_ids, tuning_labels)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"7db2c8b546155f0d3029282ca1002144fb11686d"},"cell_type":"code","source":"model.fit_generator(generator=train_gen, \n                    epochs=25, \n                    steps_per_epoch=np.ceil(len(train_ids)/BATCH_SIZE),\n                   callbacks=callbacks,\n                    validation_data=valid_gen,\n                    validation_steps=np.ceil(len(valid_ids) / BATCH_SIZE))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"31e1dfd6751caad16712fe823d1491c00efa9060"},"cell_type":"markdown","source":"# Test Predict"},{"metadata":{"trusted":true,"_uuid":"bb14c2eba4fa2f0a278998fc420c4cd4b56d6298"},"cell_type":"code","source":"test_preds = model.predict(test_batch[0])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b39a869b62a5e39d0b5b2efa01e5b47ee93944fd"},"cell_type":"code","source":"fig = plt.figure(figsize=(20, 8))\npred_cutoff = 0.2\nfor sample_idx in range(BATCH_SIZE):\n    ax = fig.add_subplot(3,3, sample_idx + 1)\n    ax.set_title(','.join([class_descriptions.loc[idx_2_label[i]]['description'] for i in (np.argwhere(test_preds[sample_idx]>pred_cutoff)).flatten()]))\n    ax.imshow(test_batch[0][sample_idx])\n    ax.set_axis_off()\nplt.show()\n    ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"eed8deb929e6430a047d4c1594eab25ed19f2d89"},"cell_type":"markdown","source":""},{"metadata":{"trusted":true,"_uuid":"73f62f5b35b7655377461b9b1a25cf8760432716"},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}