{"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":"code","source":"#LIBRARIES IMPORT\nimport os\nimport re\nimport jax\nimport flax\nimport optax\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport jax.numpy as jnp\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nimport tqdm\nfrom tqdm.notebook import trange\nfrom sklearn.metrics import f1_score\n\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-18T16:18:20.168576Z","iopub.execute_input":"2022-12-18T16:18:20.16921Z","iopub.status.idle":"2022-12-18T16:18:33.848346Z","shell.execute_reply.started":"2022-12-18T16:18:20.169106Z","shell.execute_reply":"2022-12-18T16:18:33.84728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For occupying required amount of GPU memory\nos.environ[\"XLA_PYTHON_CLIENT_ALLOCATOR\"] = 'platform'","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:33.85051Z","iopub.execute_input":"2022-12-18T16:18:33.851478Z","iopub.status.idle":"2022-12-18T16:18:33.856341Z","shell.execute_reply.started":"2022-12-18T16:18:33.851437Z","shell.execute_reply":"2022-12-18T16:18:33.855005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/google/jax/issues/10989#issuecomment-1147710198\ntf.config.experimental.set_visible_devices([], 'GPU')","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:33.857913Z","iopub.execute_input":"2022-12-18T16:18:33.858725Z","iopub.status.idle":"2022-12-18T16:18:34.33802Z","shell.execute_reply.started":"2022-12-18T16:18:33.858684Z","shell.execute_reply":"2022-12-18T16:18:34.337018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ==============================\n# Resource info\n# ==============================\nprint(\"local devices: \", jax.local_devices())\nprint(\"local device count: \", jax.local_device_count())","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:34.340889Z","iopub.execute_input":"2022-12-18T16:18:34.341352Z","iopub.status.idle":"2022-12-18T16:18:34.654396Z","shell.execute_reply.started":"2022-12-18T16:18:34.341285Z","shell.execute_reply":"2022-12-18T16:18:34.653382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs","metadata":{}},{"cell_type":"code","source":"DEVICE_COUNT = jax.local_device_count()\nAUTO = tf.data.experimental.AUTOTUNE\nIMAGE_SIZE = [192, 192]\nSEED = 42\nNUM_CLASSES = 102\nBATCH_SIZE = 64 * DEVICE_COUNT\nEPOCHS = 100\n\ndtype = jnp.float16  # half precision\n\nTRAIN_FILES = tf.io.gfile.glob(\n    '../input/flower-classification-with-tpus/tfrecords-jpeg-192x192/train/*.tfrec'\n)\nVAL_FILES = tf.io.gfile.glob(\n    '../input/flower-classification-with-tpus/tfrecords-jpeg-192x192/val/*.tfrec'\n)\nTEST_FILES = tf.io.gfile.glob(\n    '../input/flower-classification-with-tpus/tfrecords-jpeg-192x192/test/*.tfrec'\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:34.655508Z","iopub.execute_input":"2022-12-18T16:18:34.656055Z","iopub.status.idle":"2022-12-18T16:18:34.68897Z","shell.execute_reply.started":"2022-12-18T16:18:34.656018Z","shell.execute_reply":"2022-12-18T16:18:34.688111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# function for decoding and reshaping image\ndef decode_image(image_data):\n    \"\"\"\n    Decodes jpeg image and reshapes \n    it to (IMAGE_HEIGHT, IMAGE_WIDTH, 3)\n    \"\"\"\n    image = tf.image.decode_jpeg(image_data, channels=3)\n    image = tf.cast(image, tf.float32)\n    image = tf.reshape(image, [*IMAGE_SIZE, 3])\n    return image\n\n\n# function for reading labeled tf_record\n# useful for loading train and valid data\ndef read_labeled_tfrecord(example):\n    \"\"\"\n    Reads labeled tf_record for train/valid\n    datasets\n    \"\"\"\n    LABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"class\": tf.io.FixedLenFeature([], tf.int64),\n    }\n    example = tf.io.parse_single_example(example, LABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    label = tf.cast(example['class'], tf.int32)\n    return {\"image\": image, \"label\": label}\n\n\n# function for reading un_labeled tf_record\n# useful for loading test data\ndef read_unlabeled_tfrecord(example):\n    \"\"\"\n    Reads unlabeled tf_record for \n    test dataset\n    \"\"\"\n    UNLABELED_TFREC_FORMAT = {\n        \"image\": tf.io.FixedLenFeature([], tf.string),\n        \"id\": tf.io.FixedLenFeature([], tf.string),\n    }\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'])\n    idnum = example['id']\n    return {\"image\": image, \"image_id\": idnum}\n\n\n# function to load datasets\ndef load_dataset(filenames, labeled=True):\n    \"\"\"\n    Loads tf.data.Dataset \n    \"\"\"\n    ignore_order = tf.data.Options()\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO)\n    dataset = dataset.with_options(ignore_order)\n    dataset = dataset.map(\n        read_labeled_tfrecord if labeled else read_unlabeled_tfrecord,\n        num_parallel_calls=AUTO)\n    return dataset\n\n\n# Based on https://github.com/google-research/vision_transformer/blob/main/vit_jax/input_pipeline.py\ndef get_data_inference(*,\n                       data,\n                       mode,\n                       num_classes,\n                       batch_size,\n                       image_size,\n                       preprocess=None):\n    \"\"\"Returns dataset for training/eval.\n    Args:\n    data: tf.data.Dataset to read data from.\n    mode: Must be \"train\" or \"test\".\n    num_classes: Number of classes (used for one-hot encoding).\n    batch_size: Global batch size. Note that the returned dataset will have\n      dimensions [local_devices, batch_size / local_devices, ...].\n    image_size: Image size list [height, width] after cropping (for training) / resizing (for\n      evaluation).\n    preprocess: Optional preprocess function. This function will be applied to\n      the dataset just after repeat/shuffling, and before the data augmentation\n      preprocess step is applied.\n    \"\"\"\n\n    def _pp(data):\n        im = data['image']\n        im = tf.image.resize(im, image_size)\n        im = (im - 127.5) / 127.5\n        data['image'] = im\n        return data\n\n    if preprocess is not None:\n        data = data.map(preprocess, tf.data.experimental.AUTOTUNE)\n    data = data.map(_pp, tf.data.experimental.AUTOTUNE)\n    data = data.batch(batch_size, drop_remainder=False)\n\n    # Shard data such that it can be distributed accross devices\n    num_devices = jax.local_device_count()\n\n    def _shard(data):\n        data['image'] = tf.reshape(\n            data['image'],\n            [num_devices, -1, *image_size, data['image'].shape[-1]])\n        data['image_id'] = tf.reshape(data['image_id'], [num_devices, -1])\n        return data\n\n    if num_devices is not None:\n        data = data.map(_shard, tf.data.experimental.AUTOTUNE)\n\n    return data.prefetch(1)\n\n\n# function for counting total items\ndef count_data_items(filenames):\n    \"\"\"\n    Counts the number of \n    data items\n    \"\"\"\n    n = [\n        int(re.compile(r\"-([0-9]*)\\.\").search(filename).group(1))\n        for filename in filenames\n    ]\n    return np.sum(n)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:34.690245Z","iopub.execute_input":"2022-12-18T16:18:34.690573Z","iopub.status.idle":"2022-12-18T16:18:34.708757Z","shell.execute_reply.started":"2022-12-18T16:18:34.69054Z","shell.execute_reply":"2022-12-18T16:18:34.707859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_TEST_IMAGES = count_data_items(TEST_FILES)\n\ntest_dataset = get_data_inference(data=load_dataset(TEST_FILES, labeled=False),\n                         mode=\"train\",\n                         num_classes=NUM_CLASSES,\n                         batch_size=BATCH_SIZE,\n                         image_size=IMAGE_SIZE,\n                         preprocess=None)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:34.709889Z","iopub.execute_input":"2022-12-18T16:18:34.710627Z","iopub.status.idle":"2022-12-18T16:18:35.003216Z","shell.execute_reply.started":"2022-12-18T16:18:34.710451Z","shell.execute_reply":"2022-12-18T16:18:35.002289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show images","metadata":{}},{"cell_type":"code","source":"batch = next(iter(test_dataset.as_numpy_iterator()))\nplt.figure(figsize=(12, 12))\nfor i, img in enumerate(batch[\"image\"][0]):\n    if i == 16:\n        break\n    plt.subplot(4, 4, i + 1)\n    plt.imshow((img * 127.5 + 127.5).astype(int))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:18:35.004763Z","iopub.execute_input":"2022-12-18T16:18:35.00514Z","iopub.status.idle":"2022-12-18T16:18:37.138325Z","shell.execute_reply.started":"2022-12-18T16:18:35.005106Z","shell.execute_reply":"2022-12-18T16:18:37.136979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Implementation of ResNet","metadata":{}},{"cell_type":"code","source":"# From https://github.com/google/flax/blob/main/examples/imagenet/models.py\n\n# Copyright 2022 The Flax Authors.\n#\n# Licensed under the Apache License, Version 2.0 (the \"License\");\n# you may not use this file except in compliance with the License.\n# You may obtain a copy of the License at\n#\n#     http://www.apache.org/licenses/LICENSE-2.0\n#\n# Unless required by applicable law or agreed to in writing, software\n# distributed under the License is distributed on an \"AS IS\" BASIS,\n# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.\n# See the License for the specific language governing permissions and\n# limitations under the License.\n\"\"\"Flax implementation of ResNet V1.\"\"\"\n\n# See issue #620.\n# pytype: disable=wrong-arg-count\n\nfrom functools import partial\nfrom typing import Any, Callable, Sequence, Tuple\n\nfrom flax import linen as nn\nimport jax.numpy as jnp\n\nModuleDef = Any\n\n\nclass ResNetBlock(nn.Module):\n    \"\"\"ResNet block.\"\"\"\n    filters: int\n    conv: ModuleDef\n    norm: ModuleDef\n    act: Callable\n    strides: Tuple[int, int] = (1, 1)\n\n    @nn.compact\n    def __call__(\n        self,\n        x,\n    ):\n        residual = x\n        y = self.conv(self.filters, (3, 3), self.strides)(x)\n        y = self.norm()(y)\n        y = self.act(y)\n        y = self.conv(self.filters, (3, 3))(y)\n        y = self.norm(scale_init=nn.initializers.zeros)(y)\n\n        if residual.shape != y.shape:\n            residual = self.conv(self.filters, (1, 1),\n                                 self.strides,\n                                 name='conv_proj')(residual)\n            residual = self.norm(name='norm_proj')(residual)\n\n        return self.act(residual + y)\n\n\nclass BottleneckResNetBlock(nn.Module):\n    \"\"\"Bottleneck ResNet block.\"\"\"\n    filters: int\n    conv: ModuleDef\n    norm: ModuleDef\n    act: Callable\n    strides: Tuple[int, int] = (1, 1)\n\n    @nn.compact\n    def __call__(self, x):\n        residual = x\n        y = self.conv(self.filters, (1, 1))(x)\n        y = self.norm()(y)\n        y = self.act(y)\n        y = self.conv(self.filters, (3, 3), self.strides)(y)\n        y = self.norm()(y)\n        y = self.act(y)\n        y = self.conv(self.filters * 4, (1, 1))(y)\n        y = self.norm(scale_init=nn.initializers.zeros)(y)\n\n        if residual.shape != y.shape:\n            residual = self.conv(self.filters * 4, (1, 1),\n                                 self.strides,\n                                 name='conv_proj')(residual)\n            residual = self.norm(name='norm_proj')(residual)\n\n        return self.act(residual + y)\n\n\nclass ResNet(nn.Module):\n    \"\"\"ResNetV1.\"\"\"\n    stage_sizes: Sequence[int]\n    block_cls: ModuleDef\n    num_classes: int\n    num_filters: int = 64\n    dtype: Any = jnp.float32\n    act: Callable = nn.relu\n    conv: ModuleDef = nn.Conv\n\n    @nn.compact\n    def __call__(self, x, train: bool = True):\n        conv = partial(self.conv, use_bias=False, dtype=self.dtype)\n        norm = partial(nn.BatchNorm,\n                       use_running_average=not train,\n                       momentum=0.9,\n                       epsilon=1e-5,\n                       dtype=self.dtype)\n\n        x = conv(self.num_filters, (7, 7), (2, 2),\n                 padding=[(3, 3), (3, 3)],\n                 name='conv_init')(x)\n        x = norm(name='bn_init')(x)\n        x = nn.relu(x)\n        x = nn.max_pool(x, (3, 3), strides=(2, 2), padding='SAME')\n        for i, block_size in enumerate(self.stage_sizes):\n            for j in range(block_size):\n                strides = (2, 2) if i > 0 and j == 0 else (1, 1)\n                x = self.block_cls(self.num_filters * 2**i,\n                                   strides=strides,\n                                   conv=conv,\n                                   norm=norm,\n                                   act=self.act)(x)\n        x = jnp.mean(x, axis=(1, 2))\n        x = nn.Dense(self.num_classes, dtype=self.dtype)(x)\n        x = jnp.asarray(x, self.dtype)\n        return x\n\n\nResNet18 = partial(ResNet, stage_sizes=[2, 2, 2, 2], block_cls=ResNetBlock)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:19:06.391543Z","iopub.execute_input":"2022-12-18T16:19:06.391927Z","iopub.status.idle":"2022-12-18T16:19:06.415171Z","shell.execute_reply.started":"2022-12-18T16:19:06.391894Z","shell.execute_reply":"2022-12-18T16:19:06.413911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions for inference","metadata":{}},{"cell_type":"code","source":"def inference(params_repl, steps, dataset):\n    pred_list = []\n    image_id_list = []\n    for i, batch in zip(tqdm.trange(steps), dataset.as_numpy_iterator()):\n        logits = model_apply_repl(params_repl, batch[\"image\"])\n\n        image_id_list.extend(batch[\"image_id\"])\n        pred_list.extend(\n            np.argmax(jax.device_get(logits), -1).flatten().tolist())\n    return pred_list, sum([i.tolist() for i in image_id_list], [])","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:23:10.656029Z","iopub.execute_input":"2022-12-18T16:23:10.656509Z","iopub.status.idle":"2022-12-18T16:23:10.673663Z","shell.execute_reply.started":"2022-12-18T16:23:10.656465Z","shell.execute_reply":"2022-12-18T16:23:10.672693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Init models","metadata":{}},{"cell_type":"code","source":"# ==========================\n# model\n# ==========================\nmodel = ResNet18(num_classes=NUM_CLASSES, dtype=dtype)\n\nparams = jax.jit(lambda: model.init(\n    jax.random.PRNGKey(0), batch[\"image\"][0, :1], train=False),\n                 backend='cpu')()","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:19:08.798159Z","iopub.execute_input":"2022-12-18T16:19:08.798585Z","iopub.status.idle":"2022-12-18T16:19:17.01369Z","shell.execute_reply.started":"2022-12-18T16:19:08.798545Z","shell.execute_reply":"2022-12-18T16:19:17.012578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make replication","metadata":{}},{"cell_type":"code","source":"params_repl = flax.jax_utils.replicate(params)\nmodel_apply_repl = jax.pmap(\n    lambda params, inputs: model.apply(params, inputs, train=False),\n    axis_name='batch')","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:19:18.631326Z","iopub.execute_input":"2022-12-18T16:19:18.632511Z","iopub.status.idle":"2022-12-18T16:19:18.724157Z","shell.execute_reply.started":"2022-12-18T16:19:18.632455Z","shell.execute_reply":"2022-12-18T16:19:18.723178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load weight","metadata":{}},{"cell_type":"code","source":"from flax.serialization import to_bytes, from_bytes\nfrom flax.linen import FrozenDict\n\ndef load_params(params: FrozenDict, path: str) -> FrozenDict:\n    with open(path, 'rb') as f:\n        serialized_params = f.read()\n\n    return FrozenDict(from_bytes(params, serialized_params))","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:19:20.792714Z","iopub.execute_input":"2022-12-18T16:19:20.793109Z","iopub.status.idle":"2022-12-18T16:19:20.798779Z","shell.execute_reply.started":"2022-12-18T16:19:20.793058Z","shell.execute_reply":"2022-12-18T16:19:20.797593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"params_repl = load_params(\n    params_repl,\n    \"/kaggle/input/flax-training-tutorial-by-inoichan/resnet.weight\")\n\n# re-replicate for the device used now\nparams = flax.jax_utils.unreplicate(params_repl)\nparams_repl = flax.jax_utils.replicate(params)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:19:23.206013Z","iopub.execute_input":"2022-12-18T16:19:23.206972Z","iopub.status.idle":"2022-12-18T16:19:25.029868Z","shell.execute_reply.started":"2022-12-18T16:19:23.206917Z","shell.execute_reply":"2022-12-18T16:19:25.028872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"pred_list, image_id_list = inference(params_repl=params_repl,\n                                     steps=(NUM_TEST_IMAGES // BATCH_SIZE) + 1,\n                                     dataset=test_dataset)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:23:15.195028Z","iopub.execute_input":"2022-12-18T16:23:15.195386Z","iopub.status.idle":"2022-12-18T16:23:25.462752Z","shell.execute_reply.started":"2022-12-18T16:23:15.195357Z","shell.execute_reply":"2022-12-18T16:23:25.461766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv(\n    '/kaggle/input/flower-classification-with-tpus/sample_submission.csv')\nsub[\"id\"] = list(map(bytes.decode, image_id_list))\nsub[\"label\"] = pred_list\nsub.to_csv('submission.csv', index=False)\nsub","metadata":{"execution":{"iopub.status.busy":"2022-12-18T16:27:02.727026Z","iopub.execute_input":"2022-12-18T16:27:02.727387Z","iopub.status.idle":"2022-12-18T16:27:02.761806Z","shell.execute_reply.started":"2022-12-18T16:27:02.727357Z","shell.execute_reply":"2022-12-18T16:27:02.760805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}