{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-17T03:03:19.024858Z","iopub.execute_input":"2022-07-17T03:03:19.025309Z","iopub.status.idle":"2022-07-17T03:03:19.064348Z","shell.execute_reply.started":"2022-07-17T03:03:19.025202Z","shell.execute_reply":"2022-07-17T03:03:19.063782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install --upgrade pip\n#!pip install --upgrade \"jax[tpu]>=0.2.16\" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html\n#!pip install --upgrade jax jaxlib\n\n!pip install --upgrade pip\n!pip install --upgrade \"jax[cpu]\"\n!pip install --upgrade dm-haiku optax","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:08:25.841473Z","iopub.execute_input":"2022-07-17T07:08:25.842592Z","iopub.status.idle":"2022-07-17T07:08:40.468108Z","shell.execute_reply.started":"2022-07-17T07:08:25.842484Z","shell.execute_reply":"2022-07-17T07:08:40.467071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import functools as ft\n\nimport jax\nimport jax.numpy as jnp\nimport jax.random as jr\nimport jax.nn as jnn\n\nimport numpy as np\nimport optax\nimport pandas as pd\n\nimport haiku as hk\n\nimport haiku.initializers as hki\n\nimport logging","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:09:42.360406Z","iopub.execute_input":"2022-07-17T07:09:42.361240Z","iopub.status.idle":"2022-07-17T07:09:42.367850Z","shell.execute_reply.started":"2022-07-17T07:09:42.361185Z","shell.execute_reply":"2022-07-17T07:09:42.367084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dataset(filename='/kaggle/input/digit-recognizer/train.csv', filename1='/kaggle/input/digit-recognizer/test.csv'):\n    train_data = pd.read_csv(filename)\n    test = pd.read_csv(filename1).values[:, :]\n\n    train_y = train_data.values[:, 0]\n    train_x = train_data.values[:, 1:]\n\n    train_x = (train_x - 128.0) / 255.0\n    test = (test - 128.0) / 255.0\n\n    train_x = train_x.reshape((-1, 28, 28, 1))\n    test = test.reshape((-1, 28, 28, 1))\n\n    return jnp.array(train_x), jnp.array(train_y, dtype=jnp.int32), jnp.array(test)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:09:45.818916Z","iopub.execute_input":"2022-07-17T07:09:45.819338Z","iopub.status.idle":"2022-07-17T07:09:45.826972Z","shell.execute_reply.started":"2022-07-17T07:09:45.819299Z","shell.execute_reply":"2022-07-17T07:09:45.826210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_generator_parallel(x, y, rng_key, batch_size, num_devices):\n    def batch_generator():\n        n = x.shape[0]\n        key = rng_key\n        kk = batch_size // num_devices\n        while True:\n            key, k1 = jax.random.split(key)\n            perm = jax.random.choice(k1, n, shape=(batch_size,))\n\n            yield x[perm, :, :, :].reshape(num_devices, kk, *x.shape[1:]), y[perm].reshape(num_devices, kk, *y.shape[1:])\n\n    return batch_generator()\n","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:09:47.698786Z","iopub.execute_input":"2022-07-17T07:09:47.699424Z","iopub.status.idle":"2022-07-17T07:09:47.706035Z","shell.execute_reply.started":"2022-07-17T07:09:47.699387Z","shell.execute_reply":"2022-07-17T07:09:47.705251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvNetHybrid(hk.Module):\n    def __init__(self, dropout=0.5):\n        super(ConvNetHybrid, self).__init__()\n        self.dropout = dropout\n        scale_init = hki.Constant(1.0)\n        offset_init = hki.Constant(1e-8)\n        self.bn = lambda: hk.BatchNorm(True, True, 0.98)\n\n    def __call__(self, inputs, is_training=True):\n        dropout = self.dropout if is_training else 0.0\n        lc_init = hki.VarianceScaling(1.0, 'fan_in', 'truncated_normal')\n\n        # Input regularizer\n        x = hk.Conv2D(output_channels=32, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init)(inputs)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n        x = hk.MaxPool(window_shape=2, strides=2, padding=\"SAME\")(x)\n        \n        x = hk.Conv2D(output_channels=64, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init, b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n\n        x = hk.Conv2D(output_channels=64, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init, b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n        x = hk.MaxPool(window_shape=2, strides=2, padding=\"SAME\")(x)\n        \n        x = hk.Conv2D(output_channels=128, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init, b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n\n        x = hk.Conv2D(output_channels=128, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init,\n                      b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n        x = hk.MaxPool(window_shape=2, strides=2, padding=\"SAME\")(x)\n        \n        x = hk.Conv2D(output_channels=256, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init, b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n        \n        x = hk.Conv2D(output_channels=256, kernel_shape=3, stride=1, padding=\"SAME\", w_init=lc_init, b_init=hki.Constant(1e-6))(x)\n        x = self.bn()(x, is_training)\n        x = jnn.gelu(x, approximate=False)\n        x = hk.MaxPool(window_shape=2, strides=2, padding=\"SAME\")(x)\n\n        y = jnp.mean(x, axis=(1, 2))\n\n        lc_init = hki.VarianceScaling(1.0)\n        y = hk.Linear(784, w_init=lc_init, b_init=hki.Constant(1e-6))(y)\n        y = hk.dropout(hk.next_rng_key(), dropout, y)\n        y = hk.Linear(256, w_init=lc_init, b_init=hki.Constant(1e-6))(y)\n        y = hk.dropout(hk.next_rng_key(), dropout, y)\n        y = hk.Linear(10, w_init=lc_init, b_init=hki.Constant(1e-6))(y)\n\n        return y - jnn.logsumexp(y, axis=1, keepdims=True) \n","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:54:05.709468Z","iopub.execute_input":"2022-07-17T07:54:05.710160Z","iopub.status.idle":"2022-07-17T07:54:05.734468Z","shell.execute_reply.started":"2022-07-17T07:54:05.710124Z","shell.execute_reply":"2022-07-17T07:54:05.733561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_forward_fn(dropout=0.5):\n    def forward_fn(x: jnp.ndarray, is_training: bool = True) -> jnp.ndarray:\n        n = ConvNetHybrid(dropout)\n        return n(x, is_training=is_training)\n\n    return forward_fn","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:54:07.312735Z","iopub.execute_input":"2022-07-17T07:54:07.313414Z","iopub.status.idle":"2022-07-17T07:54:07.320114Z","shell.execute_reply.started":"2022-07-17T07:54:07.313366Z","shell.execute_reply":"2022-07-17T07:54:07.319123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@ft.partial(jax.jit, static_argnums=(0, 6, 7))\ndef lm_loss_fn(forward_fn, params, state, rng, x, y, is_training: bool = True, num_classes: int = 10):\n    y_pred, state = forward_fn(params, state, rng, x, is_training)\n\n    l2_loss = 0.1 * sum(jnp.sum(jnp.square(p)) for p in jax.tree_util.tree_leaves(params))\n    y_hot = jnn.one_hot(y, num_classes=num_classes)\n    #logits = jnp.exp(y_pred)\n    return -jnp.mean(y_hot * y_pred) + 1e-6 * l2_loss, state","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:54:09.740691Z","iopub.execute_input":"2022-07-17T07:54:09.741666Z","iopub.status.idle":"2022-07-17T07:54:09.749001Z","shell.execute_reply.started":"2022-07-17T07:54:09.741630Z","shell.execute_reply":"2022-07-17T07:54:09.748156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GradientUpdater:\n    def __init__(self, net_init, loss_fn, optimizer: optax.GradientTransformation):\n        self._net_init = net_init\n        self._loss_fn = loss_fn\n        self._opt = optimizer\n\n    def init(self, master_rng, x):\n        out_rng, init_rng = jax.random.split(master_rng)\n        params, state = self._net_init(init_rng, x)\n        opt_state = self._opt.init(params)\n        return jnp.array(0), out_rng, params, state, opt_state\n\n    def update(self, num_steps, rng, params, state, opt_state, x: jnp.ndarray, y: jnp.ndarray):\n        rng, new_rng = jax.random.split(rng)\n\n        (loss, state), grads = jax.value_and_grad(self._loss_fn, has_aux=True)(params, state, rng, x, y)\n\n        grads = jax.lax.pmean(grads, axis_name='j')\n\n        updates, opt_state = self._opt.update(grads, opt_state, params)\n\n        params = optax.apply_updates(params, updates)\n\n        metrics = {\n            'step': num_steps,\n            'loss': loss,\n        }\n\n        return num_steps + 1, new_rng, params, state, opt_state, metrics\n\n\ndef replicate_tree(t, num_devices):\n    return jax.tree_map(lambda x: jnp.array([x] * num_devices), t)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:54:12.524671Z","iopub.execute_input":"2022-07-17T07:54:12.525065Z","iopub.status.idle":"2022-07-17T07:54:12.539046Z","shell.execute_reply.started":"2022-07-17T07:54:12.525031Z","shell.execute_reply":"2022-07-17T07:54:12.537963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logging.getLogger().setLevel(logging.INFO)\ngrad_clip_value = 1.0\nlearning_rate = 0.003\nbatch_size = 168\ndropout = 0.5\nmax_steps = 1000\nnum_devices = jax.local_device_count()\nrng = jr.PRNGKey(111)\n\nx, y, test = load_dataset()\n\nprint(\"Number of training examples :::::: \", x.shape[0])\n\nrng, rng_key = jr.split(rng)\n\ntrain_dataset = get_generator_parallel(x, y, rng_key, batch_size, num_devices)\n\nforward_fn = build_forward_fn(dropout)\nforward_fn = hk.transform_with_state(forward_fn)\n\nforward_apply = forward_fn.apply\nloss_fn = ft.partial(lm_loss_fn, forward_apply)\n\nscheduler = optax.exponential_decay(init_value=learning_rate, transition_steps=200, decay_rate=0.99)\n\noptimizer = optax.chain(\n    optax.adaptive_grad_clip(grad_clip_value),\n    #optax.sgd(learning_rate=learning_rate, momentum=0.95, nesterov=True),\n    optax.scale_by_radam(),\n    #optax.scale_by_adam(),\n    optax.scale_by_schedule(scheduler),\n    optax.scale(-1.0)\n)\n\nupdater = GradientUpdater(forward_fn.init, loss_fn, optimizer)\n\nprint('Initializing parameters...')\n\nrng1, rng = jr.split(rng)\na = next(train_dataset)\nw, z = a\nnum_steps, rng2, params, state, opt_state = updater.init(rng1, w[0, :, :, :, :])\n\nrng1, rng = jr.split(rng)\nparams_multi_device = params\nopt_state_multi_device = opt_state\nnum_steps_replicated = replicate_tree(num_steps, num_devices)\nrng_replicated = rng1\nstate_multi_device = state\n\nbatch_update = jax.pmap(updater.update, axis_name='j', in_axes=(0, None, None, None, None, 0, 0),\n                        out_axes=(0, None, None, None, None, 0))","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:57:09.739453Z","iopub.execute_input":"2022-07-17T07:57:09.739821Z","iopub.status.idle":"2022-07-17T07:57:15.564075Z","shell.execute_reply.started":"2022-07-17T07:57:09.739791Z","shell.execute_reply":"2022-07-17T07:57:15.562905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Starting train loop ++++++++...')\nfor i, (w, z) in zip(range(max_steps), train_dataset):\n    if (i + 1) % 50 == 0:\n        print(f'Step {i} computing forward-backward pass')\n    num_steps_replicated, rng_replicated, params_multi_device, state_multi_device, opt_state_multi_device, metrics = batch_update(\n        num_steps_replicated, rng_replicated, params_multi_device, state_multi_device, opt_state_multi_device, w, z)\n\n    if (i + 1) % 50 == 0:\n        print(f'At step {i} the loss is {metrics}')\n\nprint('Starting evaluation loop +++++++++++++++')\nrng1, rng = jr.split(rng)\nstate = state_multi_device\nrng = rng1\nparams = params_multi_device\n\nfn = jax.jit(forward_apply, static_argnames=['is_training'])\n\nprint(\"Number of testing examples ::::: \", test.shape[0])\n\nres = np.zeros(test.shape[0], dtype=np.int64)\nn1 = test.shape[0]\n\ncount = n1 // 100\nfor j in range(count):\n    (rng,) = jr.split(rng, 1)\n    a, b = j * 100, (j + 1) * 100\n    logits, _ = fn(params, state, rng, test[a:b, :, :, :], is_training=False)\n    res[a:b] = np.array(jnp.argmax(jnp.exp(logits), axis=1), dtype=np.int64)\n\ndf = pd.DataFrame({'ImageId': np.arange(1, n1 + 1, dtype=np.int64), 'Label': res})\n\ndf.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T07:57:20.571189Z","iopub.execute_input":"2022-07-17T07:57:20.571608Z","iopub.status.idle":"2022-07-17T08:02:37.705714Z","shell.execute_reply.started":"2022-07-17T07:57:20.571576Z","shell.execute_reply":"2022-07-17T08:02:37.704479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}