{"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":"<div style=\"color:white;\n       display:fill;\n       border-radius:10px;\n       background-color:#13E18D;\n       font-size:200%;\n       font-family:Verdana;\n       letter-spacing:0.9px;\n       text-align:center;\n       position:relative;\">\n    <p style=\"padding: 1px;\n          color:white;\">\n        <h1 id=\"meta\">\n        Toxic Plant Classification </h1>\n    </p>\n</div>\n\n## MultiOutput Classification Model\nHans Elliott","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np \nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom prettytable import PrettyTable\nimport random\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '2' ##silence tensorflow info&warning messages\n#For working with images:\nimport PIL\nimport cv2\n# Modeling\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nfrom keras.callbacks import ModelCheckpoint, EarlyStopping\n# Eval\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix\n# Weights & Biases\nimport wandb\nfrom wandb.keras import WandbCallback\nfrom kaggle_secrets import UserSecretsClient","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-13T19:21:17.506236Z","iopub.execute_input":"2022-08-13T19:21:17.508030Z","iopub.status.idle":"2022-08-13T19:21:17.519408Z","shell.execute_reply.started":"2022-08-13T19:21:17.507991Z","shell.execute_reply":"2022-08-13T19:21:17.517820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Random seed\nr_seed = 83","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:17.522514Z","iopub.execute_input":"2022-08-13T19:21:17.523606Z","iopub.status.idle":"2022-08-13T19:21:17.533077Z","shell.execute_reply.started":"2022-08-13T19:21:17.523560Z","shell.execute_reply":"2022-08-13T19:21:17.531035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# GPU ACCELERATOR ENABLED\ndevice_name = tf.test.gpu_device_name()\nif \"GPU\" not in device_name:\n    print(\"GPU device not found\")\nelse:\n    print('Found GPU at: {}'.format(device_name))","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:17.534916Z","iopub.execute_input":"2022-08-13T19:21:17.536376Z","iopub.status.idle":"2022-08-13T19:21:20.967721Z","shell.execute_reply.started":"2022-08-13T19:21:17.536333Z","shell.execute_reply":"2022-08-13T19:21:20.966192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"color:white;\n       display:fill;\n       border-radius:10px;\n       background-color:#5642C5;\n       font-size:200%;\n       font-family:Verdana;\n       letter-spacing:0.9px;\n       text-align:center;\n       position:relative;\">\n    <p style=\"padding: 1px;\n          color:white;\">\n        <h1 id=\"meta\">\n        Metadata </h1>\n    </p>\n</div>","metadata":{}},{"cell_type":"code","source":"# Load metadata\nmeta = pd.read_csv(\"../input/tpcherbarium-merge/full_meta.csv\")\nmeta.info()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:20.972245Z","iopub.execute_input":"2022-08-13T19:21:20.973609Z","iopub.status.idle":"2022-08-13T19:21:21.030999Z","shell.execute_reply.started":"2022-08-13T19:21:20.973574Z","shell.execute_reply":"2022-08-13T19:21:21.029504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.032853Z","iopub.execute_input":"2022-08-13T19:21:21.033626Z","iopub.status.idle":"2022-08-13T19:21:21.064173Z","shell.execute_reply.started":"2022-08-13T19:21:21.033579Z","shell.execute_reply":"2022-08-13T19:21:21.062807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Adjust metadata for training   \nThe `flow_from_dataframe()` function requires your metadata has a column specifying the path to an image (which we have already).  \nAdditionally, we need a binary outcome variable for toxicity (whether a plant is toxic or not).  \nTo use this process for a multi-output model, we need to one-hot encode the categorical labels as well.  ","metadata":{}},{"cell_type":"code","source":"# Add toxicity dummy, which is a 1 for toxic plants and 0 for nontoxic - binary\nmeta['toxicity'] = int(1)\nmeta.loc[meta.class_id==5, 'toxicity'] = int(0)\n\nmeta[meta.toxicity==int(1)]['class_id'].unique() ##class 5 is nontoxic","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.066211Z","iopub.execute_input":"2022-08-13T19:21:21.067020Z","iopub.status.idle":"2022-08-13T19:21:21.082959Z","shell.execute_reply.started":"2022-08-13T19:21:21.066975Z","shell.execute_reply":"2022-08-13T19:21:21.081162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add sparse categorical labels specifying species\nstr_class_id = [str(c) for c in meta['class_id']]\nmeta['label'] = str_class_id","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.085657Z","iopub.execute_input":"2022-08-13T19:21:21.086524Z","iopub.status.idle":"2022-08-13T19:21:21.095560Z","shell.execute_reply.started":"2022-08-13T19:21:21.086478Z","shell.execute_reply":"2022-08-13T19:21:21.093811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# One hot encode the species categorical labels\nonehot_labels = tf.keras.utils.to_categorical(meta['label'].values)\nmeta['onehot_label'] = onehot_labels.tolist()\nmeta.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.097749Z","iopub.execute_input":"2022-08-13T19:21:21.098289Z","iopub.status.idle":"2022-08-13T19:21:21.127010Z","shell.execute_reply.started":"2022-08-13T19:21:21.098247Z","shell.execute_reply":"2022-08-13T19:21:21.125273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Image counts per class\nprint(\"Total Images:\", len(meta))\n\nt = PrettyTable([\"Class\", \"Images\"])\nfor i in range(6):\n    if i == 5: c = \"(nontoxic) 5\" \n    else: c = i\n    t.add_row([c, len(meta[meta.class_id == i])])\nt.add_row([\"toxic\", len(meta[meta.toxicity == 1])])\nt","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.129619Z","iopub.execute_input":"2022-08-13T19:21:21.130865Z","iopub.status.idle":"2022-08-13T19:21:21.155795Z","shell.execute_reply.started":"2022-08-13T19:21:21.130831Z","shell.execute_reply":"2022-08-13T19:21:21.153902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training/Testing Split\nStratified by class - save 50% of images per class for test data.","metadata":{}},{"cell_type":"code","source":"train_test_split = 0.75\ntrain_meta = meta.groupby(\"class_id\", group_keys=False).apply(\n    lambda x: x.sample(frac=train_test_split, random_state=r_seed)\n)\ntest_meta = meta.loc[meta.index.difference(train_meta.index), :]","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.161053Z","iopub.execute_input":"2022-08-13T19:21:21.161444Z","iopub.status.idle":"2022-08-13T19:21:21.192561Z","shell.execute_reply.started":"2022-08-13T19:21:21.161412Z","shell.execute_reply":"2022-08-13T19:21:21.191173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"train_meta.shape {train_meta.shape}\")\nprint(f\"test_meta.shape {test_meta.shape}\")\n\nt = PrettyTable([\"Class\", \"Train Images\", \"Test Images\"])\nfor i in range(6):\n    t.add_row([i, len(train_meta[train_meta.class_id == i]), len(test_meta[test_meta.class_id == i])])\nt.add_row([\"toxic\", len(train_meta[train_meta.toxicity==1]), len(test_meta[test_meta.toxicity==1])])\nt","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.194856Z","iopub.execute_input":"2022-08-13T19:21:21.195676Z","iopub.status.idle":"2022-08-13T19:21:21.220588Z","shell.execute_reply.started":"2022-08-13T19:21:21.195629Z","shell.execute_reply":"2022-08-13T19:21:21.219194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# reset indices\ntrain_meta = train_meta.reset_index()\ntest_meta = test_meta.reset_index()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.222607Z","iopub.execute_input":"2022-08-13T19:21:21.223421Z","iopub.status.idle":"2022-08-13T19:21:21.232695Z","shell.execute_reply.started":"2022-08-13T19:21:21.223375Z","shell.execute_reply":"2022-08-13T19:21:21.231193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Null Classifier","metadata":{}},{"cell_type":"code","source":"print(\"Accuracy if...\")\nprint(\"Guess toxic for all images:\", len(train_meta[train_meta.toxicity==1])/len(train_meta))\nprint(\"Guess nontoxic for all images:\", len(train_meta[train_meta.toxicity==0])/len(train_meta))\nprint(\"Guess most common class (0) for all species:\", len(train_meta[train_meta.class_id==0])/len(train_meta))","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.235150Z","iopub.execute_input":"2022-08-13T19:21:21.235691Z","iopub.status.idle":"2022-08-13T19:21:21.251672Z","shell.execute_reply.started":"2022-08-13T19:21:21.235647Z","shell.execute_reply":"2022-08-13T19:21:21.249920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"color:white;\n       display:fill;\n       border-radius:10px;\n       background-color:#5642C5;\n       font-size:200%;\n       font-family:Verdana;\n       letter-spacing:0.9px;\n       text-align:center;\n       position:relative;\">\n    <p style=\"padding: 1px;\n          color:white;\">\n        <h1 id=\"prep\">\n        Image Preprocessing </h1>\n    </p>\n</div>","metadata":{}},{"cell_type":"code","source":"# Preprocesser function (req: cv2, numpy)\ndef preprocess_image(img, bright_threshold=0.25, bright_value=30, image_size=(199,199), resize=True, recolor=True):\n    \"\"\"\n    Applies a standardized preprocessing procedure to every image before it is passed into the model.\n    Notes:\n    If using tf.keras.preprocessing.image.ImageDataGenerator (with either flow_from_dataframe or flow_from_directory),\n    do not specify the image size (ie, resize=False) or recolor the images (recolor=False). Resizing will cause issues, so \n    instead specify the image size in the flow_from_ function. Additionally, the flow_from_ functions automatically read in images\n    as RGB (whereas cv2 reads them as BGR), so recoloring the images will have the opposite effect.\n    \"\"\"\n    #Recolor\n    if recolor:\n        im = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    else:\n        im = img\n    #Adjust brightness\n    hsv = cv2.cvtColor(im, cv2.COLOR_RGB2HSV)\n    mean_val = np.mean(hsv[:,:,2])/255 #as a percentage of maximum pixel value\n    if mean_val <= bright_threshold:\n        h, s, v = cv2.split(hsv)\n        lim = 255 - bright_value\n        v[v > lim] = 255\n        v[v <= lim] += bright_value\n        final_hsv = cv2.merge((h, s, v))\n        im = cv2.cvtColor(final_hsv, cv2.COLOR_HSV2RGB)\n    #Resize\n    if resize:\n        im = cv2.resize(im, image_size, interpolation = cv2.INTER_AREA)\n    else:\n        im = im\n    #Central Crop\n    #w = image_size[1]\n    #h = image_size[0]\n    #center = tuple([int(x/2) for x in im.shape])\n    #x = center[1] - w/2\n    #y = center[0] - h/2\n    #crop_img = im[int(y):int(y+h), int(x):int(x+w)]\n        \n    return im","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:21.254803Z","iopub.execute_input":"2022-08-13T19:21:21.256195Z","iopub.status.idle":"2022-08-13T19:21:21.268148Z","shell.execute_reply.started":"2022-08-13T19:21:21.256147Z","shell.execute_reply":"2022-08-13T19:21:21.266395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Examples\nimgs = []\nslangs = []\n## randomly select images from the metadata, read images and add to list\nrand_indices = [random.randint(0, len(meta)) for i in range(9)]\nfor idx in rand_indices:\n    imgs.append( cv2.imread(meta.loc[idx, 'path']) )\n    slangs.append( meta.loc[idx, 'slang'] )\n\nfig, axes = plt.subplots(3, 3, figsize=(16,16))\naxes = axes.flatten()\n# For each image, apply the preprocessor function and plot image to a subplot\nfor img, ax, slang in zip(imgs, axes, slangs):\n    img = preprocess_image(img, image_size=(199, 199), resize=True, recolor=True) #PREPROCESSING\n    ax.imshow(img)\n    ax.axis('off')\n    ax.set_title(slang)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:21:32.491412Z","iopub.execute_input":"2022-08-13T19:21:32.491834Z","iopub.status.idle":"2022-08-13T19:21:34.128856Z","shell.execute_reply.started":"2022-08-13T19:21:32.491800Z","shell.execute_reply":"2022-08-13T19:21:34.127590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"color:white;\n       display:fill;\n       border-radius:10px;\n       background-color:#5642C5;\n       font-size:200%;\n       font-family:Verdana;\n       letter-spacing:0.9px;\n       text-align:center;\n       position:relative;\">\n    <p style=\"padding: 1px;\n          color:white;\">\n        <h1 id=\"model\">\n        Model </h1>\n    </p>\n</div>  \n\nThis is a multi-output model: it will simultaneously learn to predict if an image is toxic or nontoxic (binary classification) and what species of plant the image is (multi-class classification).  \nThe toxicity prediction will come first in the sequence and be used as an input into the final species prediction layer. ","metadata":{}},{"cell_type":"markdown","source":"<a id=\"section-config\"></a>\n## Configuration","metadata":{}},{"cell_type":"code","source":"config = dict(\n    model_id = 'ToxicToSpecies-MultiOutput-Resnet101',\n    train_test_split = train_test_split,\n    # Image processing\n    image_size = (192, 192),\n    # Model architecture\n    resnet = True,\n    conv1 = True,\n    conv2 = False,\n    conv1_filters = 32,\n    conv2_filters = 64,\n    conv1_kernel = (2,2),\n    conv2_kernel = (3,3),\n    conv1_strides = (1,1),\n    conv2_strides = (2,2),\n    maxpool1_size = (2,2),\n    maxpool2_size = (2,2),\n    dense1_size=1860,\n    dense2_size=None,\n    #Initialization\n    relu_init = \"He\",\n    #Regularization\n    dropout1 = 0.45,\n    dropout2 = 0.65,\n    l1_reg=0.0,\n    l2_reg=0.0,\n    # Optimizer\n    learn_rate = 0.001,\n    lr_decay = 0.1,\n    beta_1 = 0.9,\n    beta_2 = 0.999,\n    epsilon = 1e-07,\n    amsgrad = True,\n    # Training\n    epochs = 200,\n    batch_size = 48,\n    toxicity_lossweight = 0.85,\n    species_lossweight = 1.0,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:48:24.694205Z","iopub.execute_input":"2022-08-13T19:48:24.694641Z","iopub.status.idle":"2022-08-13T19:48:24.708183Z","shell.execute_reply.started":"2022-08-13T19:48:24.694608Z","shell.execute_reply":"2022-08-13T19:48:24.706492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"0.001 * (1/ (1 + 0.1*100))\n#learn_rate * (1 / (1 + lr_decay*iteration))","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:39:36.748945Z","iopub.execute_input":"2022-08-13T19:39:36.749636Z","iopub.status.idle":"2022-08-13T19:39:36.761143Z","shell.execute_reply.started":"2022-08-13T19:39:36.749579Z","shell.execute_reply":"2022-08-13T19:39:36.759320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Login to Wandb\nuser_secrets = UserSecretsClient()\n# I have saved my API key as \"wandb_api\" (Add-ons -> Secrets)\nwandb_api = user_secrets.get_secret(\"wandb_api\") \nwandb.login(key=wandb_api)\n\n# Initialize Wandb run\nwandb.init(anonymous='allow', project=\"toxic-plants\", config=config)\nconfig = wandb.config","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Architecture\nFor the MultiOutput model, I use Keras's 'Functional API'.","metadata":{}},{"cell_type":"code","source":"def build_multi_model(config):\n    # Prep\n    input_shape = (*config[\"image_size\"], 3)\n    ## Define Regularization\n    l1_reg = tf.keras.regularizers.L1(l1=config['l1_reg'])\n    l2_reg = tf.keras.regularizers.L2(l2=config['l2_reg'])\n    l1l2_reg = tf.keras.regularizers.L1L2(l1=config['l1_reg'], l2=config['l2_reg'])\n    ## Define ReLU Weight initialization\n    if config['relu_init'] == 'He': init = tf.keras.initializers.HeNormal(seed=r_seed)\n\n    # Structure\n    ## Input layer\n    img_input = tf.keras.Input(input_shape)\n    x = img_input\n    ## ResNet101\n    if config[\"resnet\"]:\n        x = tf.keras.applications.ResNet101(\n            input_shape=input_shape,\n            weights=\"imagenet\",\n            include_top=False ##do not include a full-dense layer at top of resnet\n        )(x)\n    ## Conv 1\n    if config[\"conv1\"]:\n        x = tf.keras.layers.Conv2D(filters=config['conv1_filters'], kernel_size=config['conv1_kernel'], \n                                     strides=config['conv1_strides'], padding='same', activation = 'relu',\n                                     kernel_regularizer=l1l2_reg, bias_regularizer=l1l2_reg,\n                                     kernel_initializer=init, name='Conv1_1')(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.Conv2D(filters=config['conv1_filters'], kernel_size=config['conv1_kernel'], \n                                     strides=config['conv1_strides'], padding='same', activation = 'relu',\n                                     kernel_regularizer=l1l2_reg, bias_regularizer=l1l2_reg,\n                                     kernel_initializer=init)(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.MaxPool2D(pool_size=config['maxpool1_size'])(x)\n        x = tf.keras.layers.Dropout(config['dropout1'])(x)\n    # Conv 2\n    if config[\"conv2\"]:\n        x = tf.keras.layers.Conv2D(filters=config['conv2_filters'], kernel_size=config['conv2_kernel'], \n                                     strides=config['conv2_strides'], padding='same', activation = 'relu',\n                                     kernel_regularizer=l1l2_reg, bias_regularizer=l1l2_reg, \n                                     kernel_initializer=init, name='Conv2_1')(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.Conv2D(filters=config['conv2_filters'], kernel_size=config['conv2_kernel'], \n                                     strides=config['conv2_strides'], padding='same', activation = 'relu',\n                                     kernel_regularizer=l1l2_reg, bias_regularizer=l1l2_reg, \n                                     kernel_initializer=init)(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.MaxPool2D(pool_size=config['maxpool2_size'])(x)\n        x = tf.keras.layers.Dropout(config['dropout2'])(x)\n    ## Dense Layers\n    x = tf.keras.layers.Flatten()(x)\n    x = tf.keras.layers.Dense(config['dense1_size'], activation='relu',\n                               kernel_initializer=init, name='Dense1')(x)\n    if config[\"dense2_size\"] is not None:\n        x = tf.keras.layers.Dense(config['dense2_size'], activation='relu',\n                                   kernel_initializer=init, name='Dense2')(x)\n    # Output Layers\n    ## Toxic or Nontoxic (binary)\n    o_toxic = tf.keras.layers.Dense(1, activation='sigmoid', name='toxicity', \n                                    kernel_initializer=tf.keras.initializers.GlorotUniform(seed=r_seed))(x)\n    o_concat = tf.keras.layers.concatenate([o_toxic, x])\n    ## Species Classifier\n    o_species = tf.keras.layers.Dense(6, activation='softmax', name='species', \n                                      kernel_initializer=tf.keras.initializers.GlorotUniform(seed=r_seed))(o_concat)\n    \n    # Finalize\n    model = tf.keras.models.Model(inputs=img_input, outputs=[o_toxic, o_species])   \n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:48:30.736424Z","iopub.execute_input":"2022-08-13T19:48:30.736839Z","iopub.status.idle":"2022-08-13T19:48:30.761660Z","shell.execute_reply.started":"2022-08-13T19:48:30.736806Z","shell.execute_reply":"2022-08-13T19:48:30.759928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss, Optimizer, Metrics","metadata":{}},{"cell_type":"code","source":"# Build Model\nmodel = build_multi_model(config)\ninit_weights = model.get_weights() #save random weights\n\n\n# Loss\nloss_binary = tf.keras.losses.BinaryCrossentropy(from_logits=False)\nloss_categ = tf.keras.losses.CategoricalCrossentropy(from_logits=False)\n# Optimizer \nlr_schedule = tf.keras.optimizers.schedules.ExponentialDecay(initial_learning_rate=config['learn_rate'],\n                                                             decay_steps=config['epochs'],\n                                                             decay_rate=config['lr_decay'])\noptim_adam = tf.keras.optimizers.Adam(learning_rate=lr_schedule,\n                                      beta_1=config['beta_1'], beta_2=config['beta_2'],\n                                      epsilon=config['epsilon'],\n                                      amsgrad=config['amsgrad'])\n# Metric Tracking\nmetric_list = [\n    tf.keras.metrics.Accuracy(name='acc'),\n    #tf.keras.metrics.TopKCategoricalAccuracy(k=topk, name=f'top_{topk}_acc'),\n]\n\n\n# Compile\nlosses = {'toxicity': loss_binary,\n          'species': loss_categ}\nloss_weights = {'toxicity': config[\"toxicity_lossweight\"], \n                'species': config[\"species_lossweight\"]}\n\nmodel.compile(optimizer=optim_adam,\n              loss=losses, \n              loss_weights=loss_weights,\n              metrics=metric_list)\n\n\n#View model\nmodel.summary()\ntf.keras.utils.plot_model(model, to_file='multioutput-mod_plt.png', \n                          show_shapes=True, show_layer_names=True, dpi=60)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:48:32.779960Z","iopub.execute_input":"2022-08-13T19:48:32.781437Z","iopub.status.idle":"2022-08-13T19:48:38.365900Z","shell.execute_reply.started":"2022-08-13T19:48:32.781387Z","shell.execute_reply":"2022-08-13T19:48:38.364213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add model info to config\nconv_layers = BN_layers = pool_layers = dropout_layers = dense_layers = 0\nfor layer in model.layers:\n    name = layer.__class__.__name__\n    if name == \"Conv2D\":\n        conv_layers += 1\n    elif name == \"BatchNormalization\":\n        BN_layers += 1\n    elif name == \"MaxPooling2D\":\n        pool_layers += 1\n    elif name == \"Dropout\":\n        dropout_layers += 1\n    elif name == \"Dense\":\n        dense_layers += 1\ndense_layers -= 2 ##subtract output layers from dense count\n\nconfig[\"Layers\"] = {\"Resnet\":config[\"resnet\"], \"Conv_Layers\":conv_layers, \"BatchNorm_layers\":BN_layers, \n                    \"MaxPool_layers\":pool_layers, \"Dropout_layers\":dropout_layers, \n                    \"Dense_layers\":dense_layers}","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:40:06.094296Z","iopub.execute_input":"2022-08-13T19:40:06.096535Z","iopub.status.idle":"2022-08-13T19:40:06.107531Z","shell.execute_reply.started":"2022-08-13T19:40:06.096472Z","shell.execute_reply":"2022-08-13T19:40:06.105834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ImageDataGenerator","metadata":{}},{"cell_type":"code","source":"# Adjust preprocessing function for the DataGen\ndef _prep_img(img):\n    return preprocess_image(img, \n                            bright_threshold=0.25, bright_value=30,\n                            image_size=config['image_size'], \n                            resize=False, recolor=False) ##specify with flow_from_dataframe instead\n\n\n# Image DataGenerator\ntrain_datagen = tf.keras.preprocessing.image.ImageDataGenerator(\n    preprocessing_function = _prep_img,\n    rescale = 1.0/255, #scale pixels values to [0, 1]\n    #augmentation\n    rotation_range=30,\n    width_shift_range=0.1,\n    height_shift_range=0.1,\n    shear_range=0.1,\n    zoom_range=0.2,\n    horizontal_flip=True,\n    fill_mode='reflect',\n     #cval = 235,   ##a bright constant value for the fill mode \"constant\"\n    validation_split=0.2\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:32:15.876429Z","iopub.execute_input":"2022-08-13T19:32:15.876845Z","iopub.status.idle":"2022-08-13T19:32:15.887183Z","shell.execute_reply.started":"2022-08-13T19:32:15.876812Z","shell.execute_reply":"2022-08-13T19:32:15.885279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Image Augmentation Test\nex = cv2.imread(train_meta.loc[12, 'path'])\nplt.figure(figsize = (2,2))\nplt.imshow(ex)\nplt.show()\n\n# Image after augmentation, preprocessing is applied\nx = ex\nx = x.reshape((1,) + x.shape)\naug_x = train_datagen.flow(x)\naug_images = [next(aug_x)[0] for i in range(12)]\n\nfig, axes = plt.subplots(2,5,figsize=(8,8)) ##3 img rows, 4 img cols\naxes = axes.flatten()\nfor img, ax in zip(aug_images, axes): ##zip image to its subplot\n    ax.imshow(img)\n    ax.axis('off')\n    \nplt.subplots_adjust(wspace=0.1, hspace=0.1)\nplt.suptitle(\"Example image after processing & augmentation (pre-recoloring/resizing)\", fontsize=12)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:32:21.325238Z","iopub.execute_input":"2022-08-13T19:32:21.325686Z","iopub.status.idle":"2022-08-13T19:32:23.860484Z","shell.execute_reply.started":"2022-08-13T19:32:21.325653Z","shell.execute_reply":"2022-08-13T19:32:23.859180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"color:white;\n       display:fill;\n       border-radius:10px;\n       background-color:#5642C5;\n       font-size:200%;\n       font-family:Verdana;\n       letter-spacing:0.9px;\n       text-align:center;\n       position:relative;\">\n    <p style=\"padding: 1px;\n          color:white;\">\n        <h1 id=\"train\">\n        Training </h1>\n    </p>\n</div>\n\n[←config](#section-config)","metadata":{}},{"cell_type":"code","source":"my_callbacks = [\n    WandbCallback(log_weights=True), ##sends metrics to Wandb\n    EarlyStopping(\n        monitor='val_species_acc', \n        patience=50,\n        min_delta=0.01,\n        mode='max',\n        verbose=1),\n    ModelCheckpoint(\n        filepath='./tpc-simple_cnn.h5',\n        monitor='val_species_acc', \n        save_best_only=True, \n        initial_value_threshold=0.80,\n        mode='max',\n        verbose=0)\n]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Reset model weights to randomly initialized\nreset_model = lambda model, weights: model.set_weights(weights) \nreset_model(model, init_weights)\n\n# TRAIN\nhistory = model.fit(\n    train_datagen.flow_from_dataframe(dataframe=train_meta,\n                                      x_col='path',\n                                      y_col=['toxicity', 'onehot_label'],\n                                      class_mode='multi_output',\n                                      target_size=config['image_size'],\n                                      interpolation=\"bilinear\",\n                                      color_mode=\"rgb\",\n                                      batch_size=config['batch_size'],\n                                      shuffle=True,\n                                      seed=r_seed,\n                                      subset=\"training\"\n                                     ),\n    validation_data = train_datagen.flow_from_dataframe(dataframe=train_meta,\n                                                        x_col='path',\n                                                        y_col=['toxicity', 'onehot_label'],\n                                                        class_mode='multi_output',                                                        \n                                                        target_size=config['image_size'],\n                                                        interpolation=\"bilinear\",\n                                                        color_mode=\"rgb\",                                                      \n                                                        batch_size=config['batch_size'],\n                                                        shuffle=True,\n                                                        seed=r_seed,\n                                                        subset=\"validation\"\n                                                       ),\n    epochs = config['epochs'],\n    callbacks=my_callbacks\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-13T19:40:16.911069Z","iopub.execute_input":"2022-08-13T19:40:16.911534Z","iopub.status.idle":"2022-08-13T19:46:49.976645Z","shell.execute_reply.started":"2022-08-13T19:40:16.911501Z","shell.execute_reply":"2022-08-13T19:46:49.975152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish() ##tell wandb to stop tracking the session","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plotting the fit history\ndef plot_fit_history(history, y_name=None):\n    if y_name is not None:\n        y_name_ = y_name+\"_\"\n    else:\n        y_name_=''\n        y_name=''\n    fig, ax = plt.subplots(nrows=1, ncols=2, figsize=(15,5))\n    \n    ax[0].plot(history.history[f'{y_name_}loss'], label='loss')\n    ax[0].plot(history.history[f'val_{y_name_}loss'], label='val loss')\n    ax[0].set_title(y_name + \" Loss\")\n    ax[0].legend()\n\n    ax[1].plot(history.history[f'{y_name_}acc'], label='acc')\n    ax[1].plot(history.history[f'val_{y_name_}acc'], label='val acc')\n    ax[1].set_title(y_name+\" Accuracy\")\n    ax[1].legend()\n\n    return plt.show()\n    \n\nplot_fit_history(history, y_name=\"toxicity\")\nplot_fit_history(history, y_name=\"species\")","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:01:18.536790Z","iopub.execute_input":"2022-08-12T21:01:18.537703Z","iopub.status.idle":"2022-08-12T21:01:19.253750Z","shell.execute_reply.started":"2022-08-12T21:01:18.537651Z","shell.execute_reply":"2022-08-12T21:01:19.252806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"# Use data generator to flow in test images in batches & apply the same preprocessing for prediction \ntest_datagen = tf.keras.preprocessing.image.ImageDataGenerator(\n    rescale=1.0/255,\n    preprocessing_function=_prep_img\n).flow_from_dataframe(\n    dataframe=test_meta,\n    x_col='path',\n    y_col=None,\n    class_mode=None,\n    target_size=config['image_size'],\n    interpolation=\"bilinear\",\n    batch_size=128,\n    shuffle=False\n)\ny_pred = model.predict(test_datagen)","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:01:21.435451Z","iopub.execute_input":"2022-08-12T21:01:21.436076Z","iopub.status.idle":"2022-08-12T21:01:55.980760Z","shell.execute_reply.started":"2022-08-12T21:01:21.436040Z","shell.execute_reply":"2022-08-12T21:01:55.979523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Unpack the 2 prediction arrays produced from model.predict\ntox_pred = y_pred[0]\nspecies_pred = y_pred[1]\n\nprint(tox_pred.shape)\nprint(species_pred.shape)","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:09.097683Z","iopub.execute_input":"2022-08-12T21:02:09.098078Z","iopub.status.idle":"2022-08-12T21:02:09.105315Z","shell.execute_reply.started":"2022-08-12T21:02:09.098042Z","shell.execute_reply":"2022-08-12T21:02:09.104087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Determine class predictions\n##where prob > 0.5, predict toxic\ntox_pred = np.where(tox_pred > 0.5, 1, 0).reshape(len(test_meta),) \n##for species, pick class with highest prediced prob\nspecies_pred = np.argmax(species_pred, axis=1) ","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:11.067268Z","iopub.execute_input":"2022-08-12T21:02:11.067852Z","iopub.status.idle":"2022-08-12T21:02:11.073904Z","shell.execute_reply.started":"2022-08-12T21:02:11.067815Z","shell.execute_reply":"2022-08-12T21:02:11.072423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_species_labels = [pred for pred in species_pred]\npred_toxic_labels = [pred for pred in tox_pred]\ny_test_species = [int(tru) for tru in test_meta['label']]\ny_test_toxic = [int(tru) for tru in test_meta['toxicity']]\n\n\nprint(\"species pred:\", pred_species_labels[0:5])\nprint(\"species true:\", y_test_species[0:5])\nprint(\"\")\nprint(\"toxic pred:\", pred_toxic_labels[0:5])\nprint(\"toxic true:\", y_test_toxic[0:5])","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:13.585280Z","iopub.execute_input":"2022-08-12T21:02:13.585643Z","iopub.status.idle":"2022-08-12T21:02:13.594513Z","shell.execute_reply.started":"2022-08-12T21:02:13.585612Z","shell.execute_reply":"2022-08-12T21:02:13.593438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test Metrics","metadata":{}},{"cell_type":"code","source":"t = PrettyTable([\"Species Test Accuracy\", \"Toxicity Test Accuracy\"])\nt.add_row([round(accuracy_score(y_test_species, pred_species_labels), 6), \n           round(accuracy_score(y_test_toxic, pred_toxic_labels), 6)])\nt","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:15.453342Z","iopub.execute_input":"2022-08-12T21:02:15.454354Z","iopub.status.idle":"2022-08-12T21:02:15.466476Z","shell.execute_reply.started":"2022-08-12T21:02:15.454316Z","shell.execute_reply":"2022-08-12T21:02:15.465230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Species Confusion Matrix\ndf_cm = pd.DataFrame(confusion_matrix(y_test_species, pred_species_labels),\n                     index = [i for i in [0,1,2,3,4,\"True 5\"]],\n                     columns = [i for i in [\"Pred 0\",1,2,3,4,5]])\nprint(\"SPECIES\")\nprint(pd.read_csv(\"../input/toxic-plant-classification/tpc-imgs/metadata.csv\").loc[:,['class_id', 'slang']])\n\nplt.figure(figsize = (10,7))\nm = sns.heatmap(df_cm, annot=True)\nm.xaxis.set_ticks_position(\"top\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:16.941376Z","iopub.execute_input":"2022-08-12T21:02:16.941760Z","iopub.status.idle":"2022-08-12T21:02:17.493906Z","shell.execute_reply.started":"2022-08-12T21:02:16.941726Z","shell.execute_reply":"2022-08-12T21:02:17.492853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Toxicity Confusion Matrix\ndf_cm2 = pd.DataFrame(confusion_matrix(y_test_toxic, pred_toxic_labels),\n                     index = [i for i in [\"true nontoxic\", \"true toxic\"]],\n                     columns = [i for i in [\"pred nontoxic\", \"pred toxic\"]])\n\nprint(\"TOXICITY\")\nplt.figure(figsize = (10,7))\nm = sns.heatmap(df_cm2, annot=True)\nm.xaxis.set_ticks_position(\"top\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:19.099359Z","iopub.execute_input":"2022-08-12T21:02:19.099724Z","iopub.status.idle":"2022-08-12T21:02:19.752358Z","shell.execute_reply.started":"2022-08-12T21:02:19.099693Z","shell.execute_reply":"2022-08-12T21:02:19.751056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Break down of model predictions\nprint(f\"Total Test Images = {len(test_meta)}\")\n\nt = PrettyTable(['Class', 'Times Predicted', 'Actual Images'])\nfor i in range(7):\n    if i < 6:\n        t.add_row([i, len(species_pred[species_pred==i]), len(test_meta[test_meta.class_id==i])])\n    if i == 6:\n        t.add_row([\"-\",\"-\",\"-\"])\n        t.add_row([\"Toxic\", np.sum(tox_pred==1), len(test_meta[test_meta.toxicity==1])])\n\nt","metadata":{"execution":{"iopub.status.busy":"2022-08-12T21:02:25.128277Z","iopub.execute_input":"2022-08-12T21:02:25.128722Z","iopub.status.idle":"2022-08-12T21:02:25.175312Z","shell.execute_reply.started":"2022-08-12T21:02:25.128682Z","shell.execute_reply":"2022-08-12T21:02:25.174384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Notes:\n**MultiOutput Model Resources**  \nQ&A: https://github.com/keras-team/keras-preprocessing/issues/197  \nGist tutorial: https://gist.github.com/rragundez/8b1b9e190abef8ffb6d93f03f8e0e091 (best resource)  \nNotebook example with multiple output layers: https://www.kaggle.com/code/venkatkumar001/fgvc-9-starter-eda-tensorflow.  \n\n\nBatch Norm: https://github.com/christianversloot/machine-learning-articles/blob/main/how-to-use-batch-normalization-with-keras.md  \nHe Initialization for ReLU layers: https://github.com/christianversloot/machine-learning-articles/blob/main/he-xavier-initialization-activation-functions-choose-wisely.md  \n","metadata":{}}]}