{"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 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-18T10:20:18.602436Z","iopub.execute_input":"2022-12-18T10:20:18.602851Z","iopub.status.idle":"2022-12-18T10:20:25.565801Z","shell.execute_reply.started":"2022-12-18T10:20:18.602771Z","shell.execute_reply":"2022-12-18T10:20:25.564880Z"},"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-18T10:20:25.567821Z","iopub.execute_input":"2022-12-18T10:20:25.568472Z","iopub.status.idle":"2022-12-18T10:20:25.575382Z","shell.execute_reply.started":"2022-12-18T10:20:25.568435Z","shell.execute_reply":"2022-12-18T10:20:25.574357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Detect hardware, return appropriate distribution strategy\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver(\n    )  # TPU detection. No parameters necessary if TPU_NAME environment variable is set. On Kaggle this is always the case.\n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n    # https://github.com/google/jax/issues/10989#issuecomment-1147710198\n    tf.config.experimental.set_visible_devices([], 'GPU')","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:25.577040Z","iopub.execute_input":"2022-12-18T10:20:25.577652Z","iopub.status.idle":"2022-12-18T10:20:25.852786Z","shell.execute_reply.started":"2022-12-18T10:20:25.577613Z","shell.execute_reply":"2022-12-18T10:20:25.851606Z"},"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-18T10:20:25.855839Z","iopub.execute_input":"2022-12-18T10:20:25.856201Z","iopub.status.idle":"2022-12-18T10:20:26.150857Z","shell.execute_reply.started":"2022-12-18T10:20:25.856165Z","shell.execute_reply":"2022-12-18T10:20:26.149990Z"},"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.bfloat16 if tpu else 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)\n","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:26.152008Z","iopub.execute_input":"2022-12-18T10:20:26.152674Z","iopub.status.idle":"2022-12-18T10:20:26.178140Z","shell.execute_reply.started":"2022-12-18T10:20:26.152632Z","shell.execute_reply":"2022-12-18T10:20:26.177212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset func","metadata":{}},{"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# augmentation\ndef aug(data):\n    img = data[\"image\"]\n    img = tf.image.random_flip_left_right(img)\n    img = tf.image.random_hue(img, 0.1)\n    img = tf.image.random_saturation(img, 0.9, 1.1)\n    img = tf.image.random_contrast(img, 0.9, 1.1)\n    img = tf.image.random_brightness(img, 0.9)\n    img = tf.image.random_flip_left_right(img)\n    return dict(image=img, label=data[\"label\"])\n\n\n# Based on https://github.com/google-research/vision_transformer/blob/main/vit_jax/input_pipeline.py\ndef get_data(*,\n             data,\n             mode,\n             num_classes,\n             repeats,\n             batch_size,\n             image_size,\n             shuffle_buffer,\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    repeats: How many times the dataset should be repeated. For indefinite\n      repeats specify None.\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    shuffle_buffer: Number of elements to preload the shuffle buffer with.\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        label = tf.one_hot(data['label'], num_classes)  # pylint: disable=no-value-for-parameter\n        return {'image': im, 'label': label}\n\n    data = data.repeat(repeats)\n    if mode == 'train':\n        data = data.shuffle(shuffle_buffer)\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,\n                      drop_remainder=True if mode == \"train\" else 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['label'] = tf.reshape(data['label'],\n                                   [num_devices, -1, num_classes])\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-18T10:20:26.179414Z","iopub.execute_input":"2022-12-18T10:20:26.179786Z","iopub.status.idle":"2022-12-18T10:20:26.203370Z","shell.execute_reply.started":"2022-12-18T10:20:26.179749Z","shell.execute_reply":"2022-12-18T10:20:26.202238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Dataset","metadata":{}},{"cell_type":"code","source":"NUM_TRAIN_IMAGES = count_data_items(TRAIN_FILES)\nNUM_VAL_IMAGES = count_data_items(VAL_FILES)\n\ntrain_dataset = get_data(data=load_dataset(TRAIN_FILES, labeled=True),\n                         mode=\"train\",\n                         num_classes=NUM_CLASSES,\n                         repeats=None,\n                         batch_size=BATCH_SIZE,\n                         image_size=IMAGE_SIZE,\n                         shuffle_buffer=256,\n                         preprocess=aug)\n\nval_dataset = get_data(data=load_dataset(VAL_FILES, labeled=True),\n                       mode=\"test\",\n                       num_classes=NUM_CLASSES,\n                       repeats=None,\n                       batch_size=BATCH_SIZE,\n                       image_size=IMAGE_SIZE,\n                       shuffle_buffer=1,\n                       preprocess=None)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:26.204891Z","iopub.execute_input":"2022-12-18T10:20:26.205720Z","iopub.status.idle":"2022-12-18T10:20:26.631920Z","shell.execute_reply.started":"2022-12-18T10:20:26.205682Z","shell.execute_reply":"2022-12-18T10:20:26.631008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show images","metadata":{}},{"cell_type":"code","source":"batch = next(iter(train_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-18T10:20:26.633143Z","iopub.execute_input":"2022-12-18T10:20:26.633773Z","iopub.status.idle":"2022-12-18T10:20:29.472838Z","shell.execute_reply.started":"2022-12-18T10:20:26.633737Z","shell.execute_reply":"2022-12-18T10:20:29.467878Z"},"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-18T10:20:29.474794Z","iopub.execute_input":"2022-12-18T10:20:29.475286Z","iopub.status.idle":"2022-12-18T10:20:29.506609Z","shell.execute_reply.started":"2022-12-18T10:20:29.475250Z","shell.execute_reply":"2022-12-18T10:20:29.505546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions for training","metadata":{}},{"cell_type":"code","source":"def cross_entropy_loss(*, logits, labels):\n    logp = jax.nn.log_softmax(logits)\n    return -jnp.mean(jnp.sum(logp * labels, axis=1))\n\n\n# # optax has softmax_cross_entropy func\n# def cross_entropy_loss(*, logits, labels):\n#     return optax.softmax_cross_entropy(logits=logits, labels=labels).mean()\n\n\ndef get_loss_and_score(params_repl, steps, dataset):\n    loss_list = []\n    pred_list = []\n    label_list = []\n    for i, batch in zip(tqdm.trange(steps), dataset.as_numpy_iterator()):\n        logits = model_apply_repl(params_repl, batch)\n        label = batch['label']\n\n        pred_list.extend(jax.device_get(logits).tolist())\n        label_list.extend(jax.device_get(label).tolist())\n        loss_list.append(\n            jax.device_get(cross_entropy_loss(logits=logits, labels=label)))\n    return np.mean(loss_list), f1_score(np.argmax(label_list,\n                                                  axis=-1).flatten(),\n                                        np.argmax(pred_list,\n                                                  axis=-1).flatten(),\n                                        average=\"macro\")\n\n\ndef make_update_fn(*, apply_fn, loss_fn, tx):\n    \"\"\"Returns update step for data parallel training.\"\"\"\n\n    def update_fn(params, opt_state, batch):\n\n        def calc_loss_fn(params, images, labels):\n            logits, new_model_state = apply_fn(params,\n                                               images,\n                                               train=True,\n                                               mutable=[\"batch_stats\"])\n            return loss_fn(logits=logits, labels=labels), new_model_state\n\n        grad_fn = jax.value_and_grad(calc_loss_fn, has_aux=True)\n        (loss, new_model_state), grad = grad_fn(params, batch[\"image\"],\n                                                batch[\"label\"])\n\n        grad = jax.tree_map(lambda x: jax.lax.pmean(x, axis_name='batch'),\n                            grad)\n        updates, opt_state = tx.update(grad, opt_state)\n        params = optax.apply_updates(params, updates)\n        loss = jax.lax.pmean(loss, axis_name='batch')\n\n        # update batch_stats\n        params = flax.core.unfreeze(params)\n        params[\"batch_stats\"] = new_model_state[\"batch_stats\"]\n        params = flax.core.freeze(params)\n\n        return params, opt_state, loss\n\n    return jax.pmap(update_fn, axis_name='batch', donate_argnums=(0, ))\n","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:29.510820Z","iopub.execute_input":"2022-12-18T10:20:29.511505Z","iopub.status.idle":"2022-12-18T10:20:29.525536Z","shell.execute_reply.started":"2022-12-18T10:20:29.511467Z","shell.execute_reply":"2022-12-18T10:20:29.524720Z"},"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')()\n\n# ==========================\n# optimizer\n# ==========================\ntx = optax.sgd(learning_rate=0.1,\n               momentum=0.9,\n               accumulator_dtype='bfloat16' if tpu else None)\nupdate_fn_repl = make_update_fn(apply_fn=model.apply,\n                                loss_fn=cross_entropy_loss,\n                                tx=tx)\nopt_state = tx.init(params)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:29.527147Z","iopub.execute_input":"2022-12-18T10:20:29.527878Z","iopub.status.idle":"2022-12-18T10:20:37.601160Z","shell.execute_reply.started":"2022-12-18T10:20:29.527840Z","shell.execute_reply":"2022-12-18T10:20:37.600063Z"},"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['image'], train=False),\n    axis_name='batch')\nopt_state_repl = flax.jax_utils.replicate(opt_state)","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:37.602878Z","iopub.execute_input":"2022-12-18T10:20:37.603604Z","iopub.status.idle":"2022-12-18T10:20:37.746338Z","shell.execute_reply.started":"2022-12-18T10:20:37.603554Z","shell.execute_reply":"2022-12-18T10:20:37.745339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model training","metadata":{}},{"cell_type":"code","source":"train_losses_list = []\nval_losses_list = []\nval_f1_list = []\nsteps_per_epoch = NUM_TRAIN_IMAGES // BATCH_SIZE\n\nfor e in range(EPOCHS):\n    for step, batch in zip(\n            tqdm.trange(1, steps_per_epoch + 1),\n            train_dataset.as_numpy_iterator(),\n    ):\n\n        params_repl, opt_state_repl, loss_repl = update_fn_repl(\n            params_repl, opt_state_repl, batch)\n        train_losses_list.append(loss_repl[0])\n\n    val_loss, val_f1 = get_loss_and_score(params_repl=params_repl,\n                                          steps=NUM_VAL_IMAGES // BATCH_SIZE,\n                                          dataset=val_dataset)\n    val_losses_list.append(val_loss)\n    val_f1_list.append(val_f1)\n    print(\n        f\"Epoch: {e} train_loss: {np.mean(train_losses_list[-steps_per_epoch:]):.5f} val_loss: {val_loss}, val_f1: {val_f1}\"\n    )\n\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 3, 1)\nplt.plot(train_losses_list)\n\nplt.subplot(1, 3, 2)\nplt.plot(val_losses_list)\n\nplt.subplot(1, 3, 3)\nplt.plot(val_f1_list)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-18T10:20:37.747873Z","iopub.execute_input":"2022-12-18T10:20:37.748219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save params","metadata":{}},{"cell_type":"code","source":"from flax.serialization import to_bytes, from_bytes\nfrom flax.linen import FrozenDict\n\n\ndef save_params(params: FrozenDict, path: str) -> None:\n    serialized_params = to_bytes(params)\n    with open(path, 'wb') as f:\n        f.write(serialized_params)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_params(params=params_repl, path=\"/kaggle/working/resnet.weight\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}