{"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-08T16:34:24.687621Z","iopub.execute_input":"2023-01-08T16:34:24.688295Z","iopub.status.idle":"2023-01-08T16:34:24.696336Z","shell.execute_reply.started":"2023-01-08T16:34:24.688264Z","shell.execute_reply":"2023-01-08T16:34:24.694195Z"},"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-08T16:34:24.721799Z","iopub.execute_input":"2023-01-08T16:34:24.722238Z","iopub.status.idle":"2023-01-08T16:34:24.729182Z","shell.execute_reply.started":"2023-01-08T16:34:24.722208Z","shell.execute_reply":"2023-01-08T16:34:24.727340Z"},"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-08T16:34:24.755676Z","iopub.execute_input":"2023-01-08T16:34:24.756638Z","iopub.status.idle":"2023-01-08T16:34:24.761095Z","shell.execute_reply.started":"2023-01-08T16:34:24.756605Z","shell.execute_reply":"2023-01-08T16:34:24.760378Z"},"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-08T16:34:24.915499Z","iopub.execute_input":"2023-01-08T16:34:24.915910Z","iopub.status.idle":"2023-01-08T16:34:34.683583Z","shell.execute_reply.started":"2023-01-08T16:34:24.915863Z","shell.execute_reply":"2023-01-08T16:34:34.682795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 1024\nEMBEDD_DIM = 64\nN_EPOCHS = 3\nMAX_LEN  = 30\nMAX_NEG_SAMPLES=30","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.685249Z","iopub.execute_input":"2023-01-08T16:34:34.686647Z","iopub.status.idle":"2023-01-08T16:34:34.691381Z","shell.execute_reply.started":"2023-01-08T16:34:34.686594Z","shell.execute_reply":"2023-01-08T16:34:34.690426Z"},"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([X, np.zeros(diff_len)])\n        return X\n    \n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        row = self.M[idx]\n        X = row.data[-MAX_LEN:]\n        \n        slen = len(X)\n        X = self.padsequence(X, slen)\n        X = torch.tensor(X, dtype=torch.long)\n        mask = torch.zeros(MAX_LEN, dtype=torch.long)\n        mask[:slen] = 1\n        \n        return (X, mask, slen)\n    \n    def __len__(self):\n        return len(self.indices)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.692763Z","iopub.execute_input":"2023-01-08T16:34:34.693925Z","iopub.status.idle":"2023-01-08T16:34:34.712393Z","shell.execute_reply.started":"2023-01-08T16:34:34.693858Z","shell.execute_reply":"2023-01-08T16:34:34.710548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model","metadata":{}},{"cell_type":"code","source":"class OttoItem2VecModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.embedding = nn.Embedding(1+len(aid2idx), EMBEDD_DIM)\n        self.out = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(2*EMBEDD_DIM, 1)\n        )\n    \n    def forward(self, xanchor, xpos=None, xneg=None):\n        xanchor = self.embedding(xanchor)\n        if xpos is not None:\n            xpos = self.embedding(xpos)\n            x = torch.cat([xanchor, xpos], dim=-1)\n        \n        if xneg is not None:\n            neg_count_ = xneg.shape[1]\n            xneg = self.embedding(xneg)\n            xanchor = xanchor.unsqueeze(dim=1).expand(-1, neg_count_, -1)\n            x = torch.cat([xanchor, xneg], dim=-1)\n        return self.out(x)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.717151Z","iopub.execute_input":"2023-01-08T16:34:34.718099Z","iopub.status.idle":"2023-01-08T16:34:34.728879Z","shell.execute_reply.started":"2023-01-08T16:34:34.718046Z","shell.execute_reply":"2023-01-08T16:34:34.727482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# negative sampling","metadata":{}},{"cell_type":"code","source":"def negative_sampling(batch_size, X):\n    X = X.unique()\n    X = X[X!=0]\n    n_samples = len(X)\n    \n    perm_indices = torch.randint(0, n_samples, (batch_size, n_samples), device=device)\n    xneg = X[perm_indices]\n    xneg = xneg[:, :MAX_NEG_SAMPLES]\n    return xneg","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.730208Z","iopub.execute_input":"2023-01-08T16:34:34.730553Z","iopub.status.idle":"2023-01-08T16:34:34.748868Z","shell.execute_reply.started":"2023-01-08T16:34:34.730526Z","shell.execute_reply":"2023-01-08T16:34:34.746510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# generate pairs","metadata":{}},{"cell_type":"code","source":"def generate_pairs(X, mask):\n    xanchor1 = X[:, :-1].flatten()\n    xpos1 = X[:, 1:].flatten()\n    mask1 = mask[:, :-1].flatten()\n    \n    xanchor1 = xanchor1[mask1]\n    xpos1 = xpos1[mask1]\n    \n    if X.shape[1] > 2:\n        xanchor2 = X[:, :-2].flatten()\n        xpos2 = X[:, 2:].flatten()\n        mask2 = mask[:, :-2].flatten()\n        xanchor2 = xanchor2[mask2]; xpos2=xpos2[mask2]\n    \n    xanchor = torch.cat([xanchor1, xanchor2])\n    xpos = torch.cat([xpos1, xpos2])\n    return xanchor, xpos","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.751031Z","iopub.execute_input":"2023-01-08T16:34:34.752212Z","iopub.status.idle":"2023-01-08T16:34:34.763227Z","shell.execute_reply.started":"2023-01-08T16:34:34.752172Z","shell.execute_reply":"2023-01-08T16:34:34.762095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = OttoDataset(np.arange(mat.shape[0]), mat)\ntrain_dataloader = torch.utils.data.DataLoader(train_dataset, \n                                               batch_size=BATCH_SIZE, \n                                               shuffle=True)\n\nprint(len(train_dataset), len(train_dataloader))","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.764916Z","iopub.execute_input":"2023-01-08T16:34:34.765597Z","iopub.status.idle":"2023-01-08T16:34:34.802555Z","shell.execute_reply.started":"2023-01-08T16:34:34.765557Z","shell.execute_reply":"2023-01-08T16:34:34.800382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = OttoItem2VecModel().to(device)\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)\nschedular = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, \n                                                       N_EPOCHS * len(train_dataloader),\n                                                       eta_min = 5e-6)","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:34.804637Z","iopub.execute_input":"2023-01-08T16:34:34.805123Z","iopub.status.idle":"2023-01-08T16:34:35.598462Z","shell.execute_reply.started":"2023-01-08T16:34:34.805091Z","shell.execute_reply":"2023-01-08T16:34:35.596081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for e in range(N_EPOCHS):\n    epoch_pos_loss=[]; epoch_neg_loss=[]\n    \n    for (X, mask, slen) in train_dataloader:\n        max_seqlen = slen.max()\n        X = X[:, :max_seqlen]\n        mask = mask[:, :max_seqlen]\n        \n        X = X.to(device)\n        mask = mask.to(device)\n        \n        (xanchor, xpos) = generate_pairs(X, mask)\n\n        batch_size = len(xanchor)\n        xneg = negative_sampling(batch_size, X)\n\n        #xanchor = xanchor.to(device)\n        #xpos = xpos.to(device)\n        #xneg = xneg.to(device)\n        \n        model.train()\n        \n        optimizer.zero_grad()\n        \n        ypos = model(xanchor, xpos=xpos)\n        pos_loss = -torch.log(1e-8+ypos.sigmoid()).mean()\n        pos_loss.backward()\n        \n        yneg = model(xanchor, xneg=xneg)\n        neg_loss = -torch.log(1e-8+1-yneg.sigmoid()).mean()\n        neg_loss.backward()\n        \n        optimizer.step()\n        schedular.step()\n        \n        epoch_pos_loss.append(pos_loss.item())\n        epoch_neg_loss.append(neg_loss.item())\n    \n    print(\"Epoch:{} | pos loss:{:.4f} | negloss:{:.4f}\".format(e, np.mean(epoch_pos_loss), np.mean(epoch_neg_loss)))\n    plt.title(\"epoch pos loss\")\n    plt.plot(epoch_pos_loss)\n    plt.show()\n    plt.title(\"epoch neg loss\")\n    plt.plot(epoch_neg_loss)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:34:35.599880Z","iopub.execute_input":"2023-01-08T16:34:35.600271Z","iopub.status.idle":"2023-01-08T16:35:05.706824Z","shell.execute_reply.started":"2023-01-08T16:34:35.600236Z","shell.execute_reply":"2023-01-08T16:35:05.705716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, \"model.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-01-08T16:35:05.709225Z","iopub.execute_input":"2023-01-08T16:35:05.710229Z","iopub.status.idle":"2023-01-08T16:35:06.446867Z","shell.execute_reply.started":"2023-01-08T16:35:05.710199Z","shell.execute_reply":"2023-01-08T16:35:06.446114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}