{"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":"<img src=\"https://storage.googleapis.com/kaggle-competitions/kaggle/37333/logos/header.png?t=2022-06-29-00-47-20\">\n\n<h1><center>STRIP AI TPU TensorFlow - Training</center></h1>\n\nThis is a pretty basic training baseline for a quickstart with your experiments. It takes only **10 minutes** to train a model. Guaranteed.\n\nThe **[STRIP AI - 256x256 PNG Tiles][1]** dataset used here is prepared in a separate notebook: **[STRIP AI - EDA & Data Preparation][2]**. \n\n# Setup\n\n[1]: https://www.kaggle.com/datasets/nickuzmenkov/strip-ai-256x256-png-tiles\n[2]: https://www.kaggle.com/code/nickuzmenkov/strip-ai-eda-data-preparation","metadata":{"_uuid":"eed65e8d-e4af-46ac-b225-5c8275b70f82","_cell_guid":"4b845e3a-857b-4c52-8a86-450dcb11af9a","trusted":true}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/efficientnet-keras-dataset/efficientnet_kaggle')","metadata":{"_uuid":"fdae1408-da2c-4649-b8b7-f6194a0ac065","_cell_guid":"e240b01a-edfc-4bc8-ac7e-d55a8f28753f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T04:37:49.489416Z","iopub.execute_input":"2022-07-14T04:37:49.489786Z","iopub.status.idle":"2022-07-14T04:37:49.522169Z","shell.execute_reply.started":"2022-07-14T04:37:49.489684Z","shell.execute_reply":"2022-07-14T04:37:49.521071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport typing\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport plotly.graph_objects as go\nimport efficientnet.tfkeras as efn\nimport matplotlib.pyplot as plt\nfrom sklearn.utils import shuffle\nfrom kaggle_datasets import KaggleDatasets\nfrom sklearn.model_selection import KFold","metadata":{"_uuid":"038ac895-f8ab-4d39-a03f-4ceb972ff7b9","_cell_guid":"74f71145-f150-42a0-8b2a-952e70c5b952","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T04:57:28.525132Z","iopub.execute_input":"2022-07-14T04:57:28.525429Z","iopub.status.idle":"2022-07-14T04:57:28.689292Z","shell.execute_reply.started":"2022-07-14T04:57:28.525397Z","shell.execute_reply":"2022-07-14T04:57:28.688343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RANDOM_STATE = 42\nN_SPLITS = 5\nIMG_SIZE = 256\nGCS_PATH = KaggleDatasets().get_gcs_path(\"strip-ai-256x256-png-tiles\")\nVERBOSE = 1 if os.environ[\"KAGGLE_KERNEL_RUN_TYPE\"] == \"Interactive\" else 2\n\ntry:\n    TPU = tf.distribute.cluster_resolver.TPUClusterResolver()\n    tf.config.experimental_connect_to_cluster(TPU)\n    tf.tpu.experimental.initialize_tpu_system(TPU)\n    STRATEGY = tf.distribute.experimental.TPUStrategy(TPU)\n    BATCH_SIZE = 8 * STRATEGY.num_replicas_in_sync\nexcept Exception:\n    TPU = None\n    STRATEGY = tf.distribute.get_strategy()\n    BATCH_SIZE = 4\n\nprint(\"TensorFlow\", tf.__version__)\n\nif TPU is not None:\n    print(\"Using TPU v3-8\")\nelse:\n    print(\"Using GPU/CPU\")\n\nprint(\"Batch size:\", BATCH_SIZE)","metadata":{"_uuid":"552e7a48-3309-4384-a65d-57774515a8a8","_cell_guid":"8f61d9bf-8dac-48b3-b58a-78bb8e1b388b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T04:37:57.532492Z","iopub.execute_input":"2022-07-14T04:37:57.533109Z","iopub.status.idle":"2022-07-14T04:38:03.921037Z","shell.execute_reply.started":"2022-07-14T04:37:57.533076Z","shell.execute_reply":"2022-07-14T04:38:03.920257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def count_samples(filenames: typing.List[str]) -> int:\n    return sum(\n        int(os.path.basename(x).split(\".\")[0].split(\"-\")[-1]) for x in filenames\n    )\n\n\ndef transform(image: tf.Tensor) -> tf.Tensor:\n    image = tf.image.random_flip_left_right(image, seed=RANDOM_STATE)\n    image = tf.image.random_flip_up_down(image, seed=RANDOM_STATE)\n    if tf.random.uniform([], 0, 1, dtype=tf.float32) > 0.5:\n        image = tf.image.transpose(image)\n    image = tf.image.random_brightness(image, 0.2, seed=RANDOM_STATE)\n    image = tf.image.random_contrast(image, 0.8, 1.2, seed=RANDOM_STATE)\n    image = tf.image.random_hue(image, 0.2, seed=RANDOM_STATE)\n    image = tf.image.random_saturation(image, 0.8, 1.2, seed=RANDOM_STATE)\n    return image\n\n\ndef parse_sample(sample: tf.Tensor, augmented: bool) -> typing.Tuple[tf.Tensor, tf.Tensor]:\n    features = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"label\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    sample = tf.io.parse_single_example(sample, features)\n    \n    image = tf.image.decode_png(sample[\"image\"])\n    if augmented:\n        image = transform(image)\n    image = tf.reshape(image, (IMG_SIZE, IMG_SIZE, 3))\n    image = tf.cast(image, tf.float32) / 255.0\n    \n    laa = tf.cast(sample[\"label\"], tf.float32)\n    ce = 1.0 - laa\n    \n    return image, [ce, laa]\n\n\ndef get_dataset(\n    filenames: typing.List[str],\n    augmented: bool = True,\n    ordered: bool = False,\n    repeated: bool = True,\n    cached: bool = False,\n    distributed: bool = True,\n) -> tf.data.Dataset:\n    auto = tf.data.experimental.AUTOTUNE\n    dataset = tf.data.TFRecordDataset(filenames)\n    if not ordered:\n        ignore_order = tf.data.Options()\n        ignore_order.experimental_deterministic = False\n        dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(\n        lambda x: parse_sample(x, augmented),\n        num_parallel_calls=auto,\n    )\n    if not ordered:\n        dataset = dataset.shuffle(1024, seed=RANDOM_STATE)\n    if repeated:\n        dataset = dataset.repeat()\n    dataset = dataset.batch(BATCH_SIZE, drop_remainder=True)\n    if cached:\n        dataset = dataset.cache()\n    dataset = dataset.prefetch(auto)\n    if distributed:\n        return STRATEGY.experimental_distribute_dataset(dataset)\n    return dataset\n\n\ndef get_model() -> tf.keras.Model:\n    model = tf.keras.models.Sequential(name='EfficientNetB4')\n    model.add(efn.EfficientNetB4(\n        include_top=False,\n        input_shape=(IMG_SIZE, IMG_SIZE, 3),\n        weights='noisy-student',\n        pooling='avg'))\n    model.add(tf.keras.layers.Dense(2, activation=\"softmax\"))\n    model.compile(\n        loss=tf.keras.losses.CategoricalCrossentropy(),\n        optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3),\n        metrics=[\"accuracy\"],\n    )\n    \n    return model","metadata":{"_uuid":"71208075-95bb-477d-9b19-924d631c39a2","_cell_guid":"03060730-5d7b-4481-9112-e39c77fd16b4","collapsed":false,"_kg_hide-input":true,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T06:00:25.310304Z","iopub.execute_input":"2022-07-14T06:00:25.310661Z","iopub.status.idle":"2022-07-14T06:00:25.336933Z","shell.execute_reply.started":"2022-07-14T06:00:25.310601Z","shell.execute_reply":"2022-07-14T06:00:25.335768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Augmented train images look like this","metadata":{"_uuid":"b2fa4b7d-5991-40d2-bb1d-36bcad69992b","_cell_guid":"e8619f62-4294-44fc-9f82-e84f4af4ae54","trusted":true}},{"cell_type":"code","source":"filenames = tf.io.gfile.glob(os.path.join(GCS_PATH, \"tfrec\", \"0\", \"*.tfrec\"))\ndataset = get_dataset(filenames, distributed=False)\n\nfigure, axes = plt.subplots(5, 5, figsize=(15, 15))\naxes = np.ravel(axes)\n\nfor i, (image, label) in enumerate(dataset.unbatch().take(len(axes)).as_numpy_iterator()):\n    axes[i].imshow(image)\n    axes[i].set_title(\"LAA\" if label[1] else \"CE\")\n    axes[i].axis(\"off\")\n    \nfigure.show()","metadata":{"_uuid":"9ac23a34-b725-4cb8-804a-288a8f2fe04e","_cell_guid":"faf74d7d-8eb4-4c2d-a619-a9fe2ca986cf","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T06:00:26.099452Z","iopub.execute_input":"2022-07-14T06:00:26.099806Z","iopub.status.idle":"2022-07-14T06:00:33.684604Z","shell.execute_reply.started":"2022-07-14T06:00:26.099771Z","shell.execute_reply":"2022-07-14T06:00:33.683695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"_uuid":"171d4192-095e-4285-abc0-7e0590e6f5c5","_cell_guid":"e350557f-cec0-4540-b2d1-d2f02c462a07","trusted":true}},{"cell_type":"code","source":"kfold = KFold(n_splits=5, shuffle=True, random_state=RANDOM_STATE)\n\nfor i, (train_index, val_index) in enumerate(kfold.split(range(N_SPLITS))):\n    train_filenames = [\n        tf.io.gfile.glob(os.path.join(GCS_PATH, \"tfrec\", str(x), \"*.tfrec\"))\n        for x in train_index\n    ]\n    train_filenames = np.concatenate(train_filenames)\n    steps_per_epoch = count_samples(train_filenames) // BATCH_SIZE\n    train_dataset = get_dataset(train_filenames)\n    \n    val_filenames = [\n        tf.io.gfile.glob(os.path.join(GCS_PATH, \"tfrec\", str(x), \"*.tfrec\"))\n        for x in val_index\n    ]\n    val_filenames = np.concatenate(val_filenames)\n    validation_steps = count_samples(val_filenames) // BATCH_SIZE\n    val_dataset = get_dataset(\n        val_filenames,\n        augmented=False,\n        ordered=True,\n        repeated=False,\n        cached=True,\n    )\n    \n    with STRATEGY.scope():\n        model = get_model()\n        \n    metrics = model.fit(\n        train_dataset,\n        steps_per_epoch=steps_per_epoch,\n        validation_data=val_dataset,\n        validation_steps=validation_steps,\n        callbacks=[\n            tf.keras.callbacks.ReduceLROnPlateau(\n                monitor=\"val_accuracy\",\n                factor=0.1,\n                patience=3,\n                mode=\"max\",\n                min_lr=1e-6,\n                verbose=1,\n            ),\n            tf.keras.callbacks.EarlyStopping(\n                monitor=\"val_accuracy\",\n                patience=10,\n                mode=\"max\",\n                restore_best_weights=True,\n            ),\n        ],\n        epochs=100,\n        verbose=VERBOSE,\n    ).history\n    model.save(f\"model_{i}.h5\")\n    break","metadata":{"_uuid":"6ee78d3d-acfc-4743-b458-5f2af805ed37","_cell_guid":"65217966-a8a4-4571-ae4c-6b7b04936773","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T06:00:33.68659Z","iopub.execute_input":"2022-07-14T06:00:33.68686Z","iopub.status.idle":"2022-07-14T06:07:42.867678Z","shell.execute_reply.started":"2022-07-14T06:00:33.686829Z","shell.execute_reply":"2022-07-14T06:07:42.865671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = pd.DataFrame(metrics)\n\ngo.Figure(\n    data=(\n        go.Scatter(x=metrics.index, y=metrics[\"accuracy\"], name=\"train\"),\n        go.Scatter(x=metrics.index, y=metrics[\"val_accuracy\"], name=\"validation\"),\n    ),\n    layout=dict(\n        width=600,\n        title_text=\"Model metrics\",\n        xaxis_title_text=\"Epoch\",\n        font=dict(size=16),\n    ),\n)","metadata":{"_uuid":"f3b29c83-1560-44aa-a227-ce14b9b58453","_cell_guid":"952860a6-9fe6-4b2d-af81-b414a99083d1","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-07-14T05:52:26.713571Z","iopub.execute_input":"2022-07-14T05:52:26.714359Z","iopub.status.idle":"2022-07-14T05:52:26.736564Z","shell.execute_reply.started":"2022-07-14T05:52:26.714304Z","shell.execute_reply":"2022-07-14T05:52:26.735595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}