{"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":"markdown","source":"# OTTO Recommender \n---\n\n#### <a href='#notebook-setup'> 📒 Notebook Setup </a> | <a href='#hyperparameters'> ⚙ Hyperparameters | <a href='#data-processing'> 📦 Data Processing </a> | <a href='#model-architecture'> 🏗 Model Architecture </a> | <a href='#training'> 🔥 Training </a> | <a href='#prediction'> 🎯 Prediction </a>","metadata":{}},{"cell_type":"code","source":"\"\"\"\n- SSL: User embedding = Average of last hidden state. Contrastive loss between training and last 16 items to train user embeddings\n- Multi embedding user recommendation\n- Aggregate two hop item neighborhood with GCN?? Add \n- Start with last item prediction + SSL w. larger seq len on test data??\n- Last item prediction + SSL + MLM on Cart/Orders\n- Both NIP and LIP\n- Use ListNet loss??\n\"\"\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 📒 Notebook Setup\n---\n\n<a name='notebook-setup'>","metadata":{}},{"cell_type":"code","source":"!pip install --upgrade jaxlib jax flax optax -q\n!pip install --upgrade omegaconf -q\n!pip install --upgrade datasets -q","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import requests\nimport flax\nimport jax\nimport os\n\ndef kaggle_tpu_setup():\n    if 'TPU_NAME' not in os.environ:\n        print('TPU not found')\n        return\n    os.environ['TF_XLA_FLAGS'] = '--tf_xla_enable_xla_devices'\n    url = 'http:' + os.environ['TPU_NAME'].split(':')[1] + ':8475/requestversion/tpu_driver_nightly'\n    resp = requests.post(url)\n    jax.config.FLAGS.jax_xla_backend = 'tpu_driver'\n    jax.config.FLAGS.jax_backend_target = os.environ['TPU_NAME']\n    jax.config.update('jax_default_matmul_precision', 'bfloat16')\n\nkaggle_tpu_setup()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Commonly Used Libraries\nfrom collections import Counter, defaultdict\nfrom dataclasses import dataclass, asdict\nfrom datetime import datetime\nfrom functools import partial\nfrom termcolor import colored\nfrom time import time, sleep\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\nimport pandas as pd\nimport numpy as np\n\nimport dataclasses\nimport random\nimport yaml\nimport math\nimport sys\nimport os\nimport re\n\n\n# JAX Ecosystem Imports\nimport jax.numpy as jnp\nimport flax.linen as nn\nimport flax.training.common_utils\n\nimport optax\nimport flax\nimport jax\n\ntqdm.pandas()\n\n# Import OmegaConf\nfrom omegaconf import OmegaConf\n\n\n# IPython Imports\nfrom IPython.core.magic import register_line_cell_magic\nfrom IPython import get_ipython, display\nfrom IPython.display import FileLink\n\n# Login to WandB\nimport wandb\nos.system(\"wandb login 3b335317f20548af7e3b941d09a6de9f1736bd8d\")\n\n# Setup Jupyter Notebook\ndef _setup_jupyter_notebook():\n    from IPython.core.interactiveshell import InteractiveShell\n    InteractiveShell.ast_node_interactivity = 'all'\n    ipython = get_ipython()\n    ipython.magic('matplotlib inline')\n    ipython.magic('load_ext autoreload')\n    ipython.magic('autoreload 2')\n_setup_jupyter_notebook()\n\n# Hyperparameters Magic Command\nclass AttrDict(dict):\n    def __init__(self, *args, **kwargs):\n        super(AttrDict, self).__init__(*args, **kwargs)\n        self.__dict__ = self\n\ndef read_yaml(filename):\n    with open(filename, 'r') as stream:\n        return AttrDict(yaml.safe_load(stream))\n\n@register_line_cell_magic\ndef hyperparameters(hp_var_name, cell):\n    with open('experiment.yaml', 'w') as f:\n        f.write(cell)\n    HP = OmegaConf.load('experiment.yaml')\n    get_ipython().user_ns[hp_var_name] = HP","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ⚙ Hyperparamters\n---\n\n<a name='hyperparameters'>","metadata":{}},{"cell_type":"code","source":"%%hyperparameters HP\n\n## Model Architecture ##\nmax_seq_len: 16\nitem_embedding_size: 128\n\nhidden_size: 1024\nhidden_activation_fn: 'approx_gelu'\nhidden_dropout_prob: 0.10\nintermediate_size: 4096\n\ninitializer_range: 0.02\nlayer_norm_epsilon: 1e-7 # Post Layer Norm\n\nattention_types: ['I2I', 'I2T', 'T2I']\nnum_attention_heads: 16\nnum_hidden_layers: 8\n\n\n## Model Training ##\nnum_train_epochs: 1\nper_device_train_batch_size: 64\nper_device_eval_batch_size: 64\n\n\n## Loss Function ## \nloss_fn_str: 'softmax_crossentropy' # 'SSM'\n\n\n## Cosine Decay LR Scheduler ##\nwarmup_ratio: 0.03125\npeak_lr: 3e-5\n\n\n## Adan Optimizer ##\nweight_decay: 1e-2\nmax_grad_norm: 1.0\nbeta_2: 0.98\nepsilon: 1e-6\nema_decay: null\n\n## Data Processing ##\nuse_test_data_only: true\nnum_proc: 4\n\n## Model Tracking ##\nlogging_freq: 100\nrandom_state: 69420","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_ITEMS = 2000000","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 📦 Data Processing\n\n<a name='data-processing'>","metadata":{}},{"cell_type":"code","source":"import datasets\n# Load raw dataset from json files\n# data_files = {\n#     'train': '/kaggle/input/otto-recommender-system/test.jsonl',\n#     'test': '/kaggle/input/otto-recommender-system/test.jsonl'\n# }\n# raw_dataset = datasets.load_dataset('json', data_files=data_files, num_proc=4)\n# raw_dataset['train'] = datasets.concatenate_datasets([raw_dataset['train'], raw_dataset['test']])\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MIN_TS, MAX_TS = 1659304800, 1662328791\nMAX_TARGETS_PER_SESSION = 4\n\ndef build_timestamp_features(unix_timestamps):\n    utc_timestamps = [datetime.utcfromtimestamp(unix_ts//1000) for unix_ts in unix_timestamps]\n    ts_days = [int((utc_ts-datetime.utcfromtimestamp(MIN_TS)).days) for utc_ts in utc_timestamps]\n    ts_weekdays = [int(utc_ts.weekday()) for utc_ts in utc_timestamps]\n    ts_hours = [int(utc_ts.hour) for utc_ts in utc_timestamps]\n    ts_abs = [(MAX_TS-unix_ts//1000)/(MAX_TS-MIN_TS) for unix_ts in unix_timestamps]\n    ts_session = [(max(ts_abs)-ts)/(max(ts_abs)-min(ts_abs)+1e-6) for ts in ts_abs]\n    return {\n        'ts_days': ts_days,\n        'ts_weekdays': ts_weekdays,\n        'ts_hours': ts_hours,\n        'ts_abs': ts_abs,\n        'ts_session': ts_session,\n    }\n\ndef pad_or_clip(array, length, pad_value):\n    if len(array) > length:\n        return array[:length]\n    pad_width = length-len(array)\n    return array + [pad_value]*pad_width\n\ndef process_train_session(session_dict):\n    events = session_dict['events']\n    num_train_events = random.randint(1, len(events))\n    num_train_events = min(num_train_events, HP.max_seq_len)\n    train_events, test_events = events[:num_train_events], events[num_train_events:]\n    \n    item_ids = [event['aid']+1 for event in train_events]\n    interaction_type_ids = [\n        {'clicks': 0, 'carts': 1, 'orders': 2}[event['type']]\n        for event in train_events\n    ]\n    pad_width = HP.max_seq_len-len(item_ids)\n    attention_mask = [1]*len(item_ids) + [0]*pad_width\n    item_ids = item_ids + [0]*pad_width\n    interaction_type_ids = interaction_type_ids + [0]*pad_width\n    assert len(item_ids) == len(interaction_type_ids) == len(attention_mask) == HP.max_seq_len\n    \n    timestamp_features = build_timestamp_features([event['ts'] for event in train_events])\n    timestamp_features = {k: v+[0]*pad_width for k, v in timestamp_features.items()}\n    \n    model_inputs = {\n        'item_ids': np.array(item_ids, dtype=np.int32),\n        'interaction_type_ids': np.array(interaction_type_ids, dtype=np.int32),\n        'attention_mask': np.array(attention_mask, dtype=np.int32),\n    }\n    \n    target_click_ids = [event['aid'] + 1 for event in test_events if event['type']=='clicks']\n    target_cart_ids = [event['aid'] + 1 for event in test_events if event['type']=='carts']\n    target_order_ids = [event['aid'] + 1 for event in test_events if event['type']=='orders']\n    model_outputs = {\n        'target_click_ids': pad_or_clip(target_click_ids, MAX_TARGETS_PER_SESSION, 0),\n        'target_cart_ids': pad_or_clip(target_cart_ids, MAX_TARGETS_PER_SESSION, 0),\n        'target_order_ids': pad_or_clip(target_order_ids, MAX_TARGETS_PER_SESSION, 0),\n    }\n    return {**model_inputs, **timestamp_features, **model_outputs}\n\n\ndef get_dataloader(dataset, batch_size):\n    steps_per_epoch = len(dataset)//batch_size\n    for batch_idx in range(steps_per_epoch):\n        yield dataset[batch_idx:batch_idx+batch_size]\n\ndef collate_fn(batch): \n    jnp_batch = {\n        'item_ids': jnp.array(batch['item_ids']),\n        'interaction_type_ids': jnp.array(batch['interaction_type_ids']),\n        'attention_mask': jnp.array(batch['attention_mask']),\n        \n        'timestamp_features': {\n            'ts_days': jnp.array(batch['ts_days']),\n            'ts_weekdays': jnp.array(batch['ts_weekdays']),\n            'ts_hours': jnp.array(batch['ts_hours']),\n            'ts_abs': jnp.array(batch['ts_abs']),\n            'ts_session': jnp.array(batch['ts_session']),\n        },\n        \n        'target_click_ids': jnp.array(batch['target_click_ids']),\n        'target_cart_ids': jnp.array(batch['target_cart_ids']),\n        'target_order_ids': jnp.array(batch['target_order_ids']),\n    }\n    jnp_batch = flax.training.common_utils.shard(jnp_batch)\n    return jnp_batch\n\n\n# DATA_FILE = '/kaggle/input/otto-recommender-system/train.jsonl'\n# raw_dataset = datasets.load_dataset('json', data_files=DATA_FILE, split='train[:300000]')\n\n# train_dataset = raw_dataset.map(\n#     process_train_session,\n#     num_proc=2,\n#     desc='Processing raw dataset',\n#     remove_columns=['session', 'events']\n# )\n# train_dataset.set_format(type='numpy')\n# train_dataset.save_to_disk('processed_dataset')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = datasets.load_from_disk('/kaggle/input/otto-preprocessing/processed_dataset')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def collate_fn(batch): \n#     jnp_batch = {\n#         'item_ids': jnp.array(batch['item_ids']),\n#         'interaction_type_ids': jnp.array(batch['interaction_type_ids']),\n#         'attention_mask': jnp.array(batch['attention_mask']),\n        \n#         'timestamp_features': {\n#             'ts_days': jnp.array(batch['ts_days']),\n#             'ts_weekdays': jnp.array(batch['ts_weekdays']),\n#             'ts_hours': jnp.array(batch['ts_hours']),\n#             'ts_abs': jnp.array(batch['ts_abs']),\n#             'ts_session': jnp.array(batch['ts_session']),\n#         },\n        \n#         'target_click_ids': jnp.array(batch['target_click_ids']),\n#         'target_cart_ids': jnp.array(batch['target_cart_ids']),\n#         'target_order_ids': jnp.array(batch['target_order_ids']),\n#     }\n#     jnp_batch = flax.training.common_utils.shard(jnp_batch)\n#     return jnp_batch\n\ndef collate_fn(batch): \n    jnp_batch = {\n        'item_ids': jnp.array(batch['item_ids']),\n        'interaction_type_ids': jnp.array(batch['interaction_type_ids']),\n        \n        'timestamp_features': {\n            'ts_days': jnp.array(batch['ts_days']),\n            'ts_weekdays': jnp.array(batch['ts_weekdays']),\n            'ts_hours': jnp.array(batch['ts_hours']),\n            'ts_abs': jnp.array(batch['ts_abs']),\n            'ts_session': jnp.array(batch['ts_session']),\n        },\n        \n        'target_item_id': jnp.array(batch['target_item_id']),\n    }\n    jnp_batch = flax.training.common_utils.shard(jnp_batch)\n    return jnp_batch","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🏗 Model Architecture\n\n<a name='model-architecture'>","metadata":{}},{"cell_type":"code","source":"default_init = jax.nn.initializers.normal(HP.initializer_range)\n\nclass TimestampEmbeddings(nn.Module):\n    config: dict\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.num_day_embeddings = 35\n        self.num_weekday_embeddings = 7\n        self.num_hour_embeddings = 24\n        \n        self.day_embeddings = nn.Embed(\n            self.num_day_embeddings,\n            self.config.hidden_size,\n            embedding_init=default_init,\n        )\n        self.weekday_embeddings = nn.Embed(\n            self.num_weekday_embeddings,\n            self.config.hidden_size,\n            embedding_init=default_init,\n        )\n        self.hour_embeddings = nn.Embed(\n            self.num_hour_embeddings,\n            self.config.hidden_size,\n            embedding_init=default_init,\n        )\n        self.abs_timestamp_proj = nn.Dense(\n            self.config.hidden_size,\n            kernel_init=default_init,\n        )\n        self.session_timestamp_proj = nn.Dense(\n            self.config.hidden_size,\n            kernel_init=default_init,\n        )\n        self.output_proj = nn.Dense(\n            self.config.hidden_size,\n            kernel_init=default_init,\n        )\n        \n        self.LayerNorm = nn.LayerNorm(self.config.layer_norm_epsilon)\n        \n    \n    def __call__(self, ts_days, ts_weekdays, ts_hours, ts_abs, ts_session):\n        day_embeds = self.day_embeddings(ts_days.astype('i4'))\n        weekday_embeds = self.weekday_embeddings(ts_weekdays.astype('i4'))\n        hour_embeds = self.hour_embeddings(ts_hours.astype('i4'))\n        abs_embeds = self.abs_timestamp_proj(jnp.expand_dims(ts_abs, axis=-1))\n        relative_embeds = self.session_timestamp_proj(jnp.expand_dims(ts_session, axis=-1))\n        \n        timestamp_embeds = day_embeds + weekday_embeds + hour_embeds + abs_embeds + relative_embeds\n        timestamp_embeds = self.output_proj(timestamp_embeds)\n        timestamp_embeds = self.LayerNorm(timestamp_embeds)\n        return timestamp_embeds\n\n\nclass ItemEmbeddings(nn.Module):\n    config: dict\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.item_embeddings = nn.Embed(\n            NUM_ITEMS,\n            self.config.item_embedding_size,\n        )\n        self.dense = nn.Dense(\n            self.config.hidden_size,\n            kernel_init=default_init,\n        )\n        self.LayerNorm = nn.LayerNorm(self.config.layer_norm_epsilon)\n        \n    def __call__(self, item_ids, interaction_type_ids):\n        item_embeds = self.item_embeddings(item_ids.astype('i4'))\n        interaction_embeds = jax.nn.one_hot(interaction_type_ids, num_classes=4).astype(item_embeds.dtype)\n        \n        # Combine embeddings and project to hidden size\n        hidden_states = jnp.concatenate([item_embeds, interaction_embeds], axis=-1)\n        hidden_states = self.dense(hidden_states)\n        \n        hidden_states = self.LayerNorm(hidden_states)\n        return item_embeds, hidden_states\n    \n    \nclass DisentangledSelfAttention(nn.Module):\n    \"\"\"\n    Compute item contextual correlation and timestamp correlation separtely\n    with different parameterizations.\n    \"\"\"\n    config: dict\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.per_head_dim = self.config.hidden_size//self.config.num_attention_heads\n        \n        self.query_proj = nn.DenseGeneral(\n            features=(self.config.num_attention_heads, self.per_head_dim),\n            kernel_init=default_init,\n            dtype=self.dtype\n        )\n        self.key_proj = nn.DenseGeneral(\n            features=(self.config.num_attention_heads, self.per_head_dim),\n            kernel_init=default_init,\n            dtype=self.dtype\n        )\n        self.value_proj = nn.DenseGeneral(\n            features=(self.config.num_attention_heads, self.per_head_dim),\n            kernel_init=default_init,\n            dtype=self.dtype,\n        )\n        \n        if 'I2T' in self.config.attention_types:\n            self.timestamp_key_proj = nn.DenseGeneral(\n                features=(self.config.num_attention_heads, self.per_head_dim),\n                kernel_init=default_init,\n                dtype=self.dtype\n            )\n        if 'T2I' in self.config.attention_types:\n            self.timestamp_query_proj = nn.DenseGeneral(\n                features=(self.config.num_attention_heads, self.per_head_dim),\n                kernel_init=default_init,\n                dtype=self.dtype\n            )\n        \n        self.output_proj = nn.DenseGeneral(\n            features=self.config.hidden_size,\n            axis=(-2, -1), # (num_attention_heads, per_head_dim) -> (hidden_size)\n            kernel_init=default_init,\n            dtype=self.dtype,\n        )\n        self.LayerNorm = nn.LayerNorm(epsilon=self.config.layer_norm_epsilon)\n        self.dropout = nn.Dropout(rate=self.config.hidden_dropout_prob)\n        \n    def __call__(self, hidden_states, timestamp_embeddings, deterministic=True):\n        query_states = self.query_proj(hidden_states)\n        key_states = self.key_proj(hidden_states)\n        value_states = self.value_proj(hidden_states)\n        \n        ts_key_states = self.timestamp_key_proj(timestamp_embeddings)\n        ts_query_states = self.timestamp_query_proj(timestamp_embeddings)\n        \n        scale_factor = len(self.config.attention_types)\n        depth = query_states.shape[-1]\n        scale = jnp.sqrt(depth*scale_factor).astype(query_states.dtype)\n        \n        i2i_attn_score = jnp.einsum('...qhd,...khd->hqk', query_states, key_states)\n        i2t_attn_score = jnp.einsum('...qhd,...khd->hqk', query_states, ts_key_states)\n        t2i_attn_score = jnp.einsum('...qhd,...khd->hqk', ts_query_states, key_states)\n        attention_score = (i2i_attn_score + i2t_attn_score + t2i_attn_score) / scale\n        \n        attention_weights = jax.nn.softmax(attention_score, axis=-1).astype(self.dtype)\n        attention_out = jnp.einsum('...hqk,...khd->...qhd', attention_weights, value_states)\n        \n        attention_out = self.output_proj(attention_out)\n        attention_out = self.dropout(attention_out, deterministic=deterministic)\n        out = self.LayerNorm(hidden_states+attention_out)\n        return out\n\n\nclass FFN(nn.Module):\n    config: dict\n    layer_idx: int\n    dtype: jnp.dtype\n    \n    def setup(self):\n        init_scale = 1/((self.layer_idx+1)**0.5)\n        kernel_init = jax.nn.initializers.normal(self.config.initializer_range*init_scale)\n        \n        self.intermediate_proj = nn.Dense(\n            self.config.intermediate_size,\n            kernel_init=kernel_init,\n            dtype=self.dtype,\n        )\n        self.output_proj = nn.Dense(\n            self.config.hidden_size,\n            kernel_init=kernel_init,\n            dtype=self.dtype,\n        )\n        self.dropout = nn.Dropout(rate=self.config.hidden_dropout_prob)\n        \n    def __call__(self, hidden_states, deterministic=True):\n        hidden_states = self.intermediate_proj(hidden_states)\n        if self.config.hidden_activation_fn == 'gelu_approx':\n            hidden_states = jax.nn.gelu(hidden_states, approximate=True)\n        hidden_states = self.output_proj(hidden_states)\n        hidden_states = self.dropout(hidden_states, deterministic=deterministic)\n        return hidden_states\n\n\nclass TransformerLayer(nn.Module):\n    config: dict\n    layer_idx: int\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.attention = DisentangledSelfAttention(self.config, self.dtype)\n        self.ffn = FFN(self.config, self.layer_idx, self.dtype)\n        self.PostLayerNorm = nn.LayerNorm(epsilon=self.config.layer_norm_epsilon)\n    \n    def __call__(self, hidden_states, timestamp_embeddings, deterministic):\n        attention_out = self.attention(hidden_states, timestamp_embeddings, deterministic)\n        ffn_out = self.ffn(attention_out, deterministic)\n        out = self.PostLayerNorm(ffn_out+attention_out)\n        return out\n\n\nclass PredictionHead(nn.Module):\n    config: dict\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.dense = nn.Dense(\n            self.config.item_embedding_size,\n            dtype=self.dtype,\n            kernel_init=default_init,\n        )\n        self.LayerNorm = nn.LayerNorm(self.config.layer_norm_epsilon)\n\n        self.decoder = nn.Dense(\n            NUM_ITEMS,\n            dtype=self.dtype,\n            use_bias=False,\n            kernel_init=default_init,\n        )\n        bias_init = jax.nn.initializers.zeros\n        self.bias = self.param('bias', bias_init, (NUM_ITEMS,))\n    \n    def __call__(self, hidden_states, item_embeddings):\n        hidden_states = self.dense(hidden_states)\n        hidden_states = jax.nn.gelu(hidden_states, approximate=True)\n        hidden_states = self.LayerNorm(hidden_states)\n        hidden_states = self.decoder.apply({'params': {'kernel': item_embeddings.T}}, hidden_states)\n        \n        bias = jnp.asarray(self.bias, self.dtype)\n        hidden_states += bias\n        return hidden_states\n\n\nclass Model(nn.Module):\n    config: dict\n    dtype: jnp.dtype\n    \n    def setup(self):\n        self.item_embeddings = ItemEmbeddings(self.config, self.dtype)\n        self.ts_embeddings = TimestampEmbeddings(self.config, self.dtype)\n        self.layers = [\n            TransformerLayer(self.config, name=str(layer_idx), layer_idx=layer_idx, dtype=self.dtype)\n            for layer_idx in range(self.config.num_hidden_layers)\n        ]\n        self.item_pred_head = PredictionHead(self.config, self.dtype)\n        \n    def __call__(\n        self,\n        item_ids,\n        interaction_type_ids,\n        timestamp_features,\n        #attention_mask,\n        deterministic=True,\n    ):\n        \n        item_embeds, hidden_states = self.item_embeddings(item_ids, interaction_type_ids)\n        timestamp_embeddings = self.ts_embeddings(**timestamp_features)\n        for layer in self.layers:\n            hidden_states = layer(hidden_states, timestamp_embeddings, deterministic)\n        \n        item_preds = self.item_pred_head(hidden_states, self.item_embeddings.item_embeddings.embedding)\n        #click_preds = item_preds[:, 0, :]\n        #cart_preds = item_preds[:, 1, :]\n        #order_preds = item_preds[:, 2, :]\n        #return click_preds, cart_preds, order_preds\n        \n        return item_preds[:, 0, :]\n        \n\ndummy_batch_size = 2\ndummy_input = {\n    'item_ids': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n    'interaction_type_ids': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n    'timestamp_features': {\n        'ts_days': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n        'ts_weekdays': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n        'ts_hours': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n        'ts_abs': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n        'ts_session': jnp.zeros((dummy_batch_size, HP.max_seq_len)),\n    }\n}\nkey = jax.random.PRNGKey(HP.random_state)\nmodel = Model(config=HP, dtype=jnp.bfloat16)\nparams = model.init(key, **dummy_input)\n_ = model.apply(params, **dummy_input)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# a[0].shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🔥 Training\n\n<a name='training'>","metadata":{}},{"cell_type":"code","source":"import optax\nimport flax\n\ndef build_lr_scheduler(peak_lr, warmup_ratio, total_train_steps):\n    warmup_steps = int(warmup_ratio * total_train_steps)\n    decay_steps = total_train_steps - warmup_steps\n    print(f'Warmup Steps: {warmup_steps} | Decay Steps: {decay_steps}')\n    \n    lr_scheduler = optax.warmup_cosine_decay_schedule(\n        init_value=0,\n        peak_value=peak_lr,\n        warmup_steps=warmup_steps,\n        decay_steps=decay_steps,\n        end_value=0,\n    )\n    return lr_scheduler\n\ndef build_tx(\n    lr_scheduler,\n    adam_beta_2,\n    adam_epsilon,\n    weight_decay,\n    max_grad_norm,\n    ema_decay,\n):\n    def weight_decay_mask(params):\n        params = flax.traverse_util.flatten_dict(params)\n        mask = {k: (k[-1] != 'bias' and k[-2:] != ('LayerNorm', 'scale')) for k, v in params.items()}\n        params = flax.traverse_util.unflatten_dict(params)\n        return flax.traverse_util.unflatten_dict(mask)\n\n    tx = optax.adamw(\n        learning_rate=lr_scheduler, \n        b1=0.9, \n        b2=adam_beta_2, \n        eps=adam_epsilon, \n        weight_decay=weight_decay, \n        mask=weight_decay_mask,\n    )\n\n    if max_grad_norm is not None:\n        tx = optax.chain(tx, optax.clip_by_global_norm(max_grad_norm))\n    if ema_decay is not None:\n        tx = optax.chain(tx, optax.ema(decay=ema_decay)) \n    return tx\n\ntrain_batch_size = HP.per_device_train_batch_size * jax.device_count()\neval_batch_size = HP.per_device_eval_batch_size * jax.device_count()\ntotal_train_steps = HP.num_train_epochs * (len(train_dataset) // train_batch_size)\n\nlr_scheduler = build_lr_scheduler(\n    peak_lr=HP.peak_lr,\n    warmup_ratio=HP.warmup_ratio,\n    total_train_steps=total_train_steps,\n)\n\ntx = build_tx(\n    lr_scheduler=lr_scheduler,\n    adam_beta_2=HP.beta_2,\n    adam_epsilon=HP.epsilon,\n    weight_decay=HP.weight_decay,\n    max_grad_norm=HP.max_grad_norm,\n    ema_decay=HP.ema_decay,\n)\nparams = flax.core.frozen_dict.unfreeze(params)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# params","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.restore('model_lip_loss296_step8000.msgpack', 'otto/runs/1mbkn2zd')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import flax.training.checkpoints\n\nparams = flax.training.checkpoints.restore_checkpoint(\n    '/kaggle/working/model_lip_loss296_step8000.msgpack', \n    target=params\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# params","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import flax.training.train_state\n\ndef create_state(model, params, tx):\n    state = flax.training.train_state.TrainState.create(\n        apply_fn=model.apply,\n        params=params,\n        tx=tx,\n    )\n    state = flax.jax_utils.replicate(state)\n    return state\nstate = create_state(model, params, tx)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import flax.training\n\ndef get_recall(logits, labels, k):\n    logits, labels = logits[:, 1:], labels[:, 1:]\n    top_pred_ids = jax.lax.approx_max_k(logits, k, reduction_dimension=-1)[1]\n    a = jax.nn.one_hot(top_pred_ids, logits.shape[-1]).sum(axis=1)\n    true_preds = jnp.where((a==1.)&(labels==1.), 1., 0.).sum(axis=-1)\n    return true_preds.mean()\n\n@partial(jax.pmap, axis_name='batch')\ndef train_step(state, dropout_rng, batch):\n    \n    def loss_fn(params):\n        model_outputs = state.apply_fn(\n            params,\n            item_ids=batch['item_ids'],\n            interaction_type_ids=batch['interaction_type_ids'],\n            #attention_mask=batch['attention_mask'],\n            timestamp_features=batch['timestamp_features'],\n            deterministic=False,\n            rngs={'dropout': dropout_rng},\n        )\n        logits = model_outputs\n        labels = jax.nn.one_hot(batch['target_item_id'], logits.shape[-1])\n        loss = optax.softmax_cross_entropy(logits, labels).sum()\n        \n        recall_20 = get_recall(logits, labels, 20)\n        recall_1000 = get_recall(logits, labels, 1000)\n        \n        return loss, (loss, recall_20, recall_1000)\n    \n    grad_fn = jax.value_and_grad(loss_fn, has_aux=True)\n    (total_loss, aux), grad = grad_fn(state.params)\n    grad = jax.lax.pmean(grad, axis_name='batch')\n    new_state = state.apply_gradients(grads=grad)\n    \n    aux = jax.lax.pmean(aux, axis_name='batch')\n    metrics = {'loss': aux[0], 'recall@20': aux[1], 'recall@1000': aux[2]}\n    \n    dropout_rng, new_dropout_rng = jax.random.split(dropout_rng)\n    return new_state, metrics, new_dropout_rng\n    \n\nbatch_size = HP.per_device_train_batch_size * jax.device_count()\ntrain_steps_per_epoch = len(train_dataset)//batch_size\ntotal_train_steps = train_steps_per_epoch * HP.num_train_epochs\n\nrng = jax.random.PRNGKey(HP.random_state)\ndropout_rng = jax.random.split(rng, jax.device_count())\n\n\nfor epoch in range(HP.num_train_epochs):\n    train_dataloader = get_dataloader(train_dataset, batch_size)\n    running_metrics = defaultdict(int)\n    steps_progress_bar = tqdm(enumerate(train_dataloader), total=train_steps_per_epoch, desc=f'Epoch #{epoch+1}/{HP.num_train_epochs}')\n    \n    for step, batch in steps_progress_bar:\n        batch = collate_fn(batch)\n        state, step_metrics, dropout_rng = train_step(state, dropout_rng, batch)\n        \n        # Update progress bar and running metrics\n        step_metrics = flax.jax_utils.unreplicate(step_metrics)\n        steps_progress_bar.set_postfix(**step_metrics)\n        for k, v in step_metrics.items():\n            running_metrics[k] += step_metrics[k]\n        \n        # Log metrics every `logging_freq` steps\n        if (step + 1) % HP.logging_freq == 0:\n            global_step = flax.jax_utils.unreplicate(state.step)\n            print('-'*50)\n            print(f\"Step {global_step-HP.logging_freq}-{global_step} out of {total_train_steps}\")\n            for k, v in running_metrics.items():\n                print(colored(k, 'blue'), ':', colored(v / HP.logging_freq, 'red'))\n            running_metrics = defaultdict(int)\n            print()\n\n    print(colored(f'Epoch #{epoch+1} completed.'))\n    print(colored('-'*100, 'red'))\n    print('\\n\\n')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import flax.training\n\n# # def binary_loss_fn(logits, labels):\n# #     logits, labels = logits[:, 1:], labels[:, 1:]\n# #     sample_weight = labels.sum(axis=-1)\n# #     labels = labels / (labels.sum(axis=-1, keepdims=True)+1e-8)\n# #     loss = optax.softmax_cross_entropy(logits, labels) * sample_weight\n# #     return loss\n\n# def get_recall(logits, labels, k):\n#     logits, labels = logits[:, 1:], labels[:, 1:]\n#     top_pred_ids = jax.lax.approx_max_k(logits, k, reduction_dimension=-1)[1]\n#     a = jax.nn.one_hot(top_pred_ids, logits.shape[-1]).sum(axis=1)\n#     true_preds = jnp.where((a==1.)&(labels>0), 1., 0.).sum(axis=-1)\n#     ground_preds = jnp.clip(labels.sum(axis=1), 1e-8, 20)\n#     return (true_preds / ground_preds).mean()\n    \n# @partial(jax.pmap, axis_name='batch')\n# def train_step(state, dropout_rng, batch):\n    \n#     def loss_fn(params):\n#         model_outputs = state.apply_fn(\n#             params,\n#             item_ids=batch['item_ids'],\n#             interaction_type_ids=batch['interaction_type_ids'],\n#             attention_mask=batch['attention_mask'],\n#             timestamp_features=batch['timestamp_features'],\n#             deterministic=False,\n#             rngs={'dropout': dropout_rng},\n#         )\n#         click_logits, cart_logits, order_logits = model_outputs\n        \n#         # Click Loss\n#         click_labels = jax.nn.one_hot(batch['target_click_ids'], click_logits.shape[-1]).sum(axis=1)\n#         click_mask = jnp.where(batch['target_click_ids']==0, 0., 1.)\n#         click_loss = binary_loss_fn(click_logits, click_labels).sum()\n        \n#         # Cart Loss\n#         cart_labels = jax.nn.one_hot(batch['target_cart_ids'], cart_logits.shape[-1]).sum(axis=1)\n#         cart_mask = jnp.where(batch['target_cart_ids']==0, 0., 1.)\n#         cart_loss = binary_loss_fn(cart_logits, cart_labels).sum()\n        \n#         # Order Loss\n#         order_labels = jax.nn.one_hot(batch['target_order_ids'], order_logits.shape[-1]).sum(axis=1)\n#         order_mask = jnp.where(batch['target_order_ids']==0, 0., 1.)\n#         order_loss = binary_loss_fn(order_logits, order_labels).sum()\n        \n#         total_loss = (click_loss + cart_loss + order_loss) / 3.\n        \n#         click_recall = get_recall(click_logits, click_labels, 20)\n#         cart_recall = get_recall(cart_logits, click_labels, 20)\n#         order_recall = get_recall(order_logits, order_labels, 20)\n#         click_recall_100 = get_recall(click_logits, click_labels, 100)\n        \n#         return total_loss, (click_loss, cart_loss, order_loss, click_recall, cart_recall, order_recall, click_recall_100)\n    \n#     grad_fn = jax.value_and_grad(loss_fn, has_aux=True)\n#     (total_loss, aux), grad = grad_fn(state.params)\n#     grad = jax.lax.pmean(grad, axis_name='batch')\n#     new_state = state.apply_gradients(grads=grad)\n    \n#     aux = jax.lax.pmean(aux, axis_name='batch')\n#     metrics = {\n#         'click_loss': aux[0], 'cart_loss': aux[1], 'order_loss': aux[2], \n#         'click@20': aux[3], 'cart@20': aux[4], 'order@20': aux[5], 'click@100': aux[6],\n#         'score': (0.10*aux[3]+0.30*aux[4]+0.60*aux[5]),\n#     }\n    \n#     dropout_rng, new_dropout_rng = jax.random.split(dropout_rng)\n#     return new_state, metrics, new_dropout_rng\n    \n\n# batch_size = HP.per_device_train_batch_size * jax.device_count()\n# train_steps_per_epoch = len(train_dataset)//batch_size\n# total_train_steps = train_steps_per_epoch * HP.num_train_epochs\n\n# rng = jax.random.PRNGKey(0)\n# dropout_rng = jax.random.split(rng, jax.device_count())\n\n\n# for epoch in range(HP.num_train_epochs):\n#     train_dataloader = get_dataloader(train_dataset, batch_size)\n#     running_metrics = defaultdict(int)\n#     steps_progress_bar = tqdm(enumerate(train_dataloader), total=train_steps_per_epoch, desc=f'Epoch #{epoch+1}/{HP.num_train_epochs}')\n    \n#     for step, batch in steps_progress_bar:\n#         batch = collate_fn(batch)\n#         state, step_metrics, dropout_rng = train_step(state, dropout_rng, batch)\n        \n#         # Update progress bar and running metrics\n#         step_metrics = flax.jax_utils.unreplicate(step_metrics)\n#         steps_progress_bar.set_postfix(**step_metrics)\n#         for k, v in step_metrics.items():\n#             running_metrics[k] += step_metrics[k]\n        \n#         # Log metrics every `logging_freq` steps\n#         if (step + 1) % HP.logging_freq == 0:\n#             global_step = flax.jax_utils.unreplicate(state.step)\n#             print('-'*50)\n#             print(f\"Step {global_step-HP.logging_freq}-{global_step} out of {total_train_steps}\")\n#             for k, v in running_metrics.items():\n#                 print(colored(k, 'blue'), ':', colored(v / HP.logging_freq, 'red'))\n#             running_metrics = defaultdict(int)\n#             print()\n\n#     print(colored(f'Epoch #{epoch+1} completed.'))\n#     print(colored('-'*100, 'red'))\n#     print('\\n\\n')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 50 min for 6350 w. 2 layers, 1 hour w. 4 layers\n# 6 hours for 8 layers, 64 batch size.\n# 18 hours for 16 layers, 232 batch size.","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import flax.serialization\nimport flax.training.checkpoints\n\nweights_file = 'model.msgpack'\nprint(weights_file)\nif jax.process_index() == 0:\n    params = jax.device_get(flax.jax_utils.unreplicate(state.params))\n    with open(f'/kaggle/working/{weights_file}', 'wb') as f:\n        model_bytes = flax.serialization.to_bytes(params)\n        f.write(model_bytes)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.init(project='otto')\nwandb.save('/kaggle/working/model.msgpack')\nsleep(100)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import flax.training.checkpoints\n# flax.training.checkpoints.restore_checkpoint('/kaggle/working/model.msgpack', target=state.params)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🎯 Prediction\n\n<a name='prediction'>","metadata":{}},{"cell_type":"code","source":"def process_test_session(session_dict):\n    events = session_dict['events']\n    if len(events) > HP.max_seq_len:\n        events = events[-HP.max_seq_len:]\n    \n    item_ids = [event['aid']+1 for event in events]\n    interaction_type_ids = [{'clicks': 0, 'carts': 1, 'orders': 2}[event['type']] for event in events]\n    pad_width = HP.max_seq_len-len(item_ids)\n    attention_mask = [1]*len(item_ids) + [0]*pad_width\n    item_ids = item_ids + [0]*pad_width\n    interaction_type_ids = interaction_type_ids + [0]*pad_width\n    assert len(item_ids) == len(interaction_type_ids) == len(attention_mask) == HP.max_seq_len\n    \n    timestamp_features = build_timestamp_features([event['ts'] for event in events])\n    timestamp_features = {k: v+[0]*pad_width for k, v in timestamp_features.items()}\n    \n    model_inputs = {\n        'item_ids': np.array(item_ids, dtype=np.int32),\n        'interaction_type_ids': np.array(interaction_type_ids, dtype=np.int32),\n        #'attention_mask': np.array(attention_mask, dtype=np.int32),\n    }\n    \n    return {**model_inputs, **timestamp_features}\n\ndef test_collate_fn(batch): \n    jnp_batch = {\n        'item_ids': jnp.array(batch['item_ids']),\n        'interaction_type_ids': jnp.array(batch['interaction_type_ids']),\n        #'attention_mask': jnp.array(batch['attention_mask']),\n        'timestamp_features': {\n            'ts_days': jnp.array(batch['ts_days']),\n            'ts_weekdays': jnp.array(batch['ts_weekdays']),\n            'ts_hours': jnp.array(batch['ts_hours']),\n            'ts_abs': jnp.array(batch['ts_abs']),\n            'ts_session': jnp.array(batch['ts_session']),\n        },\n    }\n    jnp_batch = flax.training.common_utils.shard(jnp_batch)\n    return jnp_batch\n\nDATA_FILE = '/kaggle/input/otto-recommender-system/test.jsonl'\nraw_test_dataset = datasets.load_dataset('json', data_files=DATA_FILE, split='train')\n\n# test_dataset = raw_test_dataset.map(\n#     process_test_session,\n#     desc='Processing raw dataset',\n#     remove_columns=['session', 'events']\n# )\n# test_dataset.set_format(type='numpy')\n# test_dataset.save_to_disk('test_dataset')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = datasets.load_from_disk('test_dataset')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@partial(jax.pmap, axis_name='batch')\ndef predict_step(state, batch):\n    model_outputs = state.apply_fn(\n        params,\n        item_ids=batch['item_ids'],\n        interaction_type_ids=batch['interaction_type_ids'],\n        #attention_mask=batch['attention_mask'],\n        timestamp_features=batch['timestamp_features'],\n        deterministic=True,\n    )\n    logits = model_outputs\n    logits = logits[:, 1:]\n    print('logits:', logits.shape)\n    \n    top_click_preds = jax.lax.approx_max_k(logits, 20, reduction_dimension=-1, recall_target=0.99)[1]\n    print('top_click_preds:', top_click_preds.shape)\n#     top_cart_preds = jax.lax.approx_max_k(cart_logits, 20, reduction_dimension=-1)[1]-1\n#     top_order_preds = jax.lax.approx_max_k(order_logits, 20, reduction_dimension=-1)[1]-1\n    \n    top_cart_preds, top_order_preds = top_click_preds, top_click_preds\n    return top_click_preds, top_cart_preds, top_order_preds\n\nbatch_size = HP.per_device_eval_batch_size * jax.device_count()\ntest_steps = len(test_dataset)//batch_size\ntest_dataloader = get_dataloader(test_dataset, batch_size)\nsteps_progress_bar = tqdm(enumerate(test_dataloader), total=test_steps)\n\nall_click_preds, all_cart_preds, all_order_preds = [], [], []\nfor step, batch in steps_progress_bar:\n    batch = test_collate_fn(batch)\n    top_click_preds, top_cart_preds, top_order_preds = predict_step(state, batch)\n    all_click_preds.append(np.concatenate(top_click_preds))\n    all_cart_preds.append(np.concatenate(top_cart_preds))\n    all_order_preds.append(np.concatenate(top_order_preds))\n    \nall_click_preds = np.concatenate(all_click_preds)\nall_cart_preds = np.concatenate(all_cart_preds)\nall_order_preds = np.concatenate(all_order_preds)\n\nprint('all_click_preds:', all_click_preds.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred = {'session_type': [], 'labels': []}\n\nfor i, session_dict in tqdm(enumerate(raw_test_dataset), total=len(raw_test_dataset)):\n    session_id = session_dict['session']\n    df_pred['session_type'].append(f'{session_id}_clicks')\n    df_pred['session_type'].append(f'{session_id}_carts')\n    df_pred['session_type'].append(f'{session_id}_orders')\n    \n    if i >= len(all_click_preds):\n        df_pred['labels'] += ['', '', '']\n        continue\n    \n    df_pred['labels'].append(' '.join([str(x) for x in all_click_preds[i].tolist()]))\n    df_pred['labels'].append(' '.join([str(x) for x in all_cart_preds[i].tolist()]))\n    df_pred['labels'].append(' '.join([str(x) for x in all_order_preds[i].tolist()]))\n    \ndf_pred = pd.DataFrame(df_pred)\ndf_pred","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred.labels.value_counts()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"1028055 / len(df_pred)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}