{"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":"I tried Graph neural network (RGCN) approaches using Deep Graph Libarry DGL.  \nPlease upvote if this notebook is useful.","metadata":{}},{"cell_type":"code","source":"!conda install -c dglteam dgl-cuda11.0 -y\n!conda install -c conda-forge swifter -y","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2022-04-22T15:24:20.358002Z","iopub.execute_input":"2022-04-22T15:24:20.358694Z","iopub.status.idle":"2022-04-22T15:26:30.948685Z","shell.execute_reply.started":"2022-04-22T15:24:20.358608Z","shell.execute_reply":"2022-04-22T15:26:30.947821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import dgl\nimport dgl.nn as dglnn\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport pandas as pd\nimport numpy as np\nimport tqdm\nimport joblib\nfrom annoy import AnnoyIndex\nimport swifter\n#from scipy.spatial import cKDTree\nfrom sklearn.preprocessing import MinMaxScaler\nfrom sklearn.metrics import recall_score, roc_auc_score","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:26:30.95245Z","iopub.execute_input":"2022-04-22T15:26:30.952679Z","iopub.status.idle":"2022-04-22T15:26:35.530341Z","shell.execute_reply.started":"2022-04-22T15:26:30.95265Z","shell.execute_reply":"2022-04-22T15:26:35.52955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    transaction_path = \"../input/h-and-m-personalized-fashion-recommendations/transactions_train.csv\"\n    transaction_2020_path = \"../input/h-and-m-split-dataset-by-year/transactions_train_2020.csv\"\n    transaction_2019_path = \"../input/h-and-m-split-dataset-by-year/transactions_train_2019.csv\"\n    customer_path = \"../input/h-and-m-personalized-fashion-recommendations/customers.csv\"\n    article_path = \"../input/h-and-m-personalized-fashion-recommendations/articles.csv\"\n    image_feat_path = \"../input/h-and-m-swint-image-embedding/swin_tiny_patch4_window7_224_emb.csv.gz\"\n    sample_submission_path = \"../input/h-and-m-personalized-fashion-recommendations/sample_submission.csv\"\n\n    output_dir = \"../output/\"\n    #start_date = '2020-08-01'\n    start_date = '2020-09-01'\n\n    image_feat_dim = 768\n    text_feat_dim = 384\n    \n    # train\n    #n_fold = 2\n    n_fold = 5\n    #epoch = 50\n    epoch = 100\n\n    seed = 2022\n    \n    # graph \n    customer_node = \"customer\"\n    article_node = \"article\"\n    buy_edge = \"buy\"\n    bought_by_edge = \"bought_by\"\n\n    buy_store_edge = \"buy_store\"\n    bought_by_store_edge = \"bought_by_store\"\n    buy_online_edge = \"buy_online\"\n    bought_by_online_edge = \"bought_by_online\"\n    age_same = \"age\"\n    age_same_by = \"age_by\"\n    \n    \n    in_feat_dim = 200\n    hidden_feat_dim = 500\n    out_feat_dim = 200\n\n    #device=torch.device(\"cpu\")\n    device = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n\n    model_path = \"model.pth\"","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:26:35.53167Z","iopub.execute_input":"2022-04-22T15:26:35.53193Z","iopub.status.idle":"2022-04-22T15:26:35.563192Z","shell.execute_reply.started":"2022-04-22T15:26:35.531896Z","shell.execute_reply":"2022-04-22T15:26:35.561791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans = pd.read_csv(Config.transaction_path, dtype={'article_id': 'str'})\ndf_trans = df_trans[df_trans.t_dat >= Config.start_date]\ndf_trans.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:26:35.565322Z","iopub.execute_input":"2022-04-22T15:26:35.565769Z","iopub.status.idle":"2022-04-22T15:27:41.095786Z","shell.execute_reply.started":"2022-04-22T15:26:35.565727Z","shell.execute_reply":"2022-04-22T15:27:41.095097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.096944Z","iopub.execute_input":"2022-04-22T15:27:41.098271Z","iopub.status.idle":"2022-04-22T15:27:41.104359Z","shell.execute_reply.started":"2022-04-22T15:27:41.098231Z","shell.execute_reply":"2022-04-22T15:27:41.103419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.t_dat.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.105983Z","iopub.execute_input":"2022-04-22T15:27:41.106258Z","iopub.status.idle":"2022-04-22T15:27:41.117828Z","shell.execute_reply.started":"2022-04-22T15:27:41.106215Z","shell.execute_reply":"2022-04-22T15:27:41.11701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ranking = df_trans[[\"article_id\", \"customer_id\"]].groupby(\"article_id\").count().reset_index().sort_values(\"customer_id\", ascending=False)\ndf_ranking = df_ranking[df_ranking[\"customer_id\"] > 10]\ndf_ranking.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.120267Z","iopub.execute_input":"2022-04-22T15:27:41.120771Z","iopub.status.idle":"2022-04-22T15:27:41.357632Z","shell.execute_reply.started":"2022-04-22T15:27:41.120664Z","shell.execute_reply":"2022-04-22T15:27:41.356894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# th回以上transactionがあるやつに絞る。\n\ndef transaction_count_filter(df, th=10):\n    df_ranking = df_trans[[\"article_id\", \"customer_id\"]].groupby(\"article_id\").count().reset_index().sort_values(\"customer_id\", ascending=False)\n    df_ranking = df_ranking[df_ranking[\"customer_id\"] >= th]\n    df = df.merge(df_ranking[[\"article_id\"]], on=\"article_id\", how=\"inner\").reset_index(drop=True)\n    df = df.drop_duplicates().reset_index(drop=True)\n    return df\n","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.358995Z","iopub.execute_input":"2022-04-22T15:27:41.359424Z","iopub.status.idle":"2022-04-22T15:27:41.366329Z","shell.execute_reply.started":"2022-04-22T15:27:41.359386Z","shell.execute_reply":"2022-04-22T15:27:41.365508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.367731Z","iopub.execute_input":"2022-04-22T15:27:41.368193Z","iopub.status.idle":"2022-04-22T15:27:41.379002Z","shell.execute_reply.started":"2022-04-22T15:27:41.368152Z","shell.execute_reply":"2022-04-22T15:27:41.378245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transaction_count_filter(df_trans).shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:41.383269Z","iopub.execute_input":"2022-04-22T15:27:41.383495Z","iopub.status.idle":"2022-04-22T15:27:42.400743Z","shell.execute_reply.started":"2022-04-22T15:27:41.383457Z","shell.execute_reply":"2022-04-22T15:27:42.399988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans = transaction_count_filter(df_trans)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:42.401911Z","iopub.execute_input":"2022-04-22T15:27:42.403303Z","iopub.status.idle":"2022-04-22T15:27:43.395649Z","shell.execute_reply.started":"2022-04-22T15:27:42.403262Z","shell.execute_reply":"2022-04-22T15:27:43.394908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:43.396751Z","iopub.execute_input":"2022-04-22T15:27:43.396996Z","iopub.status.idle":"2022-04-22T15:27:43.406802Z","shell.execute_reply.started":"2022-04-22T15:27:43.396965Z","shell.execute_reply":"2022-04-22T15:27:43.405929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = pd.read_csv(Config.sample_submission_path)\ndf_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:43.408224Z","iopub.execute_input":"2022-04-22T15:27:43.408601Z","iopub.status.idle":"2022-04-22T15:27:47.764898Z","shell.execute_reply.started":"2022-04-22T15:27:43.408565Z","shell.execute_reply":"2022-04-22T15:27:47.764178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer = pd.read_csv(Config.customer_path)\ndf_customer.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:47.766208Z","iopub.execute_input":"2022-04-22T15:27:47.766617Z","iopub.status.idle":"2022-04-22T15:27:52.600078Z","shell.execute_reply.started":"2022-04-22T15:27:47.766578Z","shell.execute_reply":"2022-04-22T15:27:52.599362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"users = df_trans[\"customer_id\"].unique().tolist()\ndf_customer_node = pd.DataFrame( {\"customer_id\": users,\n                                  \"customer_node_id\": [i for i in range(len(users))]})","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:52.60139Z","iopub.execute_input":"2022-04-22T15:27:52.60179Z","iopub.status.idle":"2022-04-22T15:27:52.946533Z","shell.execute_reply.started":"2022-04-22T15:27:52.60175Z","shell.execute_reply":"2022-04-22T15:27:52.945796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"article = df_trans[\"article_id\"].unique().tolist()\ndf_article_node = pd.DataFrame( {\"article_id\": article, \"article_node_id\": [i for i in range(len(article))]})","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:52.94778Z","iopub.execute_input":"2022-04-22T15:27:52.948026Z","iopub.status.idle":"2022-04-22T15:27:53.016552Z","shell.execute_reply.started":"2022-04-22T15:27:52.947993Z","shell.execute_reply":"2022-04-22T15:27:53.015835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_article = pd.read_csv(Config.article_path, dtype={'article_id': 'str'})\ndf_article = df_article.merge(df_article_node[[\"article_id\"]], on=\"article_id\", how=\"inner\")\n","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:53.017927Z","iopub.execute_input":"2022-04-22T15:27:53.018181Z","iopub.status.idle":"2022-04-22T15:27:53.992525Z","shell.execute_reply.started":"2022-04-22T15:27:53.018147Z","shell.execute_reply":"2022-04-22T15:27:53.991768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans = df_trans.merge(df_customer_node, how='inner', on=\"customer_id\")\ndf_trans = df_trans.merge(df_article_node, how=\"inner\", on=\"article_id\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:53.993707Z","iopub.execute_input":"2022-04-22T15:27:53.993954Z","iopub.status.idle":"2022-04-22T15:27:54.559599Z","shell.execute_reply.started":"2022-04-22T15:27:53.993921Z","shell.execute_reply":"2022-04-22T15:27:54.558731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:54.560878Z","iopub.execute_input":"2022-04-22T15:27:54.561337Z","iopub.status.idle":"2022-04-22T15:27:54.585583Z","shell.execute_reply.started":"2022-04-22T15:27:54.561293Z","shell.execute_reply":"2022-04-22T15:27:54.58056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cond_store = df_trans[\"sales_channel_id\"] == 1\n# cond_online = df_trans[\"sales_channel_id\"] == 2\n\n# graph = dgl.heterograph({\n#     (Config.customer_node, Config.buy_store_edge, Config.article_node): \n#                (df_trans[cond_store].loc[:, \"customer_node_id\"].tolist(), df_trans[cond_store].loc[:, \"article_node_id\"].tolist()),\n#     (Config.article_node, Config.bought_by_store_edge, Config.customer_node): \n#               (df_trans[cond_store].loc[:, \"article_node_id\"].tolist(), df_trans[cond_store].loc[:, \"customer_node_id\"].tolist() ),\n#     (Config.customer_node, Config.buy_online_edge, Config.article_node): \n#                (df_trans[cond_online].loc[:, \"customer_node_id\"].tolist(), df_trans[cond_online].loc[:, \"article_node_id\"].tolist()),\n#     (Config.article_node, Config.bought_by_online_edge, Config.customer_node): \n#               (df_trans[cond_online].loc[:, \"article_node_id\"].tolist(), df_trans[cond_online].loc[:, \"customer_node_id\"].tolist())\n# })\n\ngraph = dgl.heterograph({\n    (Config.customer_node, Config.buy_edge, Config.article_node): \n               (df_trans.loc[:, \"customer_node_id\"].tolist(), df_trans.loc[:, \"article_node_id\"].tolist()),\n    (Config.article_node, Config.bought_by_edge, Config.customer_node): \n              (df_trans.loc[:, \"article_node_id\"].tolist(), df_trans.loc[:, \"customer_node_id\"].tolist() ),  \n})","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:54.587082Z","iopub.execute_input":"2022-04-22T15:27:54.587398Z","iopub.status.idle":"2022-04-22T15:27:55.158983Z","shell.execute_reply.started":"2022-04-22T15:27:54.587361Z","shell.execute_reply":"2022-04-22T15:27:55.158237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"graph","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.160194Z","iopub.execute_input":"2022-04-22T15:27:55.160455Z","iopub.status.idle":"2022-04-22T15:27:55.167218Z","shell.execute_reply.started":"2022-04-22T15:27:55.160421Z","shell.execute_reply":"2022-04-22T15:27:55.166363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RGCN(nn.Module):\n    def __init__(self, in_feat, hidden_feat, out_feat, rel_names):\n        super().__init__()\n        self.conv1 = dglnn.HeteroGraphConv({\n                rel : dglnn.GraphConv(in_feat, hidden_feat, norm='right')\n                for rel in rel_names\n            })\n        self.conv2 = dglnn.HeteroGraphConv({\n                rel : dglnn.GraphConv(hidden_feat, out_feat, norm='right')\n                for rel in rel_names\n            })\n\n    def forward(self, blocks, x):\n        x = self.conv1(blocks[0], x)\n        x = {key:F.relu(val)  for key, val in x.items()}\n        x = self.conv2(blocks[1], x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.168815Z","iopub.execute_input":"2022-04-22T15:27:55.169319Z","iopub.status.idle":"2022-04-22T15:27:55.178418Z","shell.execute_reply.started":"2022-04-22T15:27:55.169272Z","shell.execute_reply":"2022-04-22T15:27:55.177729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Dense(nn.Module):\n    \n#     def __init__(self, col, image_mask, text_mask, image_dim, text_dim):\n#         super().__init__()\n#         self.col = col\n#         self.image_mask\n#         self.text_mask\n#         self.image_dence = nn.Linear(image_dim, 250)\n#         self.text_dence = nn.Linear(text_dim, 250)\n#         self.other_dence = nn.Liner()\n        \n#     def forward(self, x):\n#         _image = self.image_dence(x[self.col][:, self.image_mask])\n#         _text = self.text_dence(x[self.col][:, self.text_mask])\n#         _x = torch.cat([_image, _text])\n#         x[self.col] = _x\n        \n#         return _x","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.180344Z","iopub.execute_input":"2022-04-22T15:27:55.180955Z","iopub.status.idle":"2022-04-22T15:27:55.188062Z","shell.execute_reply.started":"2022-04-22T15:27:55.180916Z","shell.execute_reply":"2022-04-22T15:27:55.187422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ScorePredictor(nn.Module):\n    def forward(self, edge_subgraph, x):\n        with edge_subgraph.local_scope():\n            edge_subgraph.ndata['h'] = x\n            for etype in edge_subgraph.canonical_etypes:\n                edge_subgraph.apply_edges(\n                    dgl.function.u_dot_v('h', 'h', 'score'), etype=etype)\n            return edge_subgraph.edata['score']\n\nclass Model(nn.Module):\n    def __init__(self, \n                 in_features, \n                 hidden_features, \n                 out_features,\n                 item_dim,\n                 user_dim,                 \n                 etypes,\n                 item_col=Config.article_node,\n                 user_col=Config.customer_node):\n        \n        super().__init__()\n        \n        self.item_dence = nn.Linear(item_dim, in_features)\n        self.user_dence = nn.Linear(user_dim, in_features)\n        self.item_col = item_col\n        self.user_col = user_col\n        \n        self.hidden_featuers = hidden_features\n        self.out_featuers = out_features\n\n        self.rgcn = RGCN(in_features, hidden_features, out_features, etypes)\n        \n        self.score = ScorePredictor()\n\n    def forward(self, blocks, x):\n        x[self.user_col] = F.relu(self.user_dence(x[self.user_col]))\n        x[self.item_col] = F.relu(self.item_dence(x[self.item_col]))\n        x = self.rgcn(blocks, x)        \n        return x    \n    ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.19116Z","iopub.execute_input":"2022-04-22T15:27:55.191362Z","iopub.status.idle":"2022-04-22T15:27:55.203438Z","shell.execute_reply.started":"2022-04-22T15:27:55.19133Z","shell.execute_reply":"2022-04-22T15:27:55.202718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_loss(pos_score, neg_score, canonical_etypes):\n    # Margin loss\n    all_losses = []\n    for given_type in canonical_etypes:\n        n_edges = pos_score[given_type].shape[0]\n        if n_edges == 0:\n            continue\n        all_losses.append((1 - neg_score[given_type].view(n_edges, -1) + pos_score[given_type].unsqueeze(1)).clamp(min=0).mean())\n    return torch.stack(all_losses, dim=0).mean()\n\ndef compute_loss_bce(pos_score, neg_score, canonical_etypes):\n    all_losses = []\n    criterion = torch.nn.BCEWithLogitsLoss()\n    for given_type in canonical_etypes:\n        _pos_score = pos_score[given_type].squeeze(1)\n        _neg_score = neg_score[given_type].squeeze(1)\n        \n        pred = torch.cat([_pos_score, _neg_score])\n        \n        label = torch.cat([torch.ones(len(_pos_score)), torch.zeros(len(_neg_score))]).to(Config.device)\n        loss = criterion(pred, label)\n        all_losses.append(loss)\n        \n    return torch.stack(all_losses ,dim=0).mean() \n\ndef compute_auc(pos_score, neg_score, canonical_etypes):\n    aucs = []\n    for etype in canonical_etypes:\n        _pos_score = pos_score[etype].squeeze(1).to(\"cpu\").detach()\n        _neg_score = neg_score[etype].squeeze(1).to(\"cpu\").detach()\n        pred = torch.cat([_pos_score, _neg_score])\n        label = torch.cat([torch.ones(len(_pos_score)), torch.zeros(len(_neg_score))])\n        \n        roc_auc = roc_auc_score(label.numpy(), pred.numpy())\n        aucs.append(roc_auc)\n        \n    return np.mean(aucs)\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:32:42.686116Z","iopub.execute_input":"2022-04-22T15:32:42.6868Z","iopub.status.idle":"2022-04-22T15:32:42.697878Z","shell.execute_reply.started":"2022-04-22T15:32:42.686763Z","shell.execute_reply":"2022-04-22T15:32:42.696745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(train_graph, X_dic, train_dataloader, model):    \n    \n    model = model.to(Config.device)\n    opt = torch.optim.Adam(model.parameters())\n    \n    model.train()\n    for i in range(Config.epoch):\n        \n        # train loop\n        for input_nodes, positive_graph, negative_graph, blocks in train_dataloader:\n            \n            blocks = [b.to(Config.device) for b in blocks]\n            positive_graph = positive_graph.to(Config.device)\n            negative_graph = negative_graph.to(Config.device)\n\n            feature = {\n                ntype: X_dic[ntype][input_nodes[ntype]].to(Config.device) for ntype in train_graph.ntypes\n            }            \n\n            emb_dict = model(blocks, feature)\n            pos_score = model.score(positive_graph, emb_dict)\n            neg_score = model.score(negative_graph, emb_dict)\n\n            loss = compute_loss_bce(pos_score, neg_score, train_graph.canonical_etypes)\n            opt.zero_grad()\n            loss.backward()\n            opt.step()\n            auc = compute_auc(pos_score, neg_score, train_graph.canonical_etypes)\n    \n        print(f\"epoch: {i} | train loss:{loss.item()} | AUC {auc}\")\n\n        #evaluate()\n    \n    return model\n\n\ndef evaluate(train_graph, graph, X_dic, model, valid_eid_dict):\n    \n    emb = inference(train_graph, X_dic, model)\n    score_list = []\n    for etype in graph.canonical_etypes:\n        src, dst = graph.find_edges(valid_eid_dict[etype], etype=etype)\n        score = (emb[etype][src] * emb[etype][dst]).sum(1)\n        score_list.append(score)\n\n    print(score)\n        \n\n\ndef inference(graph, X_dic, model):\n    model = model.to(Config.device)\n    model.eval()\n\n    dataloader = dgl.dataloading.NodeDataLoader(\n                graph,\n                {\n                    Config.article_node: torch.arange(graph.number_of_nodes(ntype=Config.article_node)),\n                    Config.customer_node: torch.arange(graph.number_of_nodes(ntype=Config.customer_node))\n                },\n                dgl.dataloading.MultiLayerFullNeighborSampler(1),\n                batch_size=1024,\n                shuffle=True,\n                drop_last=False,\n            )\n\n\n    with torch.no_grad():\n        for n_layer in range(2):\n            if n_layer == 0:\n                y = {ntype: torch.zeros(graph.number_of_nodes(ntype), model.hidden_featuers) \n                    for ntype in graph.ntypes}\n                 \n            else:\n                y = {ntype: torch.zeros(graph.number_of_nodes(ntype), model.out_featuers) \n                    for ntype in graph.ntypes}\n\n\n            for input_nodes, output_nodes, blocks in dataloader:\n                block = blocks[0].to(Config.device)            \n\n                x = {\n                    ntype: X_dic[ntype][input_nodes[ntype]].to(Config.device) for ntype in graph.ntypes\n                }                \n\n                if n_layer == 0:\n                    x[model.user_col] = F.relu(model.user_dence(x[model.user_col]))\n                    x[model.item_col] = F.relu(model.item_dence(x[model.item_col]))\n                    h = model.rgcn.conv1(block, x)  \n                    h = {key:F.relu(val) for key, val in h.items()}                    \n\n                else:\n                    h = model.rgcn.conv2(block, x)\n    \n                for ntype in graph.ntypes:\n                    y[ntype][output_nodes[ntype]] = h[ntype].cpu()                         \n                    \n            X_dic = y\n\n    return y\n    \n    \ndef validation(train_graph, graph, valid_eid_dict, x_dict, model, batch_size, fanout, num_workers):\n    scores = []\n    for src_ntype, etype, dst_ntype in graph.canonical_etypes:\n        label = torch.ones(len(valid_eid_dict[etype]))\n        src, dst = graph.find_edges(valid_eid_dict[etype], etype=etype)\n        src_emb = emb[src_ntype][src]\n        dst_emb = emb[dst_ntype][dst]\n        score = torch.sigmoid(score)\n        score = score > 0.5\n        recall = recall_score(label, score)\n        scores.append(recall)\n        \n    return np.mean(scores)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:32:27.466781Z","iopub.execute_input":"2022-04-22T15:32:27.467054Z","iopub.status.idle":"2022-04-22T15:32:27.490809Z","shell.execute_reply.started":"2022-04-22T15:32:27.467022Z","shell.execute_reply":"2022-04-22T15:32:27.490037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer_node.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.248537Z","iopub.execute_input":"2022-04-22T15:27:55.249545Z","iopub.status.idle":"2022-04-22T15:27:55.264071Z","shell.execute_reply.started":"2022-04-22T15:27:55.249503Z","shell.execute_reply":"2022-04-22T15:27:55.263259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_customer_feat(df, df_node, df_trans):\n\n    customer_drop_cols = [\"postal_code\"]\n    customer_dummy_cols = [\"club_member_status\", \"fashion_news_frequency\"]\n\n\n    df = df.drop(customer_drop_cols, axis=1)\n    df.loc[:, \"FN\"] = df[\"FN\"].fillna(0)\n    df.loc[:, \"Active\"] = df[\"Active\"].fillna(0)\n    df.loc[:, \"club_member_status\"] = df[\"club_member_status\"].fillna(\"NONE\")\n    df.loc[:, \"fashion_news_frequency\"] = df[\"fashion_news_frequency\"].fillna(\"NONE\")\n    df.loc[:, \"age\"] = df[\"age\"].fillna(0)\n    df.loc[:, \"age\"] = np.log1p(df[\"age\"])\n\n    df = pd.get_dummies(df, columns=customer_dummy_cols)\n    \n    \n    # price_mean\n    df_price_mean = df_trans[['customer_id', 'price']].groupby(\"customer_id\").mean().reset_index()\n    df_price_mean[\"price\"] = np.log(df_price_mean[\"price\"] * 1000000)\n    df = df.merge(df_price_mean, on=\"customer_id\", how=\"left\")\n    \n    # number of transaction per customer\n    cond_store = df_trans[\"sales_channel_id\"] == 1\n    cond_online = df_trans[\"sales_channel_id\"] == 2\n    \n    df_trans_count_offline = df_trans[cond_store].groupby(\"customer_id\").count().reset_index()[[\"customer_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_offline\"})\n    df_trans_count_online = df_trans[cond_online].groupby(\"customer_id\").count().reset_index()[[\"customer_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_online\"})\n    df_trans_count_both = df_trans.groupby(\"customer_id\").count().reset_index()[[\"customer_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_both\"})\n    \n    df = df.merge(df_trans_count_offline, on=\"customer_id\", how=\"left\").fillna(0)\n    df = df.merge(df_trans_count_online, on=\"customer_id\", how=\"left\").fillna(0)\n    df = df.merge(df_trans_count_both, on=\"customer_id\", how=\"left\").fillna(0)\n    \n    \n    df = df.merge(df_node, on=\"customer_id\", how=\"inner\")\n    df = df.sort_values(\"customer_node_id\").reset_index(drop=True)\n    df = df.drop([\"customer_id\", \"customer_node_id\"], axis=1)\n\n\n    return df\n\n\ndef get_article_table_feat(df):\n    #article_id_cols = [\"product_code\", \"product_type_no\", \"graphical_appearance_no\", \"colour_group_code\",\n    #         \"perceived_colour_value_id\", \"perceived_colour_master_id\", \"department_no\", \"index_group_no\",\n    #           \"section_no\", \"garment_group_no\"]\n\n    article_dummy_cols = [\"product_type_name\", \"product_group_name\", \"graphical_appearance_name\", \"colour_group_name\",\n                         \"perceived_colour_value_name\", \"perceived_colour_master_name\",\n                         #\"department_name\",\n                         \"index_name\", \"index_group_name\", \"section_name\", \"garment_group_name\"]\n\n    article_drop_cols = [\"index_code\", \"prod_name\", \"detail_desc\", \"department_name\"]\n\n    df = df.drop(article_drop_cols, axis=1)\n    df = pd.get_dummies(df, columns=article_dummy_cols)\n    return df\n\ndef get_article_image_feat(df):\n    pass\n\ndef get_article_text_feat(df):\n    pass\n\ndef create_article_feat(df, df_node, df_trans):\n    df_table_feat = get_article_table_feat(df)\n    \n    # price_mean\n    df_price_mean = df_trans[['article_id', 'price']].groupby(\"article_id\").mean().reset_index()\n    df_price_mean[\"price\"] = np.log(df_price_mean[\"price\"] * 1000000)\n    df_table_feat = df_table_feat.merge(df_price_mean, on=\"article_id\", how=\"left\")\n    \n    # number of transaction per customer\n    cond_store = df_trans[\"sales_channel_id\"] == 1\n    cond_online = df_trans[\"sales_channel_id\"] == 2\n    \n    df_trans_count_offline = df_trans[cond_store].groupby(\"article_id\").count().reset_index()[[\"article_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_offline\"})\n    df_trans_count_online = df_trans[cond_online].groupby(\"article_id\").count().reset_index()[[\"article_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_online\"})\n    df_trans_count_both = df_trans.groupby(\"article_id\").count().reset_index()[[\"article_id\", \"t_dat\"]].rename(columns={\"t_dat\": \"count_both\"})\n    \n    df_table_feat = df_table_feat.merge(df_trans_count_offline, on=\"article_id\", how=\"left\").fillna(0)\n    df_table_feat = df_table_feat.merge(df_trans_count_online, on=\"article_id\", how=\"left\").fillna(0)\n    df_table_feat = df_table_feat.merge(df_trans_count_both, on=\"article_id\", how=\"left\").fillna(0)\n    \n    df = df_table_feat\n    df = df.merge(df_node, on=\"article_id\", how=\"inner\")\n    df = df.sort_values(\"article_node_id\").reset_index()\n    df = df.drop([\"article_id\", \"article_node_id\"], axis=1)\n\n    return df\n\ndef create_graph_feature(df_article, df_article_node, df_customer, df_customer_node, df_trans):    \n    X_dic = {}\n    df_customer_feat = create_customer_feat(df_customer, df_customer_node, df_trans)\n    print(df_customer_feat.isna().any())\n    df_article_feat = create_article_feat(df_article, df_article_node, df_trans)\n    print(df_article_feat.isna().any())\n\n    scaler = MinMaxScaler()\n\n    X_dic[Config.customer_node] = torch.Tensor(scaler.fit_transform(df_customer_feat.values))\n    X_dic[Config.article_node] = torch.Tensor(scaler.fit_transform(df_article_feat.values))\n    \n    return X_dic\n","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.269998Z","iopub.execute_input":"2022-04-22T15:27:55.27023Z","iopub.status.idle":"2022-04-22T15:27:55.296788Z","shell.execute_reply.started":"2022-04-22T15:27:55.270204Z","shell.execute_reply":"2022-04-22T15:27:55.296047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_trans.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.298199Z","iopub.execute_input":"2022-04-22T15:27:55.298492Z","iopub.status.idle":"2022-04-22T15:27:55.317384Z","shell.execute_reply.started":"2022-04-22T15:27:55.298456Z","shell.execute_reply":"2022-04-22T15:27:55.316408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic = create_graph_feature(df_article, df_article_node, df_customer, df_customer_node, df_trans)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:27:55.318731Z","iopub.execute_input":"2022-04-22T15:27:55.319157Z","iopub.status.idle":"2022-04-22T15:28:04.168872Z","shell.execute_reply.started":"2022-04-22T15:27:55.319097Z","shell.execute_reply":"2022-04-22T15:28:04.168135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic[\"customer\"].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.170201Z","iopub.execute_input":"2022-04-22T15:28:04.17061Z","iopub.status.idle":"2022-04-22T15:28:04.177985Z","shell.execute_reply.started":"2022-04-22T15:28:04.170571Z","shell.execute_reply":"2022-04-22T15:28:04.177012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic[\"article\"].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.179667Z","iopub.execute_input":"2022-04-22T15:28:04.18Z","iopub.status.idle":"2022-04-22T15:28:04.1879Z","shell.execute_reply.started":"2022-04-22T15:28:04.179962Z","shell.execute_reply":"2022-04-22T15:28:04.186973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic[\"customer\"].shape, X_dic[\"article\"].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.189499Z","iopub.execute_input":"2022-04-22T15:28:04.190157Z","iopub.status.idle":"2022-04-22T15:28:04.197948Z","shell.execute_reply.started":"2022-04-22T15:28:04.190079Z","shell.execute_reply":"2022-04-22T15:28:04.196993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"args = {\n    \"in_features\": Config.in_feat_dim,\n    \"hidden_features\": Config.hidden_feat_dim ,\n    \"out_features\": Config.out_feat_dim,\n    \"item_dim\": X_dic[Config.article_node].shape[1],\n    \"user_dim\": X_dic[Config.customer_node].shape[1],\n    \"etypes\": [Config.buy_edge, Config.bought_by_edge]\n    #[Config.buy_online_edge, Config.bought_by_online_edge, Config.buy_store_edge, Config.bought_by_store_edge],\n}\n\nmodel = Model(**args)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.199559Z","iopub.execute_input":"2022-04-22T15:28:04.200108Z","iopub.status.idle":"2022-04-22T15:28:04.221439Z","shell.execute_reply.started":"2022-04-22T15:28:04.200068Z","shell.execute_reply":"2022-04-22T15:28:04.220707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://zqfang.github.io/2021-08-12-graph-linkpredict/\n\n\n# train/ validation split\n\n# https://github.com/dglai/WWW20-Hands-on-Tutorial/blob/master/_legacy/basic_apps/BasicTasks_pytorch.ipynb\n\n# def train_valid_split(graph, train_rate=0.8):\n#     _train_eid_dict = {}\n#     _valid_eid_dict = {}\n\n#     for etype in graph.canonical_etypes:\n#         eids = np.random.permutation(graph.num_edges(etype))    \n#         train_eids = eids[:int(len(eids) * train_rate)]\n#         valid_eids = eids[int(len(eids) * train_rate):]\n\n#         _train_eid_dict[etype] = train_eids\n#         _valid_eid_dict[etype] = valid_eids\n        \n#     train_graph = graph.edge_subgraph(_train_eid_dict, relabel_nodes=False, store_ids=True)\n#     valid_graph = graph.edge_subgraph(_valid_eid_dict, relabel_nodes=False, store_ids=True)\n\n#     return train_graph, valid_graph, _train_eid_dict, _valid_eid_dict\n\n# train_graph, valid_graph, _train_eid_dict, _valid_eid_dict = train_valid_split(graph)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.222706Z","iopub.execute_input":"2022-04-22T15:28:04.222973Z","iopub.status.idle":"2022-04-22T15:28:04.227592Z","shell.execute_reply.started":"2022-04-22T15:28:04.222936Z","shell.execute_reply":"2022-04-22T15:28:04.226818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#_valid_eid_dict","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.22915Z","iopub.execute_input":"2022-04-22T15:28:04.229601Z","iopub.status.idle":"2022-04-22T15:28:04.238135Z","shell.execute_reply.started":"2022-04-22T15:28:04.229562Z","shell.execute_reply":"2022-04-22T15:28:04.237379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#graph.find_edges(126293, etype=('article','bought_by_online','customer'))","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.241032Z","iopub.execute_input":"2022-04-22T15:28:04.241586Z","iopub.status.idle":"2022-04-22T15:28:04.247667Z","shell.execute_reply.started":"2022-04-22T15:28:04.241538Z","shell.execute_reply":"2022-04-22T15:28:04.246785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_eid_dict = {\n    #Config.buy_store_edge: torch.arange(graph.num_edges(Config.buy_store_edge)),\n    #Config.bought_by_store_edge: torch.arange(graph.num_edges(Config.bought_by_store_edge)),\n    #Config.buy_online_edge: torch.arange(graph.num_edges(Config.buy_online_edge)),\n    #Config.bought_by_online_edge: torch.arange(graph.num_edges(Config.bought_by_online_edge)),\n    Config.buy_edge: torch.arange(graph.num_edges(Config.buy_edge)),\n    Config.bought_by_edge: torch.arange(graph.num_edges(Config.bought_by_edge)),\n   \n}\n\nreverse_types = {\n    #Config.buy_store_edge: Config.bought_by_store_edge,\n    #Config.bought_by_store_edge: Config.buy_store_edge,\n    #Config.buy_online_edge: Config.bought_by_online_edge,\n    #Config.bought_by_online_edge: Config.buy_online_edge\n    Config.buy_edge: Config.bought_by_edge,\n    Config.bought_by_edge: Config.buy_edge,\n}\n\n\nsampler = dgl.dataloading.MultiLayerFullNeighborSampler(2)\n\nsampler = dgl.dataloading.as_edge_prediction_sampler(\n    sampler,\n    exclude='reverse_types',\n    reverse_etypes=reverse_types,\n    negative_sampler=dgl.dataloading.negative_sampler.Uniform(1),\n    \n)\n\ntrain_dataloader = dgl.dataloading.DataLoader(\n    graph, \n    train_eid_dict, \n    sampler,\n    batch_size=1024,\n    shuffle=True,\n    drop_last=False,\n    num_workers=2\n)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:41:45.36595Z","iopub.execute_input":"2022-04-22T15:41:45.366541Z","iopub.status.idle":"2022-04-22T15:41:45.40668Z","shell.execute_reply.started":"2022-04-22T15:41:45.366498Z","shell.execute_reply":"2022-04-22T15:41:45.405902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = train(graph, X_dic, train_dataloader, model)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:41:46.041163Z","iopub.execute_input":"2022-04-22T15:41:46.04169Z","iopub.status.idle":"2022-04-22T15:52:02.684774Z","shell.execute_reply.started":"2022-04-22T15:41:46.041653Z","shell.execute_reply":"2022-04-22T15:52:02.683192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.to(\"cpu\").state_dict(), Config.model_path)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.485768Z","iopub.status.idle":"2022-04-22T15:28:04.48664Z","shell.execute_reply.started":"2022-04-22T15:28:04.486401Z","shell.execute_reply":"2022-04-22T15:28:04.486427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(Config.model_path))","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.48799Z","iopub.status.idle":"2022-04-22T15:28:04.488887Z","shell.execute_reply.started":"2022-04-22T15:28:04.488649Z","shell.execute_reply":"2022-04-22T15:28:04.488674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"emb = inference(graph, X_dic, model)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.490023Z","iopub.status.idle":"2022-04-22T15:28:04.49076Z","shell.execute_reply.started":"2022-04-22T15:28:04.490521Z","shell.execute_reply":"2022-04-22T15:28:04.490547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic[Config.article_node].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.492136Z","iopub.status.idle":"2022-04-22T15:28:04.492986Z","shell.execute_reply.started":"2022-04-22T15:28:04.492748Z","shell.execute_reply":"2022-04-22T15:28:04.492773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"emb[Config.article_node].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.494295Z","iopub.status.idle":"2022-04-22T15:28:04.495167Z","shell.execute_reply.started":"2022-04-22T15:28:04.494911Z","shell.execute_reply":"2022-04-22T15:28:04.494936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"graph.num_nodes(Config.article_node)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.496341Z","iopub.status.idle":"2022-04-22T15:28:04.497339Z","shell.execute_reply.started":"2022-04-22T15:28:04.497095Z","shell.execute_reply":"2022-04-22T15:28:04.497134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_article_node.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.498514Z","iopub.status.idle":"2022-04-22T15:28:04.505582Z","shell.execute_reply.started":"2022-04-22T15:28:04.505188Z","shell.execute_reply":"2022-04-22T15:28:04.505215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_dic[Config.customer_node].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.506979Z","iopub.status.idle":"2022-04-22T15:28:04.507397Z","shell.execute_reply.started":"2022-04-22T15:28:04.507183Z","shell.execute_reply":"2022-04-22T15:28:04.507204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"graph.num_nodes(Config.customer_node)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.509072Z","iopub.status.idle":"2022-04-22T15:28:04.509488Z","shell.execute_reply.started":"2022-04-22T15:28:04.509273Z","shell.execute_reply":"2022-04-22T15:28:04.509296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"emb[Config.customer_node].shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.510958Z","iopub.status.idle":"2022-04-22T15:28:04.511858Z","shell.execute_reply.started":"2022-04-22T15:28:04.511488Z","shell.execute_reply":"2022-04-22T15:28:04.511514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer_node.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.513926Z","iopub.status.idle":"2022-04-22T15:28:04.514856Z","shell.execute_reply.started":"2022-04-22T15:28:04.51452Z","shell.execute_reply":"2022-04-22T15:28:04.514545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_emb_dataframe(emb, df_article_node, df_customer_node):\n    df_article_emb = pd.DataFrame(emb[Config.article_node].numpy())\n    df_article_emb = pd.concat([df_article_node, df_article_emb], axis=1)\n\n    df_customer_emb = pd.DataFrame(emb[Config.customer_node].numpy())\n    df_customer_emb = pd.concat([df_customer_node, df_customer_emb], axis=1)\n\n    return df_article_emb, df_customer_emb","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.516418Z","iopub.status.idle":"2022-04-22T15:28:04.517392Z","shell.execute_reply.started":"2022-04-22T15:28:04.517106Z","shell.execute_reply":"2022-04-22T15:28:04.517145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_article_emb, df_customer_emb = create_emb_dataframe(emb, df_article_node, df_customer_node)\n\ndf_article_emb.to_pickle(\"article_emb.pkl\")\ndf_customer_emb.to_pickle(\"customer_emb.pkl\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.518687Z","iopub.status.idle":"2022-04-22T15:28:04.519661Z","shell.execute_reply.started":"2022-04-22T15:28:04.519384Z","shell.execute_reply":"2022-04-22T15:28:04.519409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_article_emb.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.520986Z","iopub.status.idle":"2022-04-22T15:28:04.521919Z","shell.execute_reply.started":"2022-04-22T15:28:04.521681Z","shell.execute_reply":"2022-04-22T15:28:04.521706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer_emb.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.523242Z","iopub.status.idle":"2022-04-22T15:28:04.524224Z","shell.execute_reply.started":"2022-04-22T15:28:04.523881Z","shell.execute_reply":"2022-04-22T15:28:04.523967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer_emb.isna().any()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.525426Z","iopub.status.idle":"2022-04-22T15:28:04.52624Z","shell.execute_reply.started":"2022-04-22T15:28:04.52599Z","shell.execute_reply":"2022-04-22T15:28:04.526016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class NearestNeighborSearch:\n    \n#     def __init__(self, n_dim, seed):\n#         self.t = AnnoyIndex(n_dim, 'angular')  \n#         self.t.set_seed(seed)\n\n#     def create_nearest_neighbor_search(self, df_emb, emb_col, n_trees):                \n#         self.df_emb = df_emb\n\n#         for i, v in tqdm.tqdm(enumerate(df_emb[emb_col].values), total=len(df_emb)):\n#             self.t.add_item(i, v)\n\n#         self.t.build(n_trees)\n\n#     def get_nerest_negihbor(self, x: np.array, n: int, target_col=\"article_id\"):\n#         nn_index_list = self.t.get_nns_by_vector(x, n)\n#         return self.df_emb.iloc[nn_index_list, :][target_col].tolist()\n\n\n\n\n# def create_submission(df_submission, df_article_emb, df_customer_emb, n_trees=10):\n#     nns = NearestNeighborSearch(Config.out_feat_dim, seed=Config.seed)\n#     nns.create_nearest_neighbor_search(df_article_emb, list(range(Config.out_feat_dim)), n_trees)\n\n#     for customer_id, emb in tqdm.tqdm(zip(df_customer_emb[\"customer_id\"], df_customer_emb.loc[:, list(range(Config.out_feat_dim))].values), total=len(df_customer_emb)):        \n#         nn_article_list = nns.get_nerest_negihbor(emb, 12)        \n#         df_submission.loc[df_submission[\"customer_id\"] == customer_id, \"prediciton\"] = \" \".join([str(x) for x in nn_article_list])\n    \n#     return df_submission\n\n\n# use ckdtree for multi processing        \n\n#from joblib import wrap_non_picklable_objects\n\n#@wrap_non_picklable_objects\ndef _task(customer_id, emb, df_article_id, annoy_path, n_dim):\n    u = AnnoyIndex(n_dim, 'dot')\n    u.load(annoy_path)     \n    nn_index_list = u.get_nns_by_vector(emb, 12)\n    nn_article_list = df_article_id.iloc[nn_index_list, :][\"article_id\"].tolist()\n    return (customer_id, \" \".join([str(x) for x in nn_article_list]))\n\ndef create_submission_mp(df_article_emb, df_customer_emb, annoy_path=\"h_and_m.ann\", n_trees=10):\n    \n    t = AnnoyIndex(Config.out_feat_dim, 'dot')  \n    t.set_seed(Config.seed)\n\n    emb_col = list(range(Config.out_feat_dim))    \n    for i, v in tqdm.tqdm(enumerate(df_article_emb[emb_col].values), total=len(df_article_emb)):\n         t.add_item(i, v)\n\n    t.build(n_trees, n_jobs=1)\n    t.save(annoy_path)\n    \n    df_article_id = df_article_emb[[\"article_id\"]]\n\n    #result_list = [_task(customer_id, emb, df_article_id, annoy_path, Config.out_feat_dim) \n    #                for customer_id, emb in zip(df_customer_emb.head(10)[\"customer_id\"].tolist(), df_customer_emb.head(10).loc[:, emb_col].values)]\n    \n    \n    # https://github.com/spotify/annoy/issues/499\n    # https://github.com/pavlin-policar/openTSNE/blob/872e8df89d7700bc650e1c2b40a41c0c5a9c1a54/openTSNE/nearest_neighbors.py#L264-L268\n    result_list = joblib.Parallel(n_jobs=-1, verbose=1, require=\"sharedmem\")(\n        joblib.delayed(_task)(customer_id, emb, df_article_id, annoy_path, Config.out_feat_dim) \n        for customer_id, emb in zip(df_customer_emb[\"customer_id\"].tolist(), df_customer_emb.loc[:, emb_col].values)\n    )\n\n    customer_id_list, prediction_list = [], []\n    for customer_id, prediction in tqdm.tqdm(result_list):\n        customer_id_list.append(customer_id)\n        prediction_list.append(prediction)\n\n    df_pred = pd.DataFrame(\n        {\"customer_id\": customer_id_list, \"prediction_2\": prediction_list}\n    )\n\n    return df_pred\n\n#def create_submission()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.5279Z","iopub.status.idle":"2022-04-22T15:28:04.528712Z","shell.execute_reply.started":"2022-04-22T15:28:04.528371Z","shell.execute_reply":"2022-04-22T15:28:04.528396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndf_article_emb = pd.read_pickle(\"article_emb.pkl\")\ndf_customer_emb = pd.read_pickle(\"customer_emb.pkl\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.530259Z","iopub.status.idle":"2022-04-22T15:28:04.53124Z","shell.execute_reply.started":"2022-04-22T15:28:04.530987Z","shell.execute_reply":"2022-04-22T15:28:04.531012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_customer_emb.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.532603Z","iopub.status.idle":"2022-04-22T15:28:04.533623Z","shell.execute_reply.started":"2022-04-22T15:28:04.533349Z","shell.execute_reply":"2022-04-22T15:28:04.533375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_prediction = create_submission_mp(df_article_emb, df_customer_emb)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.53494Z","iopub.status.idle":"2022-04-22T15:28:04.535753Z","shell.execute_reply.started":"2022-04-22T15:28:04.535395Z","shell.execute_reply":"2022-04-22T15:28:04.535419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_prediction.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.537295Z","iopub.status.idle":"2022-04-22T15:28:04.538233Z","shell.execute_reply.started":"2022-04-22T15:28:04.537982Z","shell.execute_reply":"2022-04-22T15:28:04.538008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_prediction.head(1).T","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.539635Z","iopub.status.idle":"2022-04-22T15:28:04.540615Z","shell.execute_reply.started":"2022-04-22T15:28:04.540378Z","shell.execute_reply":"2022-04-22T15:28:04.540403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission = df_submission.merge(df_prediction, on=\"customer_id\", how=\"left\")","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.541755Z","iopub.status.idle":"2022-04-22T15:28:04.542629Z","shell.execute_reply.started":"2022-04-22T15:28:04.542365Z","shell.execute_reply":"2022-04-22T15:28:04.542395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.543714Z","iopub.status.idle":"2022-04-22T15:28:04.54426Z","shell.execute_reply.started":"2022-04-22T15:28:04.544023Z","shell.execute_reply":"2022-04-22T15:28:04.544047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission[_df_submission.prediction_2.isna()].shape, _df_submission.shape","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.545546Z","iopub.status.idle":"2022-04-22T15:28:04.546086Z","shell.execute_reply.started":"2022-04-22T15:28:04.54586Z","shell.execute_reply":"2022-04-22T15:28:04.545884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission[\"prediction\"] = _df_submission.swifter.apply(lambda row: row[2] if row[2] is not np.NaN else row[1], axis=1)\n_df_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.547163Z","iopub.status.idle":"2022-04-22T15:28:04.547689Z","shell.execute_reply.started":"2022-04-22T15:28:04.547465Z","shell.execute_reply":"2022-04-22T15:28:04.547489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission = _df_submission[[\"customer_id\", \"prediction\"]]\n_df_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.548726Z","iopub.status.idle":"2022-04-22T15:28:04.549288Z","shell.execute_reply.started":"2022-04-22T15:28:04.549045Z","shell.execute_reply":"2022-04-22T15:28:04.54907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_df_submission.to_csv(\"submission.csv\",index=None)","metadata":{"execution":{"iopub.status.busy":"2022-04-22T15:28:04.550359Z","iopub.status.idle":"2022-04-22T15:28:04.550912Z","shell.execute_reply.started":"2022-04-22T15:28:04.550676Z","shell.execute_reply":"2022-04-22T15:28:04.5507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}