{"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":"# Pretraining with Merlin's Transformers4Rec 🧙\n\n> Transformers4Rec is a powerful library to leverage transformers for session-based recommender systems.\n\nLink : https://github.com/NVIDIA-Merlin/Transformers4Rec","metadata":{}},{"cell_type":"markdown","source":"## 0. Initialization","metadata":{}},{"cell_type":"code","source":"try:\n    import transformers4rec\nexcept:\n    print(\"Install packages\\n\\n\")\n    !pip install transformers4rec[pytorch,nvtabular]\n    !pip install -U nvtabular==1.3.3\n    !pip install beartype\n    !pip install -U pytorch_lightning","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-11-09T17:26:27.040957Z","iopub.execute_input":"2022-11-09T17:26:27.041937Z","iopub.status.idle":"2022-11-09T17:26:27.055189Z","shell.execute_reply.started":"2022-11-09T17:26:27.041851Z","shell.execute_reply":"2022-11-09T17:26:27.054483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport ast\nimport json\nimport glob\nimport torch\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport nvtabular as nvt\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom datetime import datetime\nfrom collections import Counter\n\nwarnings.simplefilter(action='ignore', category=FutureWarning)\nwarnings.simplefilter(action='ignore', category=UserWarning)\n\nfrom nvtabular.ops import *\nfrom merlin.schema.tags import Tags\nfrom merlin_standard_lib import Schema\nfrom transformers4rec import torch as tr\nfrom transformers4rec.torch import Trainer\nfrom transformers4rec.torch.ranking_metric import RecallAt\nfrom nvtabular.loader.torch import TorchAsyncItr, DLDataLoader\nfrom transformers4rec.config.trainer import T4RecTrainingArguments","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-11-09T17:26:27.056421Z","iopub.execute_input":"2022-11-09T17:26:27.057414Z","iopub.status.idle":"2022-11-09T17:26:34.085098Z","shell.execute_reply.started":"2022-11-09T17:26:27.057377Z","shell.execute_reply":"2022-11-09T17:26:34.084141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create a validation set\n\n- Using the code provided by the hosts : https://github.com/otto-de/recsys-dataset\n- Strategy is the same asthe other shared by @radek1 [here](https://www.kaggle.com/competitions/otto-recommender-system/discussion/364991) : Use the last week for val!\n- You do not need to rerun this part, precomputed jsons are available here : https://www.kaggle.com/datasets/theoviel/split-otto  (please upvote !)","metadata":{}},{"cell_type":"code","source":"import json\nimport random\nimport argparse\nimport pandas as pd\n\nfrom tqdm import tqdm\nfrom typing import List\nfrom pathlib import Path\nfrom copy import deepcopy\nfrom beartype import beartype\nfrom pandas.io.json._json import JsonReader\n\n\nclass setEncoder(json.JSONEncoder):\n    def default(self, obj):\n        return list(obj)\n\n@beartype\ndef ground_truth(events: List[dict]):\n    prev_labels = {\"clicks\": None, \"carts\": set(), \"orders\": set()}\n\n    for event in reversed(events):\n        event[\"labels\"] = {}\n\n        for label in ['clicks', 'carts', 'orders']:\n            if prev_labels[label]:\n                if label != 'clicks':\n                    event[\"labels\"][label] = prev_labels[label].copy()\n                else:\n                    event[\"labels\"][label] = prev_labels[label]\n\n        if event[\"type\"] == \"clicks\":\n            prev_labels['clicks'] = event[\"aid\"]\n        if event[\"type\"] == \"carts\":\n            prev_labels['carts'].add(event[\"aid\"])\n        elif event[\"type\"] == \"orders\":\n            prev_labels['orders'].add(event[\"aid\"])\n\n    return events[:-1]\n\n\n@beartype\ndef split_events(events: List[dict], split_idx=None):\n    test_events = ground_truth(deepcopy(events))\n    if not split_idx:\n        split_idx = random.randint(1, len(test_events))\n    test_events = test_events[:split_idx]\n    labels = test_events[-1]['labels']\n    for event in test_events:\n        del event['labels']\n    return test_events, labels\n\n\n@beartype\ndef create_kaggle_testset(sessions: pd.DataFrame, sessions_output: Path, labels_output: Path):\n    last_labels = []\n    splitted_sessions = []\n\n    for _, session in tqdm(sessions.iterrows(), desc=\"Creating trimmed testset\", total=len(sessions)):\n        session = session.to_dict()\n        splitted_events, labels = split_events(session['events'])\n        last_labels.append({'session': session['session'], 'labels': labels})\n        splitted_sessions.append({'session': session['session'], 'events': splitted_events})\n\n    with open(sessions_output, 'w') as f:\n        for session in splitted_sessions:\n            f.write(json.dumps(session) + '\\n')\n\n    with open(labels_output, 'w') as f:\n        for label in last_labels:\n            f.write(json.dumps(label, cls=setEncoder) + '\\n')\n\n\n@beartype\ndef trim_session(session: dict, max_ts: int) -> dict:\n    session['events'] = [event for event in session['events'] if event['ts'] < max_ts]\n    return session\n\n\n@beartype\ndef get_max_ts(sessions_file: Path) -> int:\n    max_ts = float('-inf')\n    with open(sessions_file) as f:\n        for line in tqdm(f, desc=\"Finding max timestamp\"):\n            session = json.loads(line)\n            max_ts = max(max_ts, session['events'][-1]['ts'])\n    return max_ts\n\n\n@beartype\ndef train_test_split(session_chunks: JsonReader, train_file: Path, test_file: Path, max_ts: int, test_days: int):\n    split_millis = test_days * 24 * 60 * 60 * 1000\n    split_ts = max_ts - split_millis\n    Path(train_file).parent.mkdir(parents=True, exist_ok=True)\n    train_file = open(train_file, \"w\")\n    Path(test_file).parent.mkdir(parents=True, exist_ok=True)\n    test_file = open(test_file, \"w\")\n    for chunk in tqdm(session_chunks, desc=\"Splitting sessions\"):\n        for _, session in chunk.iterrows():\n            session = session.to_dict()\n            if session['events'][0]['ts'] > split_ts:\n                test_file.write(json.dumps(session, cls=setEncoder) + \"\\n\")\n            else:\n                session = trim_session(session, split_ts)\n                train_file.write(json.dumps(session, cls=setEncoder) + \"\\n\")\n    train_file.close()\n    test_file.close()\n\n\n@beartype\ndef main(train_set: Path, output_path: Path, days: int, seed: int):\n    random.seed(seed)\n    max_ts = get_max_ts(train_set)\n\n    session_chunks = pd.read_json(train_set, lines=True, chunksize=100000)\n    train_file = output_path / 'train_sessions.jsonl'\n    test_file_full = output_path / 'test_sessions_full.jsonl'\n    train_test_split(session_chunks, train_file, test_file_full, max_ts, days)\n\n    test_sessions = pd.read_json(test_file_full, lines=True)\n    test_sessions_file = output_path / 'test_sessions.jsonl'\n    test_labels_file = output_path / 'test_labels.jsonl'\n    create_kaggle_testset(test_sessions, test_sessions_file, test_labels_file)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-09T17:26:34.086521Z","iopub.execute_input":"2022-11-09T17:26:34.087337Z","iopub.status.idle":"2022-11-09T17:26:34.208589Z","shell.execute_reply.started":"2022-11-09T17:26:34.087297Z","shell.execute_reply":"2022-11-09T17:26:34.207618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@beartype\ndef val_split(train_set: Path, output_path: Path, days: int, seed: int):\n    random.seed(seed)\n    max_ts = get_max_ts(train_set)\n\n    session_chunks = pd.read_json(train_set, lines=True, chunksize=100000)\n    train_file = output_path / 'train_sessions.jsonl'\n    test_file_full = output_path / 'val_sessions.jsonl'\n    train_test_split(session_chunks, train_file, test_file_full, max_ts, days)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.211711Z","iopub.execute_input":"2022-11-09T17:26:34.212434Z","iopub.status.idle":"2022-11-09T17:26:34.221597Z","shell.execute_reply.started":"2022-11-09T17:26:34.212391Z","shell.execute_reply":"2022-11-09T17:26:34.220456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# val_split(\n#     Path('../input/otto-recommender-system/train.jsonl'),\n#     Path(\"\"),\n#     7,\n#     42\n# )","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.225138Z","iopub.execute_input":"2022-11-09T17:26:34.225425Z","iopub.status.idle":"2022-11-09T17:26:34.230123Z","shell.execute_reply.started":"2022-11-09T17:26:34.225399Z","shell.execute_reply":"2022-11-09T17:26:34.229034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Prepare the Data\n- I convert the jsonl to parquet and already group by sessions.\n- You do not need to rerun this part, precomputed parquets are available here : https://www.kaggle.com/datasets/theoviel/otto-parquet/  (please upvote !)","metadata":{}},{"cell_type":"code","source":"CLASSES = ['', 'clicks', 'carts', 'orders']\n\n\ndef jsonl_to_df(fn, total=1290, test=False):\n    \n    chunks = pd.read_json(fn, lines=True, chunksize=10000)\n\n    sessions, aids, tss, types, labels_clicks, labels_carts, labels_orders = [], [], [], [], [], [], []\n    for chunk in tqdm(chunks, total=total):\n        for row_idx, session_data in chunk.iterrows():\n            aids_, tss_, types_, labels_clicks_, labels_carts_, labels_orders_ = [], [], [], [], [], []\n            if len(session_data['events']) > 1 and not test:\n                events = ground_truth(session_data.events)\n                for event in events:\n                    aids_.append(event['aid'])\n                    tss_.append(event['ts'])\n                    types_.append(event['type'])\n                    labels_clicks_.append(event['labels'].get(\"clicks\", None))\n                    labels_carts_.append(list(event['labels'].get(\"carts\", [])))\n                    labels_orders_.append(list(event['labels'].get(\"orders\", [])))\n            else:\n                for event in session_data.events:\n                    aids_.append(event['aid'])\n                    tss_.append(event['ts'])\n                    types_.append(event['type'])\n\n            sessions.append(session_data.session)\n            aids.append(aids_)\n            tss.append(tss_)\n            types.append(types_)\n            labels_clicks.append(labels_clicks_)\n            labels_carts.append(labels_carts_)\n            labels_orders.append(labels_orders_)\n\n    df = pd.DataFrame(data={\n        'session': sessions,\n        'aid': aids,\n        'ts': tss,\n        'type': types,\n        \"labels_clicks\": labels_clicks,\n        \"labels_carts\": labels_carts,\n        \"labels_orders\": labels_orders,\n    })\n    df['target'] = df['type'].apply(lambda x: [CLASSES.index(c) for c in x])\n    \n    return df\n\n\ndef jsonl_to_df_train(fn, total=1290, test=False):\n    \"\"\"\n    This function processes the data in chunks to avoid OOM.\n    \"\"\"\n    chunks = pd.read_json(fn, lines=True, chunksize=200000)\n\n    for i, chunk in tqdm(enumerate(chunks), total=total):\n        sessions, aids, tss, types, labels_clicks, labels_carts, labels_orders = [], [], [], [], [], [], []\n\n        for row_idx, session_data in chunk.iterrows():\n            aids_, tss_, types_, labels_clicks_, labels_carts_, labels_orders_ = [], [], [], [], [], []\n            if len(session_data['events']) > 1 and not test:\n                events = ground_truth(session_data.events)\n                for event in events:\n                    aids_.append(event['aid'])\n                    tss_.append(event['ts'])\n                    types_.append(event['type'])\n                    labels_clicks_.append(event['labels'].get(\"clicks\", None))\n                    labels_carts_.append(list(event['labels'].get(\"carts\", [])))\n                    labels_orders_.append(list(event['labels'].get(\"orders\", [])))\n            else:\n                for event in session_data.events:\n                    aids_.append(event['aid'])\n                    tss_.append(event['ts'])\n                    types_.append(event['type'])\n\n            sessions.append(session_data.session)\n            aids.append(aids_)\n            tss.append(tss_)\n            types.append(types_)\n            labels_clicks.append(labels_clicks_)\n            labels_carts.append(labels_carts_)\n            labels_orders.append(labels_orders_)\n\n        df = pd.DataFrame(data={\n            'session': sessions,\n            'aid': aids,\n            'ts': tss,\n            'type': types,\n            \"labels_clicks\": labels_clicks,\n            \"labels_carts\": labels_carts,\n            \"labels_orders\": labels_orders,\n        })\n        df['target'] = df['type'].apply(lambda x: [CLASSES.index(c) for c in x])\n        \n        df.to_parquet(f'../output/train_{i}.parquet', index=False)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-11-09T17:26:34.231932Z","iopub.execute_input":"2022-11-09T17:26:34.232801Z","iopub.status.idle":"2022-11-09T17:26:34.258365Z","shell.execute_reply.started":"2022-11-09T17:26:34.232758Z","shell.execute_reply":"2022-11-09T17:26:34.257473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_PATH = \"../input/split-otto/train_sessions.jsonl\"\nVAL_PATH = \"../input/split-otto/val_sessions.jsonl\"\nTEST_PATH = \"../input/otto-recommender-system/test.jsonl\"","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.260025Z","iopub.execute_input":"2022-11-09T17:26:34.260381Z","iopub.status.idle":"2022-11-09T17:26:34.271273Z","shell.execute_reply.started":"2022-11-09T17:26:34.260346Z","shell.execute_reply":"2022-11-09T17:26:34.270332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# jsonl_to_df_train(TRAIN_PATH, total=56)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.273390Z","iopub.execute_input":"2022-11-09T17:26:34.273736Z","iopub.status.idle":"2022-11-09T17:26:34.281158Z","shell.execute_reply.started":"2022-11-09T17:26:34.273709Z","shell.execute_reply":"2022-11-09T17:26:34.280242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# df_val = jsonl_to_df(VAL_PATH, total=181)\n# df_val.to_parquet('val.parquet', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.282551Z","iopub.execute_input":"2022-11-09T17:26:34.283085Z","iopub.status.idle":"2022-11-09T17:26:34.289849Z","shell.execute_reply.started":"2022-11-09T17:26:34.283049Z","shell.execute_reply":"2022-11-09T17:26:34.288987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%time\n# df_test = jsonl_to_df(TEST_PATH, total=168, test=True)\n# df_test.to_parquet('test.parquet', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.291646Z","iopub.execute_input":"2022-11-09T17:26:34.292060Z","iopub.status.idle":"2022-11-09T17:26:34.301533Z","shell.execute_reply.started":"2022-11-09T17:26:34.292000Z","shell.execute_reply":"2022-11-09T17:26:34.300603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Pretrain a Language Model using Transformers4Rec\n- Adapted from https://github.com/NVIDIA-Merlin/Transformers4Rec/blob/main/examples/getting-started-session-based/02-session-based-XLNet-with-PyT.ipynb","metadata":{}},{"cell_type":"markdown","source":"### Schema\n- The schema is usually generated by the NVTabular pipeline. In this case I define it manually since I only use two features & don't resort to NVT.","metadata":{}},{"cell_type":"code","source":"%%writefile schema.pb\n\nfeature {\n  name: \"aid\"\n  type: INT\n  int_domain {\n    name: \"aid\"\n    min: 0\n    max: 1855610 \n    is_categorical: true\n  }\n  annotation {\n    tag: \"item_id\"\n    tag: \"list\"\n    tag: \"categorical\"\n    tag: \"item\"\n  }\n  value_count {\n    min: 2\n    max: 500\n  }\n}\n\nfeature {\n  name: \"target\"\n  type: INT\n  int_domain {\n    name: \"target\"\n    min: 1\n    max: 3\n    is_categorical: true\n  }\n  annotation {\n    tag: \"list\"\n    tag: \"categorical\"\n    tag: \"item\"\n  }\n  value_count {\n    min: 2\n    max: 500\n  }\n}","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.303328Z","iopub.execute_input":"2022-11-09T17:26:34.303865Z","iopub.status.idle":"2022-11-09T17:26:34.314401Z","shell.execute_reply.started":"2022-11-09T17:26:34.303825Z","shell.execute_reply":"2022-11-09T17:26:34.313308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"schema = Schema().from_proto_text(\"schema.pb\")","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.316339Z","iopub.execute_input":"2022-11-09T17:26:34.316850Z","iopub.status.idle":"2022-11-09T17:26:34.340426Z","shell.execute_reply.started":"2022-11-09T17:26:34.316816Z","shell.execute_reply":"2022-11-09T17:26:34.339560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inputs\n- `TabularSequenceFeatures.from_schema` automatically creates the embedding layers from the schema","metadata":{}},{"cell_type":"code","source":"inputs = tr.TabularSequenceFeatures.from_schema(\n    schema,\n    max_sequence_length=500,\n    d_output=100,\n    masking=\"mlm\",\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:34.344755Z","iopub.execute_input":"2022-11-09T17:26:34.345062Z","iopub.status.idle":"2022-11-09T17:26:36.595861Z","shell.execute_reply.started":"2022-11-09T17:26:34.345035Z","shell.execute_reply":"2022-11-09T17:26:36.594746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:36.597417Z","iopub.execute_input":"2022-11-09T17:26:36.598072Z","iopub.status.idle":"2022-11-09T17:26:36.607425Z","shell.execute_reply.started":"2022-11-09T17:26:36.598031Z","shell.execute_reply":"2022-11-09T17:26:36.606306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model\n- Define the config\n- Build the transformer\n- Add the next item prediction head - This is not the task of the competition !\n","metadata":{}},{"cell_type":"code","source":"# Define XLNetConfig class and set default parameters for HF XLNet config  \ntransformer_config = tr.XLNetConfig.build(\n    d_model=64, n_head=4, n_layer=2, total_seq_length=500\n)\n\n# Define the model block including: inputs, masking, projection and transformer block.\nbody = tr.SequentialBlock(\n    inputs,\n    tr.MLPBlock([64]),\n    tr.TransformerBlock(transformer_config, masking=inputs.masking)\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:36.609417Z","iopub.execute_input":"2022-11-09T17:26:36.609956Z","iopub.status.idle":"2022-11-09T17:26:36.631159Z","shell.execute_reply.started":"2022-11-09T17:26:36.609773Z","shell.execute_reply":"2022-11-09T17:26:36.630284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Defines the evaluation top-N metrics and the cut-offs\nmetrics = [\n    RecallAt(top_ks=[10, 20], labels_onehot=True)\n]\n\n# Define a head related to next item prediction task \nhead = tr.Head(\n    body,\n    tr.NextItemPredictionTask(weight_tying=True, hf_format=True, metrics=metrics),\n    inputs=inputs,\n)\n\n# Get the end-to-end Model class \nmodel = tr.Model(head)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:36.633736Z","iopub.execute_input":"2022-11-09T17:26:36.634406Z","iopub.status.idle":"2022-11-09T17:26:36.646261Z","shell.execute_reply.started":"2022-11-09T17:26:36.634368Z","shell.execute_reply":"2022-11-09T17:26:36.645253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"DEBUG = True  # Set to false to train on the whole data. Warning, it's slow !\nTRAIN = True\n\nlog_folder = \"logs/\"","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:36.647959Z","iopub.execute_input":"2022-11-09T17:26:36.648314Z","iopub.status.idle":"2022-11-09T17:26:36.654348Z","shell.execute_reply.started":"2022-11-09T17:26:36.648278Z","shell.execute_reply":"2022-11-09T17:26:36.653229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set hyperparameters for training \n\ntrain_args = T4RecTrainingArguments(\n    data_loader_engine='nvtabular', \n    dataloader_drop_last=True,\n    gradient_accumulation_steps=1,\n    per_device_train_batch_size=128, \n    per_device_eval_batch_size=128,\n    output_dir=log_folder, \n    learning_rate=0.0005,\n    lr_scheduler_type='cosine', \n    learning_rate_num_cosine_cycles_by_epoch=1,\n    num_train_epochs=1,\n    max_sequence_length=500, \n    report_to=[],\n    logging_steps=500,\n    save_steps=1000,\n    no_cuda=False,\n)\n\ntrainer = Trainer(\n    model=model,\n    args=train_args,\n    schema=schema,\n    compute_metrics=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:36.655624Z","iopub.execute_input":"2022-11-09T17:26:36.656415Z","iopub.status.idle":"2022-11-09T17:26:39.689520Z","shell.execute_reply.started":"2022-11-09T17:26:36.656379Z","shell.execute_reply":"2022-11-09T17:26:39.688533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    train_paths = sorted(glob.glob(\"../input/otto-parquet/train_54.parquet\"))\n    eval_paths = sorted(glob.glob(\"../input/otto-parquet/train_55.parquet\")) \nelse:\n    train_paths = sorted(glob.glob(\"../input/otto-parquet/train_*.parquet\"))\n    eval_paths = sorted(glob.glob(\"../input/otto-parquet/val.parquet\"))\n\ntrainer.train_dataset_or_path = train_paths\ntrainer.eval_dataset_or_path = eval_paths\n\nprint(train_paths)\nprint(eval_paths)","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:39.690868Z","iopub.execute_input":"2022-11-09T17:26:39.691331Z","iopub.status.idle":"2022-11-09T17:26:39.698904Z","shell.execute_reply.started":"2022-11-09T17:26:39.691292Z","shell.execute_reply":"2022-11-09T17:26:39.697857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN:\n    trainer.train()\n    trainer._save_model_and_checkpoint(save_model_class=True)\nelse:\n    trainer.load_model_trainer_states_from_checkpoint('/workspace/logs/2022-11-08/6/checkpoint-86707')","metadata":{"execution":{"iopub.status.busy":"2022-11-09T17:26:39.700316Z","iopub.execute_input":"2022-11-09T17:26:39.701095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Evaluate\n- On the next item prediction task","metadata":{}},{"cell_type":"code","source":"train_metrics = trainer.evaluate(eval_dataset=eval_paths, metric_key_prefix='eval')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for key in sorted(train_metrics.keys()):\n    print(\" %s = %s\" % (key, str(train_metrics[key])))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Finetune on the Competiton Task\n\n- TODO !\n- Might be left as an exercise for the reader =)","metadata":{}},{"cell_type":"markdown","source":"#### Improvements\n\n- Pretrain the MLM on the test data\n- Move all the FE to NVTabular\n- Tweak the architecture\n- Improve the recall of the pretrained model\n\n*Thanks for reading !*","metadata":{}}]}