{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.8.17","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# About\n\nA simple demonstration how to use *high-level to low-level* `keras` API while still using `model.fit`; on **TPU-VM**.","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd \nimport os\nimport warnings\nwarnings.simplefilter(action=\"ignore\")\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"3\"\n\nimport numpy as np\nfrom tqdm import tqdm\nfrom functools import partial\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits import axes_grid1\nfrom IPython.display import clear_output\nfrom sklearn.model_selection import train_test_split \n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import losses\nfrom tensorflow.keras import metrics \nfrom tensorflow.keras import optimizers\nclear_output()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-03T12:11:56.653814Z","iopub.execute_input":"2023-09-03T12:11:56.654241Z","iopub.status.idle":"2023-09-03T12:12:37.864663Z","shell.execute_reply.started":"2023-09-03T12:11:56.654200Z","shell.execute_reply":"2023-09-03T12:12:37.863681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = 'TPU' #  option: ['GPU', 'TPU']","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:37.866475Z","iopub.execute_input":"2023-09-03T12:12:37.867043Z","iopub.status.idle":"2023-09-03T12:12:37.871314Z","shell.execute_reply.started":"2023-09-03T12:12:37.867002Z","shell.execute_reply":"2023-09-03T12:12:37.870420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEVICE == 'GPU':\n    physical_devices = tf.config.list_physical_devices('GPU')\n    tf.config.optimizer.set_jit(True)\n    keras.mixed_precision.set_global_policy(\"mixed_float16\")\n    [\n        tf.config.experimental.set_memory_growth(pd, True) \\\n        for pd in physical_devices\n    ]\n    strategy = tf.distribute.MirroredStrategy()\n    \nelif DEVICE == \"TPU\":\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(\n        tpu=\"local\"\n    )\n    strategy = tf.distribute.TPUStrategy(tpu)\n    keras.mixed_precision.set_global_policy(\"mixed_bfloat16\")\n    \nclear_output()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:37.872470Z","iopub.execute_input":"2023-09-03T12:12:37.872785Z","iopub.status.idle":"2023-09-03T12:12:46.495747Z","shell.execute_reply.started":"2023-09-03T12:12:37.872757Z","shell.execute_reply":"2023-09-03T12:12:46.494837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"strategy","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:46.498021Z","iopub.execute_input":"2023-09-03T12:12:46.498363Z","iopub.status.idle":"2023-09-03T12:12:46.507385Z","shell.execute_reply.started":"2023-09-03T12:12:46.498333Z","shell.execute_reply":"2023-09-03T12:12:46.506559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**utility**","metadata":{}},{"cell_type":"code","source":"def plot_history(history):\n    # Extract values\n    gra_acc = history['gra_accuracy']\n    vow_acc = history['vow_accuracy']\n    cons_acc = history['cons_accuracy']\n\n    val_gra_acc = history['val_gra_accuracy']\n    val_vow_acc = history['val_vow_accuracy']\n    val_cons_acc = history['val_cons_accuracy']\n\n    gra_loss = history['gra_loss']\n    vow_loss = history['vow_loss']\n    cons_loss = history['cons_loss']\n\n    val_gra_loss = history['val_gra_loss']\n    val_vow_loss = history['val_vow_loss']\n    val_cons_loss = history['val_cons_loss']\n\n    epochs = range(1, len(gra_acc) + 1)\n\n    # Accuracy plot\n    plt.figure(figsize=(14, 6))\n    plt.subplot(1, 2, 1)\n    plt.plot(epochs, gra_acc, label='Grapheme Training Accuracy', marker='o')\n    plt.plot(epochs, vow_acc, label='Vowel Training Accuracy', marker='o')\n    plt.plot(epochs, cons_acc, label='Consonant Training Accuracy', marker='o')\n    plt.plot(epochs, val_gra_acc, label='Grapheme Validation Accuracy', linestyle='--')\n    plt.plot(epochs, val_vow_acc, label='Vowel Validation Accuracy', linestyle='--')\n    plt.plot(epochs, val_cons_acc, label='Consonant Validation Accuracy', linestyle='--')\n    plt.legend()\n    plt.title('Training and Validation Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n\n    # Loss plot\n    plt.subplot(1, 2, 2)\n    plt.plot(epochs, gra_loss, label='Grapheme Training Loss', marker='o')\n    plt.plot(epochs, vow_loss, label='Vowel Training Loss', marker='o')\n    plt.plot(epochs, cons_loss, label='Consonant Training Loss', marker='o')\n    plt.plot(epochs, val_gra_loss, label='Grapheme Validation Loss', linestyle='--')\n    plt.plot(epochs, val_vow_loss, label='Vowel Validation Loss', linestyle='--')\n    plt.plot(epochs, val_cons_loss, label='Consonant Validation Loss', linestyle='--')\n    plt.legend()\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:46.508535Z","iopub.execute_input":"2023-09-03T12:12:46.508837Z","iopub.status.idle":"2023-09-03T12:12:46.524020Z","shell.execute_reply.started":"2023-09-03T12:12:46.508810Z","shell.execute_reply":"2023-09-03T12:12:46.523203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"im_path = '../input/grapheme-imgs-128x128/'\ndf  = pd.read_csv('../input/bengaliai-cv19/train.csv')\ndf = df.sample(frac=1).reset_index(drop=True)\ndf['filename'] = df.image_id.apply(\n    lambda filename: im_path + filename + '.png'\n)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:46.524985Z","iopub.execute_input":"2023-09-03T12:12:46.525515Z","iopub.status.idle":"2023-09-03T12:12:46.975832Z","shell.execute_reply.started":"2023-09-03T12:12:46.525484Z","shell.execute_reply":"2023-09-03T12:12:46.974928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, valid_df = train_test_split(\n    df,\n    test_size=0.20,\n    random_state=42,\n    stratify=df[['grapheme_root', 'vowel_diacritic', 'consonant_diacritic']]\n)\ntrain_df.shape, valid_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:46.976902Z","iopub.execute_input":"2023-09-03T12:12:46.977184Z","iopub.status.idle":"2023-09-03T12:12:48.638786Z","shell.execute_reply.started":"2023-09-03T12:12:46.977159Z","shell.execute_reply":"2023-09-03T12:12:48.637710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"AUTOTUNE=tf.data.AUTOTUNE\nREPLICAS=strategy.num_replicas_in_sync\nBATCH_SIZE=128 * REPLICAS\nIMAGE_SIZE=224\n\ndef read_data(path, g, v, c):\n    # read image (x)\n    file_bytes = tf.io.read_file(path)\n    img = tf.image.decode_jpeg(file_bytes, channels = 3)\n    img = tf.image.resize(img, [IMAGE_SIZE, IMAGE_SIZE])\n    \n    # target labels\n    g = tf.one_hot(g, depth=168, dtype='float32')\n    v = tf.one_hot(v, depth=11, dtype='float32')\n    c = tf.one_hot(c, depth=7, dtype='float32')\n    \n    return img, (g, v, c)\n\ndef augment(img):\n    img = tf.image.random_flip_left_right(img)\n    img = tf.image.random_flip_up_down(img)\n    img = tf.image.random_saturation(img, 0.65, 1.05)\n    img = tf.image.random_brightness(img, 0.05)\n    img = tf.image.random_contrast(img, 0.75, 1.05)\n    img = tf.image.random_hue(img, 0.05)\n    return img","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:48.640041Z","iopub.execute_input":"2023-09-03T12:12:48.640373Z","iopub.status.idle":"2023-09-03T12:12:48.650223Z","shell.execute_reply.started":"2023-09-03T12:12:48.640347Z","shell.execute_reply":"2023-09-03T12:12:48.649288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = tf.data.Dataset.from_tensor_slices(\n    (\n        train_df['filename'].values, \n        train_df['grapheme_root'].values,\n        train_df['vowel_diacritic'].values,\n        train_df['consonant_diacritic'].values,\n    )\n)\ntrain_ds = train_ds.map(read_data, num_parallel_calls=AUTOTUNE)\ntrain_ds = train_ds.map(\n    lambda x, y: (augment(x), y), num_parallel_calls=AUTOTUNE\n)\ntrain_ds = train_ds.shuffle(BATCH_SIZE * 8)\ntrain_ds = train_ds.batch(BATCH_SIZE, drop_remainder=True)\ntrain_ds = train_ds.prefetch(AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:48.651232Z","iopub.execute_input":"2023-09-03T12:12:48.651497Z","iopub.status.idle":"2023-09-03T12:12:48.982668Z","shell.execute_reply.started":"2023-09-03T12:12:48.651473Z","shell.execute_reply":"2023-09-03T12:12:48.981574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_ds = tf.data.Dataset.from_tensor_slices(\n    (\n        valid_df['filename'].values, \n        valid_df['grapheme_root'].values,\n        valid_df['vowel_diacritic'].values,\n        valid_df['consonant_diacritic'].values,\n    )\n)\nvalid_ds = valid_ds.map(read_data, num_parallel_calls=AUTOTUNE)\nvalid_ds = valid_ds.batch(BATCH_SIZE, drop_remainder=True)\nvalid_ds = valid_ds.prefetch(AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:48.986716Z","iopub.execute_input":"2023-09-03T12:12:48.987031Z","iopub.status.idle":"2023-09-03T12:12:49.027303Z","shell.execute_reply.started":"2023-09-03T12:12:48.987002Z","shell.execute_reply":"2023-09-03T12:12:49.026394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = next(iter(train_ds))\nv, w = next(iter(valid_ds))\nx.shape, v.shape","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:49.028446Z","iopub.execute_input":"2023-09-03T12:12:49.028850Z","iopub.status.idle":"2023-09-03T12:12:54.301823Z","shell.execute_reply.started":"2023-09-03T12:12:49.028820Z","shell.execute_reply":"2023-09-03T12:12:54.300680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:54.303069Z","iopub.execute_input":"2023-09-03T12:12:54.303484Z","iopub.status.idle":"2023-09-03T12:12:54.312052Z","shell.execute_reply.started":"2023-09-03T12:12:54.303454Z","shell.execute_reply":"2023-09-03T12:12:54.311050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"w","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:54.313223Z","iopub.execute_input":"2023-09-03T12:12:54.313560Z","iopub.status.idle":"2023-09-03T12:12:54.328398Z","shell.execute_reply.started":"2023-09-03T12:12:54.313514Z","shell.execute_reply":"2023-09-03T12:12:54.327466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_cols = 3\nnum_rows = 2\n\nfig = plt.figure(figsize=(8, 8))\ngrid = axes_grid1.ImageGrid(fig, 111, nrows_ncols=(num_rows, num_cols), axes_pad=0.1)\nfor ax, im in zip(grid, x):\n    ax.imshow(tf.cast(im[..., 0], dtype=tf.uint8))\n    ax.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:54.329642Z","iopub.execute_input":"2023-09-03T12:12:54.329970Z","iopub.status.idle":"2023-09-03T12:12:55.072729Z","shell.execute_reply.started":"2023-09-03T12:12:54.329940Z","shell.execute_reply":"2023-09-03T12:12:55.071810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(8, 8))\ngrid = axes_grid1.ImageGrid(fig, 111, nrows_ncols=(num_rows, num_cols), axes_pad=0.1)\nfor ax, im in zip(grid, v):\n    ax.imshow(tf.cast(im[...,0], dtype=tf.uint8))\n    ax.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:55.073882Z","iopub.execute_input":"2023-09-03T12:12:55.074181Z","iopub.status.idle":"2023-09-03T12:12:55.685275Z","shell.execute_reply.started":"2023-09-03T12:12:55.074156Z","shell.execute_reply":"2023-09-03T12:12:55.684287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# High Level","metadata":{}},{"cell_type":"code","source":"def get_model(with_compile=True):\n    input = keras.Input(shape=(IMAGE_SIZE, IMAGE_SIZE, 3))\n    backbone = keras.applications.EfficientNetV2B0(\n        include_top=False,\n        weights='imagenet',\n        input_tensor=input,\n        include_preprocessing=True\n    )\n    gap = layers.GlobalAveragePooling2D()(backbone.output)\n    g_tensor = layers.Dense(\n        168, activation=None, name='gra',  dtype='float32'\n    )(gap)\n    v_tensor = layers.Dense(\n        11,  activation=None, name='vow',  dtype='float32'\n    )(gap)\n    c_tensor = layers.Dense(\n        7,   activation=None, name='cons', dtype='float32'\n    )(gap)\n    model = keras.Model(input, [g_tensor, v_tensor, c_tensor])\n    \n    if with_compile:\n        model.compile(\n            optimizer = optimizers.AdamW(\n                learning_rate=5e-5, weight_decay=0.01\n            ), \n            loss = {\n                'gra' : losses.CategoricalCrossentropy(from_logits=True), \n                'vow' : losses.CategoricalCrossentropy(from_logits=True), \n                'cons': losses.CategoricalCrossentropy(from_logits=True),\n            },\n            loss_weights = {\n                'gra' : 1.0,\n                'vow' : 1.0,\n                'cons': 1.0\n            },\n            metrics={\n                'gra' : metrics.CategoricalAccuracy('accuracy'), \n                'vow' : metrics.CategoricalAccuracy('accuracy'),\n                'cons': metrics.CategoricalAccuracy('accuracy')\n            }\n        )\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:55.686483Z","iopub.execute_input":"2023-09-03T12:12:55.686790Z","iopub.status.idle":"2023-09-03T12:12:55.698083Z","shell.execute_reply.started":"2023-09-03T12:12:55.686763Z","shell.execute_reply":"2023-09-03T12:12:55.697303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    model = get_model()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:12:55.699259Z","iopub.execute_input":"2023-09-03T12:12:55.699573Z","iopub.status.idle":"2023-09-03T12:13:22.634568Z","shell.execute_reply.started":"2023-09-03T12:12:55.699544Z","shell.execute_reply":"2023-09-03T12:13:22.633026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = model.fit(\n    train_ds, \n    validation_data=valid_ds, \n    epochs=10\n).history","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:13:22.636175Z","iopub.execute_input":"2023-09-03T12:13:22.636506Z","iopub.status.idle":"2023-09-03T12:23:42.713612Z","shell.execute_reply.started":"2023-09-03T12:13:22.636477Z","shell.execute_reply":"2023-09-03T12:23:42.712408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:23:42.716399Z","iopub.execute_input":"2023-09-03T12:23:42.716759Z","iopub.status.idle":"2023-09-03T12:23:43.360204Z","shell.execute_reply.started":"2023-09-03T12:23:42.716727Z","shell.execute_reply":"2023-09-03T12:23:43.359178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mid Level A","metadata":{}},{"cell_type":"code","source":"class CustomModel(keras.Model):\n    def __init__(self, model, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.model = model \n        self.output_names = model.output_names\n        self.input_names = model.input_names\n        \n    def train_step(self, data):\n        x, y = data\n\n        with tf.GradientTape() as tape:\n            y_pred = self.model(x, training=True)\n            loss = self.compiled_loss(\n                y, y_pred, regularization_losses=self.losses\n            )\n\n        trainable_vars = self.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n        self.compiled_metrics.update_state(y, y_pred)\n        return {m.name: m.result() for m in self.metrics}\n    \n    def test_step(self, data):\n        x, y = data\n        y_pred = self.model(x, training=False)\n        loss = self.compiled_loss(\n            y, y_pred, regularization_losses=self.losses\n        )\n        self.compiled_metrics.update_state(y, y_pred)\n        return {m.name: m.result() for m in self.metrics}","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:23:43.361299Z","iopub.execute_input":"2023-09-03T12:23:43.361588Z","iopub.status.idle":"2023-09-03T12:23:43.372423Z","shell.execute_reply.started":"2023-09-03T12:23:43.361562Z","shell.execute_reply":"2023-09-03T12:23:43.371555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    model = get_model()\n    custom_model = CustomModel(model)\n    \n    custom_model.compile( \n        optimizer = optimizers.AdamW(\n            learning_rate=5e-5, weight_decay=0.01\n        ), \n        loss = {\n            'gra' : losses.CategoricalCrossentropy(from_logits=True), \n            'vow' : losses.CategoricalCrossentropy(from_logits=True), \n            'cons': losses.CategoricalCrossentropy(from_logits=True),\n        },\n        loss_weights = {\n            'gra' : 1.0,\n            'vow' : 1.0,\n            'cons': 1.0\n        },\n        metrics={\n            'gra' : metrics.CategoricalAccuracy('accuracy'), \n            'vow' : metrics.CategoricalAccuracy('accuracy'),\n            'cons': metrics.CategoricalAccuracy('accuracy')\n        }\n    )\n    ","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:23:43.373466Z","iopub.execute_input":"2023-09-03T12:23:43.373743Z","iopub.status.idle":"2023-09-03T12:23:55.666632Z","shell.execute_reply.started":"2023-09-03T12:23:43.373719Z","shell.execute_reply":"2023-09-03T12:23:55.665386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = custom_model.fit(\n    train_ds, \n    validation_data=valid_ds, \n    epochs=10\n).history","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:23:55.668035Z","iopub.execute_input":"2023-09-03T12:23:55.668365Z","iopub.status.idle":"2023-09-03T12:34:32.147043Z","shell.execute_reply.started":"2023-09-03T12:23:55.668335Z","shell.execute_reply":"2023-09-03T12:34:32.145486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:34:32.149709Z","iopub.execute_input":"2023-09-03T12:34:32.150038Z","iopub.status.idle":"2023-09-03T12:34:32.786458Z","shell.execute_reply.started":"2023-09-03T12:34:32.150010Z","shell.execute_reply":"2023-09-03T12:34:32.785195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Mid Level B","metadata":{}},{"cell_type":"code","source":"class CustomModelV2(keras.Model):\n    def __init__(self, model, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n        self.model = model \n        self.output_names = model.output_names\n        self.input_names = model.input_names\n\n        # accuracy instances\n        self.gra_metric = metrics.CategoricalAccuracy('gra_accuracy')\n        self.vow_metric = metrics.CategoricalAccuracy('vow_accuracy')\n        self.cons_metric = metrics.CategoricalAccuracy('cons_accuracy')\n        \n        # loss instances\n        self.gra_loss_fn = losses.CategoricalCrossentropy(\n            from_logits=True, reduction=losses.Reduction.NONE\n        )\n        self.vow_loss_fn = losses.CategoricalCrossentropy(\n            from_logits=True, reduction=losses.Reduction.NONE\n        )\n        self.cons_loss_fn = losses.CategoricalCrossentropy(\n            from_logits=True, reduction=losses.Reduction.NONE\n        )\n        \n        # mean object\n        self.gra_loss_tracker = metrics.Mean(name=\"gra_loss\")\n        self.vow_loss_tracker = metrics.Mean(name=\"vow_loss\")\n        self.cons_loss_tracker = metrics.Mean(name=\"cons_loss\")\n\n\n    def train_step(self, data):\n        x, y = data\n\n        with tf.GradientTape() as tape:\n            y_pred = self.model(x, training=True)\n            loss = self._compute_losses(y, y_pred)\n\n        # Compute gradients\n        trainable_vars = self.model.trainable_variables\n        gradients = tape.gradient(loss, trainable_vars)\n        \n        # Update weights\n        self.optimizer.apply_gradients(zip(gradients, trainable_vars))\n        \n        # Update metrics\n        self._compute_metrics(y, y_pred)\n        \n        return {m.name: m.result() for m in self.metrics}\n    \n    def test_step(self, data):\n        x, y = data\n        \n        # Compute predictions\n        y_pred = self.model(x, training=False)\n        \n        # Updates the metrics tracking the loss\n        loss = self._compute_losses(y, y_pred)\n        \n        # Update metrics\n        self._compute_metrics(y, y_pred)\n        \n        # Return a dict mapping metric names to current value.\n        # Note that it will include the loss (tracked in self.metrics).\n        return {m.name: m.result() for m in self.metrics}\n    \n    \n    def _compute_losses(self, y, y_pred):\n        gra_loss_value = self.gra_loss_fn(y[0], y_pred[0])\n        vow_loss_value = self.vow_loss_fn(y[1], y_pred[1])\n        cons_loss_value = self.cons_loss_fn(y[2], y_pred[2])\n        \n        self.gra_loss_tracker.update_state(gra_loss_value)\n        self.vow_loss_tracker.update_state(vow_loss_value)\n        self.cons_loss_tracker.update_state(cons_loss_value)\n        \n        return [gra_loss_value, vow_loss_value, cons_loss_value]\n    \n    def _compute_metrics(self, y, y_pred):\n        self.gra_metric.update_state(y[0], y_pred[0])\n        self.vow_metric.update_state(y[1], y_pred[1])\n        self.cons_metric.update_state(y[2], y_pred[2])\n    \n    @property\n    def metrics(self):\n        return super().metrics + self.train_metrics\n    \n    @property\n    def train_metrics(self):\n        return [\n            self.gra_metric,\n            self.vow_metric, \n            self.cons_metric, \n            self.gra_loss_tracker, \n            self.vow_loss_tracker, \n            self.cons_loss_tracker, \n        ]","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:34:32.787850Z","iopub.execute_input":"2023-09-03T12:34:32.788239Z","iopub.status.idle":"2023-09-03T12:34:32.810337Z","shell.execute_reply.started":"2023-09-03T12:34:32.788211Z","shell.execute_reply":"2023-09-03T12:34:32.809172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    model = get_model()\n    custom_model_v2 = CustomModelV2(model)\n    custom_model_v2.compile(\n        optimizer = keras.optimizers.AdamW(\n            learning_rate=5e-5, weight_decay=0.01\n        ), \n    )\n    ","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:34:32.811553Z","iopub.execute_input":"2023-09-03T12:34:32.811825Z","iopub.status.idle":"2023-09-03T12:34:45.270045Z","shell.execute_reply.started":"2023-09-03T12:34:32.811801Z","shell.execute_reply":"2023-09-03T12:34:45.268660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = custom_model_v2.fit(\n    train_ds, \n    validation_data=valid_ds, \n    epochs=10\n).history","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:34:45.271345Z","iopub.execute_input":"2023-09-03T12:34:45.271653Z","iopub.status.idle":"2023-09-03T12:45:37.520926Z","shell.execute_reply.started":"2023-09-03T12:34:45.271626Z","shell.execute_reply":"2023-09-03T12:45:37.519446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:37.523513Z","iopub.execute_input":"2023-09-03T12:45:37.523850Z","iopub.status.idle":"2023-09-03T12:45:38.190346Z","shell.execute_reply.started":"2023-09-03T12:45:37.523822Z","shell.execute_reply":"2023-09-03T12:45:38.189131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Distributed Custom Training - Low Level","metadata":{}},{"cell_type":"code","source":"train_ds_dist = strategy.experimental_distribute_dataset(train_ds)\nvalid_ds_dist = strategy.experimental_distribute_dataset(valid_ds)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:38.196003Z","iopub.execute_input":"2023-09-03T12:45:38.196293Z","iopub.status.idle":"2023-09-03T12:45:38.258212Z","shell.execute_reply.started":"2023-09-03T12:45:38.196266Z","shell.execute_reply":"2023-09-03T12:45:38.257130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    #  training loss tracker\n    gra_loss_tracker = metrics.Mean(name=\"gra_loss\")\n    vow_loss_tracker = metrics.Mean(name=\"vow_loss\")\n    cons_loss_tracker = metrics.Mean(name=\"cons_loss\")\n    \n    # validation loss tracker\n    val_loss_tracker = metrics.Mean(name=\"val_loss\")\n    val_gra_loss_tracker = metrics.Mean(name=\"val_gra_loss\")\n    val_vow_loss_tracker = metrics.Mean(name=\"val_vow_loss\")\n    val_cons_loss_tracker = metrics.Mean(name=\"val_cons_loss\")\n\n    # for training accuracy\n    gra_accuracy = metrics.CategoricalAccuracy(name='gra_accuracy')\n    vow_accuracy = metrics.CategoricalAccuracy(name='vow_accuracy')\n    cons_accuracy = metrics.CategoricalAccuracy(name='cons_accuracy')\n    \n    # for validation accuracy\n    val_gra_accuracy = metrics.CategoricalAccuracy(name='val_gra_accuracy')\n    val_vow_accuracy = metrics.CategoricalAccuracy(name='val_vow_accuracy')\n    val_cons_accuracy = metrics.CategoricalAccuracy(name='val_cons_accuracy')","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:38.259337Z","iopub.execute_input":"2023-09-03T12:45:38.259646Z","iopub.status.idle":"2023-09-03T12:45:38.506199Z","shell.execute_reply.started":"2023-09-03T12:45:38.259621Z","shell.execute_reply":"2023-09-03T12:45:38.505098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    loss_object = losses.CategoricalCrossentropy(\n        from_logits=True,\n        reduction=losses.Reduction.NONE\n    )\n    \n    def compute_loss(labels, predictions, model_losses):\n        loss_1 = loss_object(labels[0], predictions[0])\n        loss_2 = loss_object(labels[1], predictions[1])\n        loss_3 = loss_object(labels[2], predictions[2])\n        \n        gra_loss_tracker.update_state(loss_1)\n        vow_loss_tracker.update_state(loss_2)\n        cons_loss_tracker.update_state(loss_3)\n        \n        loss = loss_1 + loss_2 + loss_3\n        loss = tf.nn.compute_average_loss(\n            loss, global_batch_size=BATCH_SIZE\n        )\n        \n        if model_losses:\n            loss += tf.nn.scale_regularization_loss(\n                tf.add_n(model_losses)\n            )\n            \n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:38.507507Z","iopub.execute_input":"2023-09-03T12:45:38.507838Z","iopub.status.idle":"2023-09-03T12:45:38.516850Z","shell.execute_reply.started":"2023-09-03T12:45:38.507810Z","shell.execute_reply":"2023-09-03T12:45:38.515917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with strategy.scope():\n    model = get_model(with_compile=False)\n    optimizer = keras.optimizers.AdamW(\n        learning_rate=5e-5, weight_decay=0.01\n    )","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:38.517969Z","iopub.execute_input":"2023-09-03T12:45:38.518265Z","iopub.status.idle":"2023-09-03T12:45:50.940906Z","shell.execute_reply.started":"2023-09-03T12:45:38.518238Z","shell.execute_reply":"2023-09-03T12:45:50.939676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(inputs):\n    images, labels = inputs\n\n    with tf.GradientTape() as tape:\n        predictions = model(images, training=True)\n        loss = compute_loss(labels, predictions, model.losses)\n\n    gradients = tape.gradient(loss, model.trainable_variables)\n    optimizer.apply_gradients(zip(gradients, model.trainable_variables))\n\n    gra_accuracy.update_state(labels[0], predictions[0])\n    vow_accuracy.update_state(labels[1], predictions[1])\n    cons_accuracy.update_state(labels[2], predictions[2])\n    return loss","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:50.942409Z","iopub.execute_input":"2023-09-03T12:45:50.942773Z","iopub.status.idle":"2023-09-03T12:45:50.950314Z","shell.execute_reply.started":"2023-09-03T12:45:50.942743Z","shell.execute_reply":"2023-09-03T12:45:50.949310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_step(inputs):\n    images, labels = inputs\n\n    predictions = model(images, training=False)\n    loss_1 = loss_object(labels[0], predictions[0])\n    loss_2 = loss_object(labels[1], predictions[1])\n    loss_3 = loss_object(labels[2], predictions[2])\n    loss = loss_1 + loss_2 + loss_3\n\n    val_gra_loss_tracker.update_state(loss_1)\n    val_vow_loss_tracker.update_state(loss_2)\n    val_cons_loss_tracker.update_state(loss_3)\n    val_loss_tracker.update_state(loss)\n    \n    val_gra_accuracy.update_state(labels[0], predictions[0])\n    val_vow_accuracy.update_state(labels[1], predictions[1])\n    val_cons_accuracy.update_state(labels[2], predictions[2])","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:50.951153Z","iopub.execute_input":"2023-09-03T12:45:50.951451Z","iopub.status.idle":"2023-09-03T12:45:50.966752Z","shell.execute_reply.started":"2023-09-03T12:45:50.951423Z","shell.execute_reply":"2023-09-03T12:45:50.965774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# `run` replicates the provided computation and runs it\n# with the distributed input.\n@tf.function\ndef distributed_train_step(dataset_inputs):\n    per_replica_losses = strategy.run(\n        train_step, args=(dataset_inputs,)\n    )\n    return strategy.reduce(\n        tf.distribute.ReduceOp.SUM, per_replica_losses,\n        axis=None\n    )\n\n@tf.function\ndef distributed_test_step(dataset_inputs):\n    return strategy.run(\n        test_step, args=(dataset_inputs,)\n    )","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:50.967940Z","iopub.execute_input":"2023-09-03T12:45:50.968222Z","iopub.status.idle":"2023-09-03T12:45:50.976665Z","shell.execute_reply.started":"2023-09-03T12:45:50.968198Z","shell.execute_reply":"2023-09-03T12:45:50.975757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Initialize an empty DataFrame to store log results\ncolumns = [\n    'loss', 'gra_loss', 'vow_loss', 'cons_loss',\n    'gra_accuracy', 'vow_accuracy', 'cons_accuracy',\n    'val_loss', 'val_gra_loss', 'val_vow_loss', 'val_cons_loss',\n    'val_gra_accuracy', 'val_vow_accuracy', 'val_cons_accuracy'\n]\nhistory = pd.DataFrame(columns=columns)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:50.977638Z","iopub.execute_input":"2023-09-03T12:45:50.977893Z","iopub.status.idle":"2023-09-03T12:45:50.989025Z","shell.execute_reply.started":"2023-09-03T12:45:50.977870Z","shell.execute_reply":"2023-09-03T12:45:50.988081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(10):\n    # train set\n    total_loss = 0.0\n    num_batches = 0\n    \n    for x in tqdm(train_ds_dist, desc='training set'):\n        total_loss += distributed_train_step(x)\n        num_batches += 1\n        \n    train_loss = total_loss / num_batches\n\n    # val set\n    for x in tqdm(valid_ds_dist, desc='validation set'):\n        distributed_test_step(x)\n\n\n    template = (\n        f\"Epoch {epoch + 1}\\n\"\n        \"--------------------------\\n\"\n        f\"Train Metrics:\\n\"\n        f\"loss: {train_loss:.4f}, \"\n        f\"gra_loss: {gra_loss_tracker.result():.4f}, \"\n        f\"vow_loss: {vow_loss_tracker.result():.4f}, \"\n        f\"cons_loss: {cons_loss_tracker.result():.4f}\\n\"\n        f\"gra_accuracy: {gra_accuracy.result():.4f}, \"\n        f\"vow_accuracy: {vow_accuracy.result():.4f}, \"\n        f\"cons_accuracy: {cons_accuracy.result():.4f}\\n\"\n        \"--------------------------\\n\"\n        f\"Validation Metrics:\\n\"\n        f\"val_loss: {val_loss_tracker.result():.4f}, \"\n        f\"val_gra_loss: {val_gra_loss_tracker.result():.4f}, \"\n        f\"val_vow_loss: {val_vow_loss_tracker.result():.4f}, \"\n        f\"val_cons_loss: {val_cons_loss_tracker.result():.4f}\\n\"\n        f\"val_gra_accuracy: {val_gra_accuracy.result():.4f}, \"\n        f\"val_vow_accuracy: {val_vow_accuracy.result():.4f}, \"\n        f\"val_cons_accuracy: {val_cons_accuracy.result():.4f}\\n\"\n    )\n    print(template)\n\n    \n    history.loc[epoch] = [\n        train_loss.numpy(),\n        gra_loss_tracker.result().numpy(),\n        vow_loss_tracker.result().numpy(),\n        cons_loss_tracker.result().numpy(),\n        gra_accuracy.result().numpy(),\n        vow_accuracy.result().numpy(),\n        cons_accuracy.result().numpy(),\n        val_loss_tracker.result().numpy(),\n        val_gra_loss_tracker.result().numpy(),\n        val_vow_loss_tracker.result().numpy(),\n        val_cons_loss_tracker.result().numpy(),\n        val_gra_accuracy.result().numpy(),\n        val_vow_accuracy.result().numpy(),\n        val_cons_accuracy.result().numpy()\n    ]\n    \n    gra_loss_tracker.reset_states()\n    vow_loss_tracker.reset_states()\n    cons_loss_tracker.reset_states()\n    val_loss_tracker.reset_states()\n    val_gra_loss_tracker.reset_states()\n    val_vow_loss_tracker.reset_states()\n    val_cons_loss_tracker.reset_states()\n    gra_accuracy.reset_states()\n    vow_accuracy.reset_states()\n    cons_accuracy.reset_states()\n    val_gra_accuracy.reset_states()\n    val_vow_accuracy.reset_states()\n    val_cons_accuracy.reset_states()","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:45:50.990220Z","iopub.execute_input":"2023-09-03T12:45:50.990495Z","iopub.status.idle":"2023-09-03T12:56:54.551465Z","shell.execute_reply.started":"2023-09-03T12:45:50.990471Z","shell.execute_reply":"2023-09-03T12:56:54.550057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_history(history)","metadata":{"execution":{"iopub.status.busy":"2023-09-03T12:56:54.552872Z","iopub.execute_input":"2023-09-03T12:56:54.553157Z","iopub.status.idle":"2023-09-03T12:56:55.183330Z","shell.execute_reply.started":"2023-09-03T12:56:54.553130Z","shell.execute_reply":"2023-09-03T12:56:55.182036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}