{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.15","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":84896,"databundleVersionId":10305135,"sourceType":"competition"}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Easy Tutorial on Distributed Training in JAX with PyTorch Frame Data Loading\n\nAuthor: Yiwen Yuan\nSuggested Device: TPU VM v3-8\n\n[PyTorch Frame](https://github.com/pyg-team/pytorch-frame) is a deep learning extension for PyTorch, designed for heterogeneous tabular data with different column type. It is super easy to load tables into PyTorch Frame and perform necessary stats computation for different semantic types. \n\nThis tutorial is a simple example of distributed training in JAX using PyTorch Frame as data loader.","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install pytorch_frame","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:26:28.990378Z","iopub.execute_input":"2024-12-12T06:26:28.990701Z","iopub.status.idle":"2024-12-12T06:26:32.841608Z","shell.execute_reply.started":"2024-12-12T06:26:28.990673Z","shell.execute_reply":"2024-12-12T06:26:32.840071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import necessary libraries\n\nimport time\nfrom functools import partial\n\nimport jax\nimport jax.numpy as jnp\nimport pandas as pd\nimport torch\nfrom jax import jit, random, value_and_grad, vmap\nfrom jax.nn import swish\nfrom tqdm import tqdm\n\nfrom torch_frame import categorical, numerical, timestamp\nfrom torch_frame.data import DataLoader, Dataset\nfrom torch_frame.data.stats import StatType","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:00:02.336581Z","iopub.execute_input":"2024-12-12T06:00:02.33702Z","iopub.status.idle":"2024-12-12T06:00:02.341797Z","shell.execute_reply.started":"2024-12-12T06:00:02.33699Z","shell.execute_reply":"2024-12-12T06:00:02.341171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Constants <3\n\nNUM_FEAT = 24\nNUM_DEVICES = jax.device_count()\nBATCH_SIZE = 512\n# Two layer simple MLP\nLAYER_SIZES = [24, 1024, 1024, 1]\nPARAM_SCALE = 0.01\nINIT_LR = 0.0001\nDECAY_RATE = 0.95\nDECAY_STEPS = 5\nNUM_EPOCHS = 5\nTARGET_COL = 'Premium Amount'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:00:04.976558Z","iopub.execute_input":"2024-12-12T06:00:04.976918Z","iopub.status.idle":"2024-12-12T06:00:10.670624Z","shell.execute_reply.started":"2024-12-12T06:00:04.976888Z","shell.execute_reply":"2024-12-12T06:00:10.669659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Number of devices: {NUM_DEVICES}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:00:10.672122Z","iopub.execute_input":"2024-12-12T06:00:10.6724Z","iopub.status.idle":"2024-12-12T06:00:10.676713Z","shell.execute_reply.started":"2024-12-12T06:00:10.672371Z","shell.execute_reply":"2024-12-12T06:00:10.675962Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Data Loading with PyTorch Frame**\n\nTo load data with PyTorch Frame, you need to specify the semantic column types.","metadata":{}},{"cell_type":"code","source":"col_to_stype = {\n    'Age': numerical,\n    'Annual Income': numerical,\n    'Marital Status': categorical,\n    'Number of Dependents': numerical,\n    'Education Level': categorical,\n    'Occupation': categorical,\n    'Health Score': numerical,\n    'Location': categorical,\n    'Policy Type': categorical,\n    'Previous Claims': numerical,\n    'Vehicle Age': numerical,\n    'Credit Score': numerical,\n    'Insurance Duration': numerical,\n    'Policy Start Date': timestamp,\n    'Customer Feedback': categorical,\n    'Smoking Status': categorical,\n    'Exercise Frequency': categorical,\n    'Property Type': categorical\n}\ntest_dataset = Dataset(df=pd.read_csv('/kaggle/input/playground-series-s4e12/test.csv'),\n                       col_to_stype=col_to_stype)\ncol_to_stype = col_to_stype.copy()\ncol_to_stype['Premium Amount'] = numerical\n\ndataset = Dataset(df=pd.read_csv('/kaggle/input/playground-series-s4e12/train.csv'),\n                  col_to_stype=col_to_stype, target_col=TARGET_COL)\n\n# Saves the materialized tensors for easy reuse\ndataset.materialize(path='/kaggle/working/data.pt')\ntest_dataset.materialize(path='/kaggle/working/test_data.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:00:10.677777Z","iopub.execute_input":"2024-12-12T06:00:10.678052Z","iopub.status.idle":"2024-12-12T06:04:14.819223Z","shell.execute_reply.started":"2024-12-12T06:00:10.678025Z","shell.execute_reply":"2024-12-12T06:04:14.818219Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Nan Imputation**\n\nThe original data contains nans. We need to fill the nan values. We use MEAN to impute nan for numerical columns, MODE for categorical columns and MEDIAN for timestamp columns. The stats are already calculated as part of the PyTorch Frame data materialization process.","metadata":{}},{"cell_type":"code","source":"# Use mean value to fill nans in numerical columns\n# Use mode value to fill nans in categorical columns\n# Use newest time to fill nans in timestamp columns\nfill_vals = []\nmeans = []\nstds = []\nfor stype in [numerical, categorical, timestamp]:\n    col_names = dataset.tensor_frame.col_names_dict[stype]\n    for col_name in col_names:\n        if stype == numerical:\n            fill_vals.append(dataset.col_stats[col_name][StatType.MEAN])\n            means.append(dataset.col_stats[col_name][StatType.MEAN])\n            stds.append(dataset.col_stats[col_name][StatType.STD])\n        elif stype == categorical:\n            fill_vals.append(0)\n            means.append(0.)\n            stds.append(1.)\n        elif stype == timestamp:\n            fill_vals += dataset.col_stats[col_name][StatType.NEWEST_TIME]\n            means += dataset.col_stats[col_name][StatType.MEDIAN_TIME]\n            stds += [1.] * 7\n        else:\n            raise ValueError(\"Unsupported stype\")\nfill_vals = jnp.array(fill_vals)\nmeans = jnp.array(means)\nstds = jnp.array(stds)\n\ny_mean = torch.mean(dataset.tensor_frame.y).cpu().numpy()\ny_std = torch.std(dataset.tensor_frame.y).cpu().numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:05:58.241851Z","iopub.execute_input":"2024-12-12T06:05:58.242677Z","iopub.status.idle":"2024-12-12T06:05:58.262972Z","shell.execute_reply.started":"2024-12-12T06:05:58.242637Z","shell.execute_reply":"2024-12-12T06:05:58.261754Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**2-layer MLP in JAX with distributed training capabilities**","metadata":{}},{"cell_type":"code","source":"def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):\n    def random_layer_params(m, n, key, scale=1e-2):\n        w_key, b_key = random.split(key)\n        return (scale * random.normal(w_key, (n, m)),\n                scale * random.normal(b_key, (n, )))\n\n    keys = random.split(key, len(sizes))\n    return [\n        random_layer_params(m, n, k, scale)\n        for m, n, k in zip(sizes[:-1], sizes[1:], keys)\n    ]\n\ndef predict(params, image):\n    activations = image\n    for w, b in params[:-1]:\n        outputs = jnp.dot(w, activations) + b\n        activations = swish(outputs)\n\n    final_w, final_b = params[-1]\n    logits = jnp.dot(final_w, activations) + final_b\n    return logits\n\nbatched_predict = vmap(predict, in_axes=(None, 0))\n\ndef loss(params, images, targets):\n    logits = batched_predict(params, images)\n    return jnp.mean(jnp.abs(logits - targets))\n\n@partial(jax.pmap, axis_name='devices', in_axes=(None, 0, 0, None),\n         out_axes=(None, 0))\ndef update(params, x, y, epoch_number):\n    loss_value, grads = value_and_grad(loss)(params, x, y)\n    grads = [(jax.lax.psum(dw, 'devices') / NUM_DEVICES,\n              jax.lax.psum(db, 'devices') / NUM_DEVICES)\n             for dw, db in grads]\n    lr = INIT_LR * DECAY_RATE**(epoch_number / DECAY_STEPS)\n    return [(w - lr * dw, b - lr * db)\n            for (w, b), (dw, db) in zip(params, grads)], loss_value\n\n@jit\ndef batched_mae(params, images, targets):\n    images = jnp.reshape(images, (len(images), NUM_FEAT))\n    predicted_targets = batched_predict(params, images)\n    return jnp.mean(jnp.abs(predicted_targets - targets))\n\n\ndef mae(params, data_loader):\n    maes = []\n    for tf in data_loader:\n        x = torch.cat([\n            tf.feat_dict[categorical], tf.feat_dict[numerical],\n            tf.feat_dict[timestamp].squeeze(1)\n        ], dim=1).numpy()\n        y = (tf.y.numpy() - y_mean) / y_std\n        nan_mask = jnp.isnan(x)\n        if x.any():\n            x = jnp.where(nan_mask, fill_vals, x)\n        x = (x - means) / stds\n        maes.append(batched_mae(params, x, y))\n    return jnp.mean(jnp.array(maes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:06:00.353998Z","iopub.execute_input":"2024-12-12T06:06:00.35498Z","iopub.status.idle":"2024-12-12T06:06:00.44252Z","shell.execute_reply.started":"2024-12-12T06:06:00.354941Z","shell.execute_reply":"2024-12-12T06:06:00.441549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"init_params = init_network_params(LAYER_SIZES, random.PRNGKey(0),\n                                  scale=PARAM_SCALE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:06:03.274653Z","iopub.execute_input":"2024-12-12T06:06:03.27538Z","iopub.status.idle":"2024-12-12T06:06:05.027892Z","shell.execute_reply.started":"2024-12-12T06:06:03.275332Z","shell.execute_reply":"2024-12-12T06:06:05.026703Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Actual Training**","metadata":{}},{"cell_type":"code","source":"train_loader = DataLoader(dataset.tensor_frame, batch_size=BATCH_SIZE,\n                          shuffle=True, drop_last=True)\nparams = init_params\nfor epoch in range(1, NUM_EPOCHS + 1):\n    start_time = time.time()\n    losses = []\n    for tf in tqdm(train_loader):\n        x = torch.cat([\n            tf.feat_dict[categorical], tf.feat_dict[numerical],\n            tf.feat_dict[timestamp].squeeze(1)\n        ], dim=1).numpy()\n        y = (tf.y.numpy() - y_mean) / y_std\n        nan_mask = jnp.isnan(x)\n        if nan_mask.any():\n            x = jnp.where(nan_mask, fill_vals, x)\n        x = (x - means) / stds\n        x = jnp.reshape(x, (NUM_DEVICES, BATCH_SIZE // NUM_DEVICES, NUM_FEAT))\n        y = jnp.reshape(y, (NUM_DEVICES, BATCH_SIZE // NUM_DEVICES, 1))\n        params, loss_value = update(params, x, y, epoch)\n        losses.append(jnp.sum(loss_value))\n    epoch_time = time.time() - start_time\n    train_mae = mae(params, train_loader)\n    print(f\"Epoch {epoch} in {epoch_time:0.2f} sec\")\n    print(f\"Training set loss {jnp.mean(jnp.array(losses))}\")\n    print(f\"Training set mae {train_mae}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:06:05.058806Z","iopub.execute_input":"2024-12-12T06:06:05.05978Z","iopub.status.idle":"2024-12-12T06:08:03.434577Z","shell.execute_reply.started":"2024-12-12T06:06:05.059739Z","shell.execute_reply":"2024-12-12T06:08:03.433444Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Generating Predictions on Test Data**","metadata":{}},{"cell_type":"code","source":"test_loader = DataLoader(test_dataset.tensor_frame, batch_size=BATCH_SIZE,\n                         shuffle=False, drop_last=False)\nresults = []\n\nparallel_predict = jax.pmap(vmap(predict, in_axes=(None, 0)), in_axes=(None, 0))\n\nfor tf in tqdm(test_loader):\n    x = torch.cat([\n        tf.feat_dict[categorical], tf.feat_dict[numerical],\n        tf.feat_dict[timestamp].squeeze(1)\n    ], dim=1).numpy()\n    nan_mask = jnp.isnan(x)\n    if nan_mask.any():\n        x = jnp.where(nan_mask, fill_vals, x)\n    x = (x - means) / stds\n    x = jnp.reshape(x, (NUM_DEVICES, len(x) // NUM_DEVICES, NUM_FEAT))\n    result = parallel_predict(params, x).reshape(-1) * y_std + y_mean\n    results.append(result)\n\noutputs = jnp.concatenate(results)\nsubmission = pd.read_csv(\"/kaggle/input/playground-series-s4e12/sample_submission.csv\")\nsubmission[TARGET_COL] = outputs\nsubmission = submission[['id', TARGET_COL]]\nsubmission.to_csv('/kaggle/working/final_submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-12T06:10:55.419884Z","iopub.execute_input":"2024-12-12T06:10:55.420259Z","iopub.status.idle":"2024-12-12T06:11:05.256034Z","shell.execute_reply.started":"2024-12-12T06:10:55.420229Z","shell.execute_reply":"2024-12-12T06:11:05.255071Z"}},"outputs":[],"execution_count":null}]}