{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Riiid! SAINT+ Solution\n\nThis notebook is our solution of the \"Riiid! Answer Correctness Prediction\".\n\nAll codes are available on [GitHub](https://github.com/marisakamozz/riiid)."},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls /kaggle/input/riiid-saintp-solution/","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/riiid-saintp-solution/')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import torch\nimport pandas as pd\n\nfrom saintmodel import SaintModel, SaintLightningModule, SaintHistory\nfrom saintsubmit import load_saint_config, SaintPredictor, make_submission","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"args = load_saint_config()\nmodel = SaintModel(\n    seq_len=args.seq_len, n_dim=args.n_dim, std=args.std, dropout=args.dropout, nhead=args.nhead, n_layers=args.n_layers\n)\nmodule = SaintLightningModule(args, model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\nmodule.load_state_dict(torch.load('/kaggle/input/riiid-saintp-solution/saint.ckpt')['state_dict'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\nlast_history = pd.read_pickle('/kaggle/input/riiid-saintp-solution/last_history.pickle')\nlast_timestamp = pd.read_pickle('/kaggle/input/riiid-saintp-solution/last_timestamp.pickle')\nlast_user_count = pd.read_pickle('/kaggle/input/riiid-saintp-solution/last_user_count.pickle')\ndict_lag = pd.read_pickle('/kaggle/input/riiid-saintp-solution/dict_lag.pickle')\ndict_elapsed = pd.read_pickle('/kaggle/input/riiid-saintp-solution/dict_elapsed.pickle')\ndict_user_count = pd.read_pickle('/kaggle/input/riiid-saintp-solution/dict_user_count.pickle')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodule.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"saint_history = SaintHistory(\n    args, last_history, last_timestamp, last_user_count,\n    dict_lag, dict_user_count, dict_elapsed\n)\npredictor = SaintPredictor(module.model, saint_history)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import riiideducation\nenv = riiideducation.make_env()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"make_submission(env, predictor)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}