{"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":"# GraphSAGE","metadata":{}},{"cell_type":"markdown","source":"GraphSAGE is a graph neural network (GNN) architecture that can be used to perform various machine learning tasks on graph-structured data. It is designed to be efficient and scalable, and is based on the idea of using \"sampling and aggregation\" to learn node embeddings.\n\nIn GraphSAGE, a node embedding is learned by first sampling a set of neighboring nodes for the target node, and then using those sampled nodes to compute an aggregation of their feature vectors. This aggregation is then used as the node's embedding, which can then be used in downstream machine learning tasks such as classification or regression\n","metadata":{}},{"cell_type":"markdown","source":"Put it simply, a node is a non linear transformation of it's neighbors representations. It can be illustrated by the image below ( from the GraphSAGE Paper, see the [reference](https://arxiv.org/pdf/1706.02216.pdf )\n\n![](https://user-images.githubusercontent.com/55285736/207429302-f76ceff1-675f-4ca1-88f2-6a68ba4595f6.png)","metadata":{}},{"cell_type":"markdown","source":"We have a large choice for aggregations operators, we can use mean, max, even LSTM operators.","metadata":{}},{"cell_type":"markdown","source":"### Neighborhood Sampling","metadata":{}},{"cell_type":"markdown","source":"GraphSAGE uses a 2-hop neighborhood to generates the embedding of our nodes. Which means that the mebedding of a node will be dependent of it's direct neighbors and the neighbors of each neighbors and ignore the rest of the neighbors. For this, we'll use the `NeighborSampler` of `pytorch-geometric` to use mini-batch and can be even use with GPU !\n\nBatchs are composed of **computational graphs** ; node with its appropriated neighbors.\n\n We use sampling strategy to sample at most **H** neighbors at aech hop to avoid computation complexity. But also because if a neighbors is a hub node ( a high degree node ) the complexity can explode easily. The image below ( thanks to stanford university videos ) summarize how the sampling strategy works, where each layer (stage) represents a neighbors hop.\n \n![](https://user-images.githubusercontent.com/55285736/207427440-1597c762-040e-498a-8df2-b530dec20d34.png) \n\n\n### Advantages:\n\n-  Increase computational efficiency.\n- Acts as a dropout / regularization strategy, because we ignore some neighbors.","metadata":{}},{"cell_type":"markdown","source":"\n## Training\n\nInstead of training a distinct embedding vector for each node, we train a set of aggregator functions that learn to aggregate feature information from a node’s local neighborhood. Each aggregator function aggregates information from a different number of hops, or search depth, away from a given node. For our unsupervised task of learning node embeddings, the authors design an unsupervised loss function that allows GraphSAGE to be trained without task-specific supervision.","metadata":{}},{"cell_type":"markdown","source":"### Stochastic training of GraphSAGE","metadata":{}},{"cell_type":"markdown","source":"1.  Randomly sample M nodes for N nodes.\n2. For each sampled node v\n    *  Get k-hop neighborhood using sampling strategy to reduce the complexity, and construct the computation graph.\n    *  Use the above to generate the embedding of our node of interest.\n3. Compute the reconstruction loss over M nodes ( in our case we use an autoencoder to reconstruct the adjacency matrix ).\n4. Perform the gradient update using the optimizer of your choice.","metadata":{}},{"cell_type":"markdown","source":"For further reading:\n\n- GraphSAGE Paper : https://arxiv.org/pdf/1706.02216.pdf.\n\n-  Radek local validation : https://www.kaggle.com/competitions/otto-recommender-system/discussion/364991\n\n- Radek Datasets for local validation : https://www.kaggle.com/datasets/radek1/otto-train-and-test-data-for-local-validation.\n\n- Stanford Video of Leskovec. PhD explaining Node2Vec : https://www.youtube.com/watch?v=LLUxwHc7O4A&t=213s","metadata":{"papermill":{"duration":0.009449,"end_time":"2022-12-07T20:32:30.856498","exception":false,"start_time":"2022-12-07T20:32:30.847049","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"**If you like this notebook, please upvote! That will be of great help to me and will allow me to share more materials like this with you. Thank you! 😊**","metadata":{"papermill":{"duration":0.007486,"end_time":"2022-12-07T20:32:30.872140","exception":false,"start_time":"2022-12-07T20:32:30.864654","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Libraries installation\n","metadata":{"papermill":{"duration":0.007363,"end_time":"2022-12-07T20:32:30.887313","exception":false,"start_time":"2022-12-07T20:32:30.879950","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"We will need a couple of libraries that do not come preinstalled on the Kaggle VM. Let's install them here.","metadata":{"papermill":{"duration":0.007384,"end_time":"2022-12-07T20:32:30.902584","exception":false,"start_time":"2022-12-07T20:32:30.895200","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install polars\n\n## GPU Version\n\n!pip install torch==1.12.0+cu116 -f https://download.pytorch.org/whl/cu116/torch_stable.html  \n!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric -f https://data.pyg.org/whl/torch-1.12.0+cu116.html\n!pip install pyg-lib -f https://data.pyg.org/whl/torch-1.12.0+cu116.html   \n## CPU Version\n\n#!pip install torch==1.12.0+cpu -f https://download.pytorch.org/whl/torch_stable.html\n#!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric -f https://data.pyg.org/whl/torch-1.12.0+cpu.html","metadata":{"papermill":{"duration":95.812411,"end_time":"2022-12-07T20:34:06.722803","exception":false,"start_time":"2022-12-07T20:32:30.910392","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:07:43.335918Z","iopub.execute_input":"2023-01-02T19:07:43.336331Z","iopub.status.idle":"2023-01-02T19:10:52.375966Z","shell.execute_reply.started":"2023-01-02T19:07:43.336245Z","shell.execute_reply":"2023-01-02T19:10:52.374804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport gc\nimport pandas as pd\nimport torch\nfrom torch import nn\nfrom torch_geometric.nn import GCNConv,SAGEConv,GAE\nfrom torch_geometric.data import Data\nfrom tqdm import tqdm\nimport polars as pl\nfrom collections import defaultdict\nfrom sklearn.preprocessing import StandardScaler","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":7.212061,"end_time":"2022-12-07T20:34:13.960713","exception":false,"start_time":"2022-12-07T20:34:06.748652","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:10:52.378502Z","iopub.execute_input":"2023-01-02T19:10:52.380137Z","iopub.status.idle":"2023-01-02T19:10:57.136568Z","shell.execute_reply.started":"2023-01-02T19:10:52.380089Z","shell.execute_reply":"2023-01-02T19:10:57.135528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)","metadata":{"papermill":{"duration":0.040611,"end_time":"2022-12-07T20:34:14.028594","exception":false,"start_time":"2022-12-07T20:34:13.987983","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:10:57.138018Z","iopub.execute_input":"2023-01-02T19:10:57.138956Z","iopub.status.idle":"2023-01-02T19:10:57.147577Z","shell.execute_reply.started":"2023-01-02T19:10:57.138916Z","shell.execute_reply":"2023-01-02T19:10:57.145637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build Graph","metadata":{"papermill":{"duration":0.025253,"end_time":"2022-12-07T20:34:14.080581","exception":false,"start_time":"2022-12-07T20:34:14.055328","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Nodes Features","metadata":{}},{"cell_type":"markdown","source":"Here I use features extracted by Node2Vec technique that uses Biase Random Walk to extract node embeddings. ( take a look at  the notebook [node2vec-a-biased-random-walk-approach](https://www.kaggle.com/code/rayanaay/node2vec-a-biased-random-walk-approach) )","metadata":{}},{"cell_type":"code","source":"node2vec_embeddings  = np.load(\"/kaggle/input/gnn-outputs/node2vec_embeddings.npy\")","metadata":{"papermill":{"duration":2.73573,"end_time":"2022-12-07T20:34:16.842474","exception":false,"start_time":"2022-12-07T20:34:14.106744","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:10:57.150302Z","iopub.execute_input":"2023-01-02T19:10:57.151163Z","iopub.status.idle":"2023-01-02T19:10:59.406171Z","shell.execute_reply.started":"2023-01-02T19:10:57.151127Z","shell.execute_reply":"2023-01-02T19:10:59.405110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aid_features = pl.read_parquet(\"/kaggle/input/graph-edges-features-agg/aid_features.parquet\").fill_null(0)\naid_features_agg = pl.read_parquet(\"/kaggle/input/graph-edges-features-agg/aids_features_aggregation.parquet\").fill_null(0)\naid_features_all = aid_features.join(aid_features_agg,on=\"aid\",how=\"inner\").drop(\"aid\").to_numpy()","metadata":{"papermill":{"duration":10.505369,"end_time":"2022-12-07T20:34:27.374335","exception":false,"start_time":"2022-12-07T20:34:16.868966","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:10:59.407931Z","iopub.execute_input":"2023-01-02T19:10:59.408406Z","iopub.status.idle":"2023-01-02T19:11:07.004375Z","shell.execute_reply.started":"2023-01-02T19:10:59.408362Z","shell.execute_reply":"2023-01-02T19:11:07.003639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features_and_embeddings = np.concatenate((node2vec_embeddings,aid_features_all),axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:11:07.005669Z","iopub.execute_input":"2023-01-02T19:11:07.006317Z","iopub.status.idle":"2023-01-02T19:11:08.496997Z","shell.execute_reply.started":"2023-01-02T19:11:07.006279Z","shell.execute_reply":"2023-01-02T19:11:08.495813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"std = StandardScaler().fit(features_and_embeddings)\n\nnode_features_scaled = std.transform(features_and_embeddings)","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:11:08.498665Z","iopub.execute_input":"2023-01-02T19:11:08.499050Z","iopub.status.idle":"2023-01-02T19:11:12.916739Z","shell.execute_reply.started":"2023-01-02T19:11:08.499010Z","shell.execute_reply":"2023-01-02T19:11:12.915737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del features_and_embeddings,aid_features_all,aid_features_agg,aid_features\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:11:12.918325Z","iopub.execute_input":"2023-01-02T19:11:12.919062Z","iopub.status.idle":"2023-01-02T19:11:13.187769Z","shell.execute_reply.started":"2023-01-02T19:11:12.919017Z","shell.execute_reply":"2023-01-02T19:11:13.186860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del node2vec_embeddings\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:11:13.189025Z","iopub.execute_input":"2023-01-02T19:11:13.189371Z","iopub.status.idle":"2023-01-02T19:11:13.330898Z","shell.execute_reply.started":"2023-01-02T19:11:13.189337Z","shell.execute_reply":"2023-01-02T19:11:13.329809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Interactions between aids","metadata":{}},{"cell_type":"code","source":"#edges_tensor = torch.load(\"/kaggle/input/graph-edges-features-agg/otto-graph-edges.pt\") # extracts all edges even duplicated one to biased the sampling !!\n#edges_tensor = torch.load(\"/kaggle/input/graph-edges-features-agg/all_edges_without_duplicates_90m.pt\") \nedges_tensor = torch.load(\"/kaggle/input/otto-top-neighbors/general_edges/top_k10_neighbors.pt\")","metadata":{"papermill":{"duration":14.321597,"end_time":"2022-12-07T20:34:41.940193","exception":false,"start_time":"2022-12-07T20:34:27.618596","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:12:10.678915Z","iopub.execute_input":"2023-01-02T19:12:10.679270Z","iopub.status.idle":"2023-01-02T19:12:13.275078Z","shell.execute_reply.started":"2023-01-02T19:12:10.679240Z","shell.execute_reply":"2023-01-02T19:12:13.273809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = Data(x=torch.tensor(node_features_scaled),\n            edge_index=edges_tensor,\n             )\ndata.n_id = torch.arange(data.num_nodes)","metadata":{"papermill":{"duration":0.496693,"end_time":"2022-12-07T20:34:45.040468","exception":false,"start_time":"2022-12-07T20:34:44.543775","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:12:13.477161Z","iopub.execute_input":"2023-01-02T19:12:13.477536Z","iopub.status.idle":"2023-01-02T19:12:14.393667Z","shell.execute_reply.started":"2023-01-02T19:12:13.477505Z","shell.execute_reply":"2023-01-02T19:12:14.392479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del edges_tensor,node_features_scaled#,features_and_embeddings,aid_features_all,aid_features_agg,aid_features\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:12:14.395907Z","iopub.execute_input":"2023-01-02T19:12:14.396780Z","iopub.status.idle":"2023-01-02T19:12:14.558438Z","shell.execute_reply.started":"2023-01-02T19:12:14.396739Z","shell.execute_reply":"2023-01-02T19:12:14.557246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:12:14.560202Z","iopub.execute_input":"2023-01-02T19:12:14.560596Z","iopub.status.idle":"2023-01-02T19:12:14.571720Z","shell.execute_reply.started":"2023-01-02T19:12:14.560556Z","shell.execute_reply":"2023-01-02T19:12:14.570788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Neighbor Loader","metadata":{"papermill":{"duration":0.025468,"end_time":"2022-12-07T20:34:45.096043","exception":false,"start_time":"2022-12-07T20:34:45.070575","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from torch_geometric.loader import NeighborLoader\n\ngSAGE_loader = NeighborLoader(\n    data,\n    # Sample 30 neighbors for each node for 2 iterations\n    num_neighbors=[10,10],#, neg_sample_size\n    # Use a batch size of 128 for sampling training nodes\n    batch_size=512,\n)","metadata":{"papermill":{"duration":16.615911,"end_time":"2022-12-07T20:35:01.738474","exception":false,"start_time":"2022-12-07T20:34:45.122563","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:12:14.575668Z","iopub.execute_input":"2023-01-02T19:12:14.576087Z","iopub.status.idle":"2023-01-02T19:12:16.819857Z","shell.execute_reply.started":"2023-01-02T19:12:14.576049Z","shell.execute_reply":"2023-01-02T19:12:16.818648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GraphSAGE","metadata":{"papermill":{"duration":0.02671,"end_time":"2022-12-07T20:35:01.792625","exception":false,"start_time":"2022-12-07T20:35:01.765915","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class GCNEncoder(torch.nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(GCNEncoder, self).__init__()\n        self.conv1 = SAGEConv(in_channels, 60, aggr=\"mean\",project=False) # cached only for transductive learning\n        self.conv2 = SAGEConv(60, 32, aggr=\"sum\",project=False) # cached only for transductive learning\n\n    def forward(self, x, edge_index):\n        x = torch.nn.ELU()(self.conv1(x, edge_index))\n\n        return  torch.nn.ELU()(self.conv2(x, edge_index))","metadata":{"papermill":{"duration":0.040001,"end_time":"2022-12-07T20:35:01.860869","exception":false,"start_time":"2022-12-07T20:35:01.820868","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:12:16.821446Z","iopub.execute_input":"2023-01-02T19:12:16.821876Z","iopub.status.idle":"2023-01-02T19:12:16.833291Z","shell.execute_reply.started":"2023-01-02T19:12:16.821808Z","shell.execute_reply":"2023-01-02T19:12:16.832133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out_channels = 32\nnum_features = data.x.shape[1]\nmodel = GAE(GCNEncoder(num_features, out_channels))\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = model.to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.005,weight_decay=1e-5)","metadata":{"papermill":{"duration":0.043318,"end_time":"2022-12-07T20:35:01.929978","exception":false,"start_time":"2022-12-07T20:35:01.886660","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:12:16.836074Z","iopub.execute_input":"2023-01-02T19:12:16.836940Z","iopub.status.idle":"2023-01-02T19:12:19.118989Z","shell.execute_reply.started":"2023-01-02T19:12:16.836897Z","shell.execute_reply":"2023-01-02T19:12:19.117981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" def train(loader):\n        total_loss = 0\n        for subgraph in tqdm(loader):\n            optimizer.zero_grad()\n            z = model.encode(subgraph.x.float().to(device),subgraph.edge_index.to(device))\n            loss = model.recon_loss(z, pos_edge_index=subgraph.edge_index.to(device))\n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n        return total_loss / len(loader), model\n","metadata":{"papermill":{"duration":0.04705,"end_time":"2022-12-07T20:35:02.003591","exception":false,"start_time":"2022-12-07T20:35:01.956541","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-14T10:41:02.377559Z","iopub.execute_input":"2022-12-14T10:41:02.379618Z","iopub.status.idle":"2022-12-14T10:41:02.385617Z","shell.execute_reply.started":"2022-12-14T10:41:02.379566Z","shell.execute_reply":"2022-12-14T10:41:02.384485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''for epoch in range(0,10):\n    \n    loss,model = train(gSAGE_loader)\n    print(f'Epoch: {epoch:02d}, Loss: {loss:.4f}')\ntorch.save(model,\"graphSage_model\")'''","metadata":{"papermill":{"duration":9150.366896,"end_time":"2022-12-07T23:07:32.408077","exception":false,"start_time":"2022-12-07T20:35:02.041181","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:28:25.868120Z","iopub.execute_input":"2023-01-02T19:28:25.868979Z","iopub.status.idle":"2023-01-02T19:28:25.899614Z","shell.execute_reply.started":"2023-01-02T19:28:25.868865Z","shell.execute_reply":"2023-01-02T19:28:25.898720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Encoding by mini-batch","metadata":{}},{"cell_type":"code","source":"model = torch.load(\"/kaggle/input/ottographsagemodel0523lb/graphSage_model\")","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:12:19.120981Z","iopub.execute_input":"2023-01-02T19:12:19.121351Z","iopub.status.idle":"2023-01-02T19:12:19.146819Z","shell.execute_reply.started":"2023-01-02T19:12:19.121315Z","shell.execute_reply":"2023-01-02T19:12:19.145895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np_embeddings = np.zeros((data.num_nodes,32))\nfor subgraph in tqdm(gSAGE_loader):\n    np_embeddings[subgraph.input_id] = model.encoder(subgraph.x.float().to(device),subgraph.edge_index.to(device)).cpu().detach().numpy()[:len(subgraph.input_id)]","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:12:19.376404Z","iopub.execute_input":"2023-01-02T19:12:19.376691Z","iopub.status.idle":"2023-01-02T19:13:53.317555Z","shell.execute_reply.started":"2023-01-02T19:12:19.376665Z","shell.execute_reply":"2023-01-02T19:13:53.316556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del data, model\ngc.collect()\ndel gSAGE_loader\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:13:53.319783Z","iopub.execute_input":"2023-01-02T19:13:53.320465Z","iopub.status.idle":"2023-01-02T19:13:53.632264Z","shell.execute_reply.started":"2023-01-02T19:13:53.320426Z","shell.execute_reply":"2023-01-02T19:13:53.631118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Check the sampling encoding\n'''\nembeddings_all_batchs = model.encode(data.x.float().to(device),data.edge_index.to(device)).detach().cpu().numpy()[0]\nembeddings_all_batchs[0]\nnp_embeddings[0]\n'''","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:13:53.633621Z","iopub.execute_input":"2023-01-02T19:13:53.634537Z","iopub.status.idle":"2023-01-02T19:13:53.646292Z","shell.execute_reply.started":"2023-01-02T19:13:53.634473Z","shell.execute_reply":"2023-01-02T19:13:53.644993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"For fast search, we use the `annoy` library that allows approximative neighbors search.","metadata":{"papermill":{"duration":4.922109,"end_time":"2022-12-07T23:11:54.471450","exception":false,"start_time":"2022-12-07T23:11:49.549341","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\nfrom annoy import AnnoyIndex\n\nindex = AnnoyIndex(32, 'angular')\n\nfor idx,idx_embedding in enumerate(np_embeddings):\n    index.add_item(idx, idx_embedding)\n    \nindex.build(10)\n","metadata":{"papermill":{"duration":41.168882,"end_time":"2022-12-07T23:12:40.745855","exception":false,"start_time":"2022-12-07T23:11:59.576973","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:13:53.649265Z","iopub.execute_input":"2023-01-02T19:13:53.649886Z","iopub.status.idle":"2023-01-02T19:14:26.711624Z","shell.execute_reply.started":"2023-01-02T19:13:53.649827Z","shell.execute_reply":"2023-01-02T19:14:26.706369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''del np_embeddings\ngc.collect()'''","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:14:26.713907Z","iopub.execute_input":"2023-01-02T19:14:26.714314Z","iopub.status.idle":"2023-01-02T19:14:26.723173Z","shell.execute_reply.started":"2023-01-02T19:14:26.714269Z","shell.execute_reply":"2023-01-02T19:14:26.722227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Validation / Inference","metadata":{"papermill":{"duration":5.159053,"end_time":"2022-12-07T23:12:50.931895","exception":false,"start_time":"2022-12-07T23:12:45.772842","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\ndef evaluate(path,mode=\"validation\",n_neighbors=20):\n\n\n    test = pl.read_parquet(path)\n\n    session_types = ['clicks', 'carts', 'orders']\n    test_session_AIDs = test.to_pandas().reset_index(drop=True).groupby('session')['aid'].apply(list)\n    test_session_types = test.to_pandas().reset_index(drop=True).groupby('session')['type'].apply(list)\n\n    del test\n    gc.collect()\n    labels = []\n\n    type_weight_multipliers = {0: 1, 1: 6, 2: 3}\n\n    for AIDs, types in zip(test_session_AIDs, test_session_types):\n        if len(AIDs) >= 20:\n                # if we have enough aids (over equals 20) we don't need to look for candidates! we just use the old logic\n            weights=np.logspace(0.1,1,len(AIDs),base=2, endpoint=True)-1\n            aids_temp=defaultdict(lambda: 0)\n            for aid,w,t in zip(AIDs,weights,types): \n                aids_temp[aid]+= w * type_weight_multipliers[t]\n\n            sorted_aids=[k for k, v in sorted(aids_temp.items(), key=lambda item: -item[1])]\n            labels.append(sorted_aids[:20])\n        else:\n            # here we don't have 20 aids to output -- we will use word2vec embeddings to generate candidates!\n            AIDs = list(dict.fromkeys(AIDs[::-1]))\n\n            # let's grab the most recent aid\n            most_recent_aid = AIDs[0]\n\n            # and look for some neighbors!\n            nns = [i for i in index.get_nns_by_item(most_recent_aid, n_neighbors)[1:]]\n\n\n            labels.append((AIDs+nns)[:n_neighbors])\n\n    labels_as_strings = [' '.join([str(l) for l in lls]) for lls in labels]\n\n    predictions = pd.DataFrame(data={'session_type': test_session_AIDs.index, 'labels': labels_as_strings})\n\n    prediction_dfs = []\n\n    for st in session_types:\n        modified_predictions = predictions.copy()\n        modified_predictions.session_type = modified_predictions.session_type.astype('str') + f'_{st}'\n        prediction_dfs.append(modified_predictions)\n\n    sub = pd.concat(prediction_dfs).reset_index(drop=True)\n    \n    del prediction_dfs, predictions,labels_as_strings, labels, test_session_types,test_session_AIDs\n    gc.collect()\n    if mode==\"test\":\n        sub.to_csv(\"submission.csv\",index=False)\n        return sub\n    else:\n\n        sub['labels_2'] = sub['labels'].apply(lambda x : [int(s) for s in x.split(' ')])\n        submission = pd.DataFrame()\n        submission['session'] = sub.session_type.apply(lambda x: int(x.split('_')[0]))\n        submission['type'] = sub.session_type.apply(lambda x: x.split('_')[1])\n        submission['labels'] = sub.labels_2.apply(lambda x : [item for item in x[:] ]) #.apply(lambda x: [int(i) for i in x.split(',')[:20]])\n        test_labels = pd.read_parquet('/kaggle/input/otto-train-and-test-data-for-local-validation/test_labels.parquet')\n        test_labels = test_labels.merge(submission, how='left', on=['session', 'type'])\n        del sub,submission\n        gc.collect()\n        gc.collect()\n        test_labels['hits'] = test_labels.apply(lambda df: len(set(df.ground_truth).intersection(set(df.labels))), axis=1)\n        test_labels['gt_count'] = test_labels.ground_truth.str.len().clip(0,20)\n        recall_per_type = test_labels.groupby(['type'])['hits'].sum() / test_labels.groupby(['type'])['gt_count'].sum() \n        score = (recall_per_type * pd.Series({'clicks': 0.1, 'carts': 0.30, 'orders': 0.60})).sum()\n\n        return score,recall_per_type\n","metadata":{"papermill":{"duration":5.373786,"end_time":"2022-12-07T23:13:01.444215","exception":false,"start_time":"2022-12-07T23:12:56.070429","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:21:27.649319Z","iopub.execute_input":"2023-01-02T19:21:27.649686Z","iopub.status.idle":"2023-01-02T19:21:27.670147Z","shell.execute_reply.started":"2023-01-02T19:21:27.649653Z","shell.execute_reply":"2023-01-02T19:21:27.668879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = \"/kaggle/input/otto-train-and-test-data-for-local-validation/test.parquet\"\nvalidation_score,r = evaluate(path,mode=\"validation\",n_neighbors=20)\nprint(validation_score)","metadata":{"papermill":{"duration":448.092958,"end_time":"2022-12-07T23:20:34.541978","exception":false,"start_time":"2022-12-07T23:13:06.449020","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-01-02T19:21:27.956408Z","iopub.execute_input":"2023-01-02T19:21:27.957418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(r)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save(\"graphsage_embeddings.npy\",np_embeddings)","metadata":{"execution":{"iopub.status.busy":"2023-01-02T19:20:37.604681Z","iopub.execute_input":"2023-01-02T19:20:37.605056Z","iopub.status.idle":"2023-01-02T19:20:37.974743Z","shell.execute_reply.started":"2023-01-02T19:20:37.605025Z","shell.execute_reply":"2023-01-02T19:20:37.973710Z"},"trusted":true},"execution_count":null,"outputs":[]}]}