{"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":"import os\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport scipy.sparse as sp\n\nimport pickle\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import KFold\nfrom collections import Counter","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-09T11:27:50.554725Z","iopub.execute_input":"2023-01-09T11:27:50.555189Z","iopub.status.idle":"2023-01-09T11:27:50.562312Z","shell.execute_reply.started":"2023-01-09T11:27:50.555139Z","shell.execute_reply":"2023-01-09T11:27:50.560925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:27:50.568310Z","iopub.execute_input":"2023-01-09T11:27:50.568751Z","iopub.status.idle":"2023-01-09T11:27:50.576278Z","shell.execute_reply.started":"2023-01-09T11:27:50.568710Z","shell.execute_reply":"2023-01-09T11:27:50.575055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# load data","metadata":{}},{"cell_type":"code","source":"def load_pickle_object(path):\n    with open(path, \"rb\") as file:\n        obj = pickle.load(file)\n    return obj","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:27:50.591123Z","iopub.execute_input":"2023-01-09T11:27:50.592454Z","iopub.status.idle":"2023-01-09T11:27:50.597706Z","shell.execute_reply.started":"2023-01-09T11:27:50.592408Z","shell.execute_reply":"2023-01-09T11:27:50.596565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aid2idx = load_pickle_object(\"/kaggle/input/otto-sequence-dataset/aid2idx.pkl\")\nmat =  sp.load_npz(\"/kaggle/input/otto-sequence-dataset/dataset.npz\")\nprint(\"Number of article:\", len(aid2idx))\nprint(mat.shape)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:27:50.602560Z","iopub.execute_input":"2023-01-09T11:27:50.603108Z","iopub.status.idle":"2023-01-09T11:28:04.093424Z","shell.execute_reply.started":"2023-01-09T11:27:50.603075Z","shell.execute_reply":"2023-01-09T11:28:04.092281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mat = mat[:7000000, :]\nprint(mat.shape)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:51:43.067881Z","iopub.execute_input":"2023-01-09T11:51:43.068369Z","iopub.status.idle":"2023-01-09T11:51:43.115700Z","shell.execute_reply.started":"2023-01-09T11:51:43.068302Z","shell.execute_reply":"2023-01-09T11:51:43.114501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 1024\nEMBEDD_DIM = 64\nN_EPOCHS = 8\nMAX_LEN  = 30","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:03:53.719582Z","iopub.execute_input":"2023-01-09T11:03:53.719914Z","iopub.status.idle":"2023-01-09T11:03:53.725835Z","shell.execute_reply.started":"2023-01-09T11:03:53.719884Z","shell.execute_reply":"2023-01-09T11:03:53.724379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class OttoDataset(torch.utils.data.Dataset):\n    def __init__(self, indices, M):\n        self.M=M\n        self.indices = indices\n    \n    def padsequence(self, X, slen):\n        diff_len = MAX_LEN - slen\n        X = np.concatenate([np.zeros(diff_len), X])\n        return X\n    \n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        row = self.M[idx]\n        \n        X = row.data[:-1][:MAX_LEN]\n        y = row.data[-1]\n        \n        slen = len(X)\n        X = self.padsequence(X, slen)\n        \n        X = torch.tensor(X, dtype=torch.long)\n        y = torch.tensor(y, dtype=torch.long)\n        \n        key_mask = torch.ones(MAX_LEN, dtype=torch.long)\n        key_mask[-slen:] = 0\n        \n        return (X, y, key_mask, slen)\n    \n    def __len__(self):\n        return len(self.indices)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:19:42.682839Z","iopub.execute_input":"2023-01-09T11:19:42.683597Z","iopub.status.idle":"2023-01-09T11:19:42.694544Z","shell.execute_reply.started":"2023-01-09T11:19:42.683561Z","shell.execute_reply":"2023-01-09T11:19:42.693200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class OttoGruModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.embedding = nn.Embedding(1+len(aid2idx), EMBEDD_DIM)\n        self.mha = nn.MultiheadAttention(\n            EMBEDD_DIM,\n            4,\n            dropout=0.1,\n            batch_first=True\n        )\n        \n        self.out = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(128, 1)\n        )\n    \n    def get_output(self, h):\n        #h = self.mlp(h)\n        #if len(h.shape) == 3:\n        #    h = h.permute(0, 2, 1)\n        #    h = self.bn(h)\n        #    h = h.permute(0, 2, 1)\n        #else:\n        #    h = self.bn(h)\n        y = self.out(h)\n        return y\n        \n    def forward(self, x, xneg, y, mask):\n        x = self.embedding(x)\n        xneg_embedd = self.embedding(xneg)\n        yembedd = self.embedding(y)\n        \n        \n        h, _ = self.mha(x, x, x, key_padding_mask=mask)\n        h = h[:, -1, :]\n        \n        batch_size = len(h)\n        neg_batch_size = xneg_embedd.shape[1]\n        hpos = torch.cat([h, yembedd], dim=-1)\n        hneg = torch.cat([h.unsqueeze(1).expand(-1, neg_batch_size, -1), \n                          xneg_embedd.expand(batch_size, -1, -1)], dim=-1)\n        \n        ypos = self.get_output(hpos)\n        yneg = self.get_output(hneg)\n        yneg = yneg.squeeze(-1)\n        \n        return ypos, yneg","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:19:42.696709Z","iopub.execute_input":"2023-01-09T11:19:42.697050Z","iopub.status.idle":"2023-01-09T11:19:42.710931Z","shell.execute_reply.started":"2023-01-09T11:19:42.697020Z","shell.execute_reply":"2023-01-09T11:19:42.709637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Negative Samples","metadata":{}},{"cell_type":"code","source":"def negsamples(X, k=1000, val=True):\n    (X, freq) = X.unique(return_counts=True)\n    \n    freq = freq ** 0.75\n    total_count = freq.sum()\n    p = freq/total_count\n    \n    if val:\n        np.random.seed(182382)\n        k = 5000\n    \n    u = np.random.uniform(0, 1, len(X))\n    u = torch.tensor(u, device=device, dtype=torch.float32)\n    u = u/total_count\n    mask = (u<p)\n    X = X[(mask) & (X!=0)]\n    X = X[:k]\n    return X","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:19:42.713014Z","iopub.execute_input":"2023-01-09T11:19:42.713703Z","iopub.status.idle":"2023-01-09T11:19:42.728334Z","shell.execute_reply.started":"2023-01-09T11:19:42.713668Z","shell.execute_reply":"2023-01-09T11:19:42.726950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# evaluate","metadata":{}},{"cell_type":"code","source":"def evaluate(model, val_dataloader):\n    eval_20 = 0\n    eval_50 = 0\n    \n    model.eval()\n    for X, y, key_mask, slen in val_dataloader:\n        X = X.to(device)\n        key_mask = key_mask.to(device)\n        y = y.to(device)\n        Xneg = negsamples(X, k=5000)\n        Xneg = Xneg.view(1, -1)\n        \n        max_seqlen = max(slen)\n        X = X[:, -max_seqlen:]\n        \n        with torch.no_grad():\n            yhat_pos, yhat_neg = model(X, Xneg, y, key_mask)\n        \n        \n        ydiff = yhat_pos - yhat_neg\n        ydiff = (ydiff<0).type(torch.float32).sum(dim=-1)\n        \n        eval_20 += (ydiff<20).type(torch.float32).mean()\n        eval_50 += (ydiff<50).type(torch.float32).mean()\n    \n    eval_20 = eval_20.item()/len(val_dataloader)\n    eval_50 = eval_50.item()/len(val_dataloader)\n    return (eval_20, eval_50)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:24:12.973302Z","iopub.execute_input":"2023-01-09T11:24:12.973950Z","iopub.status.idle":"2023-01-09T11:24:12.984171Z","shell.execute_reply.started":"2023-01-09T11:24:12.973914Z","shell.execute_reply":"2023-01-09T11:24:12.982735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train epoch","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, train_dataloader, optimizer, schedular):\n    epoch_loss=[]\n    epoch_reg_loss=[]\n    for (X, y, key_mask, slen) in train_dataloader:\n        X = X.to(device)\n        key_mask = key_mask.to(device)\n        y = y.to(device)\n        Xneg = negsamples(X)\n        Xneg = Xneg.view(1, -1)\n        \n        max_seqlen = max(slen)\n        X = X[:, -max_seqlen:]\n        \n        model.train()\n        yhat_pos, yhat_neg = model(X, Xneg, y, key_mask)\n        \n        reg_loss = (torch.square(yhat_pos).mean() + torch.square(yhat_neg).mean())/2\n        loss = -torch.log(1e-8 + torch.sigmoid( yhat_pos - yhat_neg ))\n        loss = loss + 1e-4*reg_loss\n        \n        loss = loss.mean()\n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        schedular.step()\n        \n        epoch_loss.append(loss.item())\n        epoch_reg_loss.append(reg_loss.item())\n    return np.mean(epoch_loss), np.mean(epoch_reg_loss)","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:24:12.986339Z","iopub.execute_input":"2023-01-09T11:24:12.986844Z","iopub.status.idle":"2023-01-09T11:24:13.004216Z","shell.execute_reply.started":"2023-01-09T11:24:12.986775Z","shell.execute_reply":"2023-01-09T11:24:13.002876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train model","metadata":{}},{"cell_type":"code","source":"def train_model(foldnum, train_dataloader, val_dataloader):\n    model = OttoGruModel().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-2, weight_decay=0.01)\n    schedular = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, \n                                                           N_EPOCHS * len(train_dataloader),\n                                                           eta_min = 1e-5\n                                                          )\n    \n    \n    #schedular = torch.optim.lr_scheduler.OneCycleLR(optimizer,\n    #                                                max_lr=0.01, \n    #                                                steps_per_epoch=len(train_dataloader), \n    #                                                epochs=N_EPOCHS\n    #                                               )\n    \n                                                           \n    best_20=None; best_50=None\n    epoch_losses = []\n    \n    for e in range(N_EPOCHS):\n        epoch_loss, epoch_reg_loss = train_epoch(model, train_dataloader, optimizer, schedular)\n        (eval_20, eval_50) = evaluate(model, val_dataloader)\n        \n        if (best_20 is None) or (eval_20>best_20):\n            best_20=eval_20\n            torch.save(model, \"model_{}.pt\".format(foldnum))\n        \n        if (best_50 is None) or (eval_50>best_50):\n            best_50=eval_50\n            torch.save(model, \"model_{}.pt\".format(foldnum))\n        \n        epoch_losses.append(epoch_loss)\n        print(\"Epoch:{} | Train loss:{:.4f}\".format(e, epoch_loss))\n        print(\"epoch reg loss:{:.4f}\".format(epoch_reg_loss))\n        print(\"eval_20:{:.4f} | eval_50:{:.4f}\".format(eval_20, eval_50))\n        print(\"best_20:{:.4f} | best_50:{:.4f}\".format(best_20, best_50))\n        print()\n        print()\n        \n    \n    plt.title(\"train losses\")\n    plt.plot(epoch_losses)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:24:13.006463Z","iopub.execute_input":"2023-01-09T11:24:13.007251Z","iopub.status.idle":"2023-01-09T11:24:13.021041Z","shell.execute_reply.started":"2023-01-09T11:24:13.007192Z","shell.execute_reply":"2023-01-09T11:24:13.020080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = KFold(n_splits=5, random_state=444, shuffle=True)\nfor foldnum, (train_index, val_index) in enumerate(kf.split(np.arange( mat.shape[0] ))):\n    print(train_index.shape, val_index.shape)\n    \n    train_dataset = OttoDataset(train_index, mat)\n    train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\n    \n    val_dataset = OttoDataset(val_index, mat)\n    val_dataloader = torch.utils.data.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, drop_last=True)\n    \n    print(len(train_dataloader), len(val_dataloader))\n    train_model(foldnum, train_dataloader, val_dataloader)\n    break","metadata":{"execution":{"iopub.status.busy":"2023-01-09T11:24:13.023028Z","iopub.execute_input":"2023-01-09T11:24:13.023458Z","iopub.status.idle":"2023-01-09T11:26:01.360989Z","shell.execute_reply.started":"2023-01-09T11:24:13.023425Z","shell.execute_reply":"2023-01-09T11:26:01.359700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}