{"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":"!pip install polars\n\n## GPU Version\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\n## CPU Version\n#!pip install pyg-lib -f https://data.pyg.org/whl/torch-1.12.0+cpu.html   \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":{"id":"AsWUrVh7vZtM","execution":{"iopub.status.busy":"2022-12-20T22:45:37.388101Z","iopub.execute_input":"2022-12-20T22:45:37.389144Z","iopub.status.idle":"2022-12-20T22:49:16.277567Z","shell.execute_reply.started":"2022-12-20T22:45:37.389043Z","shell.execute_reply":"2022-12-20T22:49:16.276319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport gc\nimport numpy as np\nfrom sklearn.preprocessing import minmax_scale\nimport polars as pl\nfrom torch import Tensor\nimport pyg_lib\nfrom tqdm import tqdm\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nimport torch_cluster\nfrom torch_geometric.nn import nearest\nimport torch.nn.functional as F\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: '{device}'\")\n\nimport pandas as pd, numpy as np\nfrom tqdm import tqdm\n","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:16.282044Z","iopub.execute_input":"2022-12-20T22:49:16.282387Z","iopub.status.idle":"2022-12-20T22:49:21.341649Z","shell.execute_reply.started":"2022-12-20T22:49:16.282352Z","shell.execute_reply":"2022-12-20T22:49:21.339709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Link Prediction on Otto Dataset\n\nThe link prediction task then tries to predict missing ratings, and can, for example, be used to recommend sessions new aid.","metadata":{"id":"vit8xKCiXAue"}},{"cell_type":"code","source":"test_df = pl.read_parquet(\"/kaggle/input/otto-full-optimized-memory-footprint/test.parquet\")\nvalid_df = pl.read_parquet(\"/kaggle/input/otto-train-and-test-data-for-local-validation/test.parquet\")","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:21.344185Z","iopub.execute_input":"2022-12-20T22:49:21.345276Z","iopub.status.idle":"2022-12-20T22:49:22.798684Z","shell.execute_reply.started":"2022-12-20T22:49:21.345233Z","shell.execute_reply":"2022-12-20T22:49:22.797290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df['ts'].max() < test_df['ts'].min()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:22.801740Z","iopub.execute_input":"2022-12-20T22:49:22.802700Z","iopub.status.idle":"2022-12-20T22:49:22.839025Z","shell.execute_reply.started":"2022-12-20T22:49:22.802656Z","shell.execute_reply":"2022-12-20T22:49:22.838000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pl.concat([valid_df, test_df],how=\"vertical\").groupby(\"session\").tail(20).sort(\"session\")","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:22.840602Z","iopub.execute_input":"2022-12-20T22:49:22.840902Z","iopub.status.idle":"2022-12-20T22:49:28.133058Z","shell.execute_reply.started":"2022-12-20T22:49:22.840875Z","shell.execute_reply":"2022-12-20T22:49:28.132030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['session'].n_unique()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:28.134794Z","iopub.execute_input":"2022-12-20T22:49:28.135392Z","iopub.status.idle":"2022-12-20T22:49:28.613339Z","shell.execute_reply.started":"2022-12-20T22:49:28.135355Z","shell.execute_reply":"2022-12-20T22:49:28.612243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Select last 3 articles for each session","metadata":{}},{"cell_type":"code","source":"del valid_df, test_df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:28.615322Z","iopub.execute_input":"2022-12-20T22:49:28.615914Z","iopub.status.idle":"2022-12-20T22:49:28.808230Z","shell.execute_reply.started":"2022-12-20T22:49:28.615875Z","shell.execute_reply":"2022-12-20T22:49:28.807274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aids_features1 = pl.read_parquet(\"/kaggle/input/graph-edges-features-agg/aids_features_aggregation.parquet\").fill_null(0)\naids_features2 = pl.read_parquet(\"/kaggle/input/graph-edges-features-agg/aid_features.parquet\").fill_null(0)\nnode2vec_embeddings = np.load(\"/kaggle/input/gnn-outputs/node2vec_embeddings.npy\")\naids_features = aids_features1.join(aids_features2,how=\"inner\",on=\"aid\")\naids_features.head()                          ","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:28.809602Z","iopub.execute_input":"2022-12-20T22:49:28.810648Z","iopub.status.idle":"2022-12-20T22:49:38.060228Z","shell.execute_reply.started":"2022-12-20T22:49:28.810609Z","shell.execute_reply":"2022-12-20T22:49:38.058413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del aids_features1,aids_features2\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:38.071480Z","iopub.execute_input":"2022-12-20T22:49:38.073944Z","iopub.status.idle":"2022-12-20T22:49:38.332672Z","shell.execute_reply.started":"2022-12-20T22:49:38.073891Z","shell.execute_reply":"2022-12-20T22:49:38.331738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"node2vec_embeddings_df = pl.DataFrame(node2vec_embeddings)\nnode2vec_embeddings_df = node2vec_embeddings_df.with_column(pl.Series(np.arange(0,len(node2vec_embeddings),1)).alias(\"aid\"))\ndel node2vec_embeddings\ngc.collect()\nnode2vec_embeddings_df = node2vec_embeddings_df.with_column(pl.col('aid').cast(pl.Int32))","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:38.336143Z","iopub.execute_input":"2022-12-20T22:49:38.336535Z","iopub.status.idle":"2022-12-20T22:49:39.378134Z","shell.execute_reply.started":"2022-12-20T22:49:38.336503Z","shell.execute_reply":"2022-12-20T22:49:39.376782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aids_f = (df.join(aids_features,on=\"aid\",how=\"left\").sort(\"aid\").unique(subset=[\"aid\"])\n         .join(node2vec_embeddings_df,on=\"aid\",how=\"left\").drop(['session','aid'])\n         )","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:49:39.379962Z","iopub.execute_input":"2022-12-20T22:49:39.380638Z","iopub.status.idle":"2022-12-20T22:50:03.380063Z","shell.execute_reply.started":"2022-12-20T22:49:39.380599Z","shell.execute_reply":"2022-12-20T22:50:03.379011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del aids_features, node2vec_embeddings_df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:50:46.369210Z","iopub.execute_input":"2022-12-20T22:50:46.370274Z","iopub.status.idle":"2022-12-20T22:50:46.634221Z","shell.execute_reply.started":"2022-12-20T22:50:46.370213Z","shell.execute_reply":"2022-12-20T22:50:46.632920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use genres as aid input features:)\naids_feat_torch = torch.from_numpy(minmax_scale(aids_f.to_numpy())).to(torch.float)                      \nassert aids_feat_torch.size() == (aids_f.shape[0], aids_f.shape[1])  # 20 genres in total.","metadata":{"id":"dd_BJMNcwJ7k","outputId":"54911dcf-687d-4d41-c1b4-be4c2ba8c4a3","execution":{"iopub.status.busy":"2022-12-20T22:51:18.315970Z","iopub.execute_input":"2022-12-20T22:51:18.316478Z","iopub.status.idle":"2022-12-20T22:51:20.871269Z","shell.execute_reply.started":"2022-12-20T22:51:18.316438Z","shell.execute_reply":"2022-12-20T22:51:20.870167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a mapping from unique session indices to range [0, num_session_nodes):\nunique_session_id = df['session'].unique()\nunique_session_id = pl.DataFrame(data={\n    'session': unique_session_id,\n    'mappedID': np.arange(0,len(unique_session_id)),\n})\n\nprint(\"Mapping of session IDs to consecutive values:\")\nprint(\"==========================================\")\n# Create a mapping from unique aid indices to range [0, num_aid_nodes):\n\nunique_aid_id = df['aid'].unique()\nunique_aid_id = pl.DataFrame(data={\n    'aid': unique_aid_id,\n    'mappedID': np.arange(0,len(unique_aid_id)),\n})\nprint(\"Mapping of aid IDs to consecutive values:\")\nprint(\"===========================================\")\nprint(unique_aid_id.head())\n\n# Perform merge to obtain the edges from sessions and aid:\nratings_session_id = df.join(unique_session_id,\n                            left_on='session', right_on='session', how='left')\nratings_session_id = torch.from_numpy(ratings_session_id['mappedID'].to_numpy())\nratings_aid_id = df.join(unique_aid_id,\n                            left_on='aid', right_on='aid', how='left')\nratings_aid_id = torch.from_numpy(ratings_aid_id['mappedID'].to_numpy())\n\n# With this, we are ready to construct our `edge_index` in COO format\n# following PyG semantics:\nedge_index_session_to_aid = torch.stack([ratings_session_id, ratings_aid_id], dim=0)\n\nprint()\nprint(\"Final edge indices pointing from sessions to aid:\")\nprint(\"=================================================\")\nprint(edge_index_session_to_aid)","metadata":{"id":"JMGYv83WzSRr","outputId":"8c613b8b-e5a6-4674-ac6c-373acfa03d1c","execution":{"iopub.status.busy":"2022-12-20T22:51:07.546710Z","iopub.execute_input":"2022-12-20T22:51:07.547142Z","iopub.status.idle":"2022-12-20T22:51:10.649156Z","shell.execute_reply.started":"2022-12-20T22:51:07.547109Z","shell.execute_reply":"2022-12-20T22:51:10.648065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"With this, we are ready to initialize our `HeteroData` object and pass the necessary information to it.\nNote that we also pass in a `node_id` vector to each node type in order to reconstruct the original node indices from sampled subgraphs.\nWe also take care of adding reverse edges to the `HeteroData` object.\nThis allows our GNN model to use both directions of the edge for message passing:","metadata":{"id":"9w9fCnjvqmd2"}},{"cell_type":"code","source":"from torch_geometric.data import HeteroData\nimport torch_geometric.transforms as T\n\ndata = HeteroData()\n# Save node indices:\ndata[\"session\"].node_id = torch.arange(len(unique_session_id))\ndata[\"aid\"].node_id = torch.arange(len(unique_aid_id))\n\n# Add the node features and edge indices:\ndata[\"aid\"].x = aids_feat_torch # TODO\ndata[\"session\", \"rates\", \"aid\"].edge_index = edge_index_session_to_aid # TODO\n\n# We also need to make sure to add the reverse edges from aid to sessions\n# in order to let a GNN be able to pass messages in both directions.\n# We can leverage the `T.ToUndirected()` transform for this from PyG:\n\n# TODO:\nimport torch_geometric.transforms as T\n\ndata = T.ToUndirected()(data)\nprint(data)\n","metadata":{"id":"I_63--974srt","outputId":"32fe707d-f9ac-4dca-8c7b-52bbb8ddfa07","execution":{"iopub.status.busy":"2022-12-20T22:51:26.589036Z","iopub.execute_input":"2022-12-20T22:51:26.589437Z","iopub.status.idle":"2022-12-20T22:51:26.736877Z","shell.execute_reply.started":"2022-12-20T22:51:26.589406Z","shell.execute_reply":"2022-12-20T22:51:26.735846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining Edge-level Training Splits\n\nSince our data is now ready-to-be-used, we can split the ratings of sessions into training, validation, and test splits.\nThis is needed in order to ensure that we leak no information about edges used during evaluation into the training phase.\nFor this, we make use of the [`transforms.RandomLinkSplit`](https://pytorch-geometric.readthedocs.io/en/latest/modules/transforms.html#torch_geometric.transforms.RandomLinkSplit) transformation from PyG.\nThis transforms randomly divides the edges in the `(\"session\", \"rates\", \"aid\")` into training, validation and test edges.\nThe `disjoint_train_ratio` parameter further separates edges in the training split into edges used for message passing (`edge_index`) and edges used for supervision (`edge_label_index`).\nNote that we also need to specify the reverse edge type `(\"aid\", \"rev_rates\", \"session\")`.\nThis allows the `RandomLinkSplit` transform to drop reverse edges accordingly to not leak any information into the training phase.","metadata":{"id":"2QGdkLAurBq9"}},{"cell_type":"code","source":"transform = T.RandomLinkSplit(\n    num_val=0.1,  # TODO\n    num_test=0.001,  # TODO\n    disjoint_train_ratio=0.3,  # TODO\n    neg_sampling_ratio=2.0,  # TODO\n    add_negative_train_samples=False,  # TODO\n    edge_types=(\"session\", \"rates\", \"aid\"),\n    rev_edge_types=(\"aid\", \"rev_rates\", \"session\"), \n)\n\ntrain_data, val_data, test_data = transform(data)\nprint(\"Training data:\")\nprint(\"==============\")\nprint(train_data)\nprint()\nprint(\"Validation data:\")\nprint(\"================\")\nprint(val_data)\n","metadata":{"id":"rwgNwoa26Eja","outputId":"12f29798-44f5-4171-e5fb-faec01945b82","execution":{"iopub.status.busy":"2022-12-20T22:51:31.394472Z","iopub.execute_input":"2022-12-20T22:51:31.394843Z","iopub.status.idle":"2022-12-20T22:51:40.480682Z","shell.execute_reply.started":"2022-12-20T22:51:31.394810Z","shell.execute_reply":"2022-12-20T22:51:40.479568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining Mini-batch Loaders\n\nWe are now ready to create a mini-batch loader that will generate subgraphs that can be used as input into our GNN.\nWhile this step is not strictly necessary for small-scale graphs, it is absolutely necessary to apply GNNs on larger graphs that do not fit onto GPU memory otherwise.\nHere, we make use of the [`loader.LinkNeighborLoader`](https://pytorch-geometric.readthedocs.io/en/latest/modules/loader.html#torch_geometric.loader.LinkNeighborLoader) which samples multiple hops from both ends of a link and creates a subgraph from it.\nHere, `edge_label_index` serves as the \"seed links\" to start sampling from.","metadata":{"id":"prKLwq6RsYoh"}},{"cell_type":"code","source":"from torch_geometric.loader import LinkNeighborLoader","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:51:40.483597Z","iopub.execute_input":"2022-12-20T22:51:40.483920Z","iopub.status.idle":"2022-12-20T22:51:40.488445Z","shell.execute_reply.started":"2022-12-20T22:51:40.483892Z","shell.execute_reply":"2022-12-20T22:51:40.487244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define seed edges:\nedge_label_index = train_data[\"session\", \"rates\", \"aid\"].edge_label_index\nedge_label = train_data[\"session\", \"rates\", \"aid\"].edge_label\n\ntrain_loader = LinkNeighborLoader(\n    data=train_data,  # TODO\n    num_neighbors=[20,10],  # TODO\n    neg_sampling_ratio=2.0,  # TODO\n    edge_label_index=((\"session\", \"rates\", \"aid\"), edge_label_index),\n    edge_label=edge_label,\n    batch_size=10000,\n    shuffle=False,\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:51:40.490118Z","iopub.execute_input":"2022-12-20T22:51:40.490489Z","iopub.status.idle":"2022-12-20T22:51:42.998534Z","shell.execute_reply.started":"2022-12-20T22:51:40.490453Z","shell.execute_reply":"2022-12-20T22:51:42.997548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Inspect a sample:\nsampled_data = next(iter(train_loader))\n\nprint(\"Sampled mini-batch:\")\nprint(\"===================\")\nprint(sampled_data)\n","metadata":{"id":"Ogh615ka9I2c","outputId":"e4b6d4d3-f59b-4f8a-9b78-a2536ed08516","execution":{"iopub.status.busy":"2022-12-20T22:51:43.000792Z","iopub.execute_input":"2022-12-20T22:51:43.001298Z","iopub.status.idle":"2022-12-20T22:51:44.166988Z","shell.execute_reply.started":"2022-12-20T22:51:43.001259Z","shell.execute_reply":"2022-12-20T22:51:44.165864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating a Heterogeneous Link-level GNN\n\nWe are now ready to create our heterogeneous GNN.\nThe GNN is responsible for learning enriched node representations from the surrounding subgraphs, which can be then used to derive edge-level predictions.\nFor defining our heterogenous GNN, we make use of [`nn.SAGEConv`](https://pytorch-geometric.readthedocs.io/en/latest/modules/nn.html#torch_geometric.nn.conv.SAGEConv) and the [`nn.to_hetero()`](https://pytorch-geometric.readthedocs.io/en/latest/modules/nn.html#torch_geometric.nn.to_hetero_transformer.to_hetero) function, which transforms a GNN defined on homogeneous graphs to be applied on heterogeneous ones.\n\nIn addition, we define a final link-level classifier, which simply takes both node embeddings of the link we are trying to predict, and applies a dot-product on them.\n\nAs sessions do not have any node-level information, we choose to learn their features jointly via a `torch.nn.Embedding` layer. In order to improve the expressiveness of aid features, we do the same for aid nodes, and simply add their shallow embeddings to the pre-defined genre features.","metadata":{"id":"uj7biOtatAmG"}},{"cell_type":"code","source":"from torch_geometric.nn import SAGEConv, to_hetero\n\nclass GNN(torch.nn.Module):\n    def __init__(self, hidden_channels):\n        super().__init__()\n\n        self.conv1 = SAGEConv(hidden_channels, hidden_channels)\n        self.conv2 = SAGEConv(hidden_channels, hidden_channels)\n\n\n    def forward(self, x: Tensor, edge_index: Tensor) -> Tensor:\n        x = torch.nn.functional.selu(self.conv1(x,edge_index))\n        x = self.conv2(x,edge_index)\n\n        return x\n\n# Our final classifier applies the dot-product between source and destination\n# node embeddings to derive edge-level predictions:\nclass Classifier(torch.nn.Module):\n    def forward(self, x_session: Tensor, x_aid: Tensor, edge_label_index: Tensor) -> Tensor:\n        # Convert node embeddings to edge-level representations:\n        edge_feat_session = x_session[edge_label_index[0]]\n        edge_feat_aid = x_aid[edge_label_index[1]]\n\n        # Apply dot-product to get a prediction per supervision edge:\n        return (edge_feat_session * edge_feat_aid).sum(dim=-1)\n\n\nclass Model(torch.nn.Module):\n    def __init__(self, hidden_channels):\n        super().__init__()\n        # Since the dataset does not come with rich features, we also learn two\n        # embedding matrices for sessions and aid:\n        self.aid_lin = torch.nn.Linear(aids_feat_torch.shape[1], hidden_channels)\n        self.session_emb = torch.nn.Embedding(data[\"session\"].num_nodes, hidden_channels)\n        self.aid_emb = torch.nn.Embedding(data[\"aid\"].num_nodes, hidden_channels)\n\n        # Instantiate homogeneous GNN:\n        self.gnn = GNN(hidden_channels)\n\n        # Convert GNN model into a heterogeneous variant:\n        self.gnn = to_hetero(self.gnn, metadata=data.metadata())\n\n        self.classifier = Classifier()\n\n    def forward(self, data: HeteroData) -> Tensor:\n        x_dict = {\n          \"session\": self.session_emb(data[\"session\"].node_id),\n          \"aid\": self.aid_lin(data[\"aid\"].x) + self.aid_emb(data[\"aid\"].node_id),\n        } \n\n        # `x_dict` holds feature matrices of all node types\n        # `edge_index_dict` holds all edge indices of all edge types\n        x_dict = self.gnn(x_dict, data.edge_index_dict)\n\n        pred = self.classifier(\n            x_dict[\"session\"],\n            x_dict[\"aid\"],\n            data[\"session\", \"rates\", \"aid\"].edge_label_index,\n        )\n    \n        return pred\n    \n    def encode_user_aid(self, data: HeteroData) -> Tensor:\n        x_dict = {\n          \"session\": self.session_emb(data[\"session\"].node_id),\n          \"aid\": self.aid_lin(data[\"aid\"].x) + self.aid_emb(data[\"aid\"].node_id),\n        } \n\n        # `x_dict` holds feature matrices of all node types\n        # `edge_index_dict` holds all edge indices of all edge types\n        x_dict = self.gnn(x_dict, data.edge_index_dict)\n\n\n        return x_dict['session'], x_dict['aid']\nmodel = Model(hidden_channels=32)\n\nprint(model)","metadata":{"id":"ebsFf-Pr_4LF","outputId":"5756285f-592d-4783-a31d-1bbccf8fd454","execution":{"iopub.status.busy":"2022-12-20T22:51:46.723726Z","iopub.execute_input":"2022-12-20T22:51:46.724166Z","iopub.status.idle":"2022-12-20T22:51:47.959268Z","shell.execute_reply.started":"2022-12-20T22:51:46.724129Z","shell.execute_reply":"2022-12-20T22:51:47.958144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training a Heterogeneous Link-level GNN\n\nTraining our GNN is then similar to training any PyTorch model.\nWe move the model to the desired device, and initialize an optimizer that takes care of adjusting model parameters via stochastic gradient descent.\n\nThe training loop then iterates over our mini-batches, applies the forward computation of the model, computes the loss from ground-truth labels and obtained predictions (here we make use of binary cross entropy), and adjusts model parameters via back-propagation and stochastic gradient descent.","metadata":{"id":"05dfew-WuWHN"}},{"cell_type":"code","source":"model = model.to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\n\nfor epoch in range(0, 10):\n    total_loss = total_examples = 0\n    for sampled_data in tqdm(train_loader):\n        optimizer.zero_grad()\n        sampled_data = sampled_data.to(device)\n        ground_truth = sampled_data[\"session\", \"rates\", \"aid\"].edge_label\n\n        pred = model(sampled_data)\n        \n        loss = F.binary_cross_entropy_with_logits(pred,ground_truth)\n\n        loss.backward()\n        optimizer.step()\n        total_loss += float(loss) * pred.numel()\n        total_examples += pred.numel()\n    print(f\"Epoch: {epoch:03d}, Loss: {total_loss / total_examples:.4f}\")","metadata":{"id":"OqLuXEcrAMru","outputId":"65d515e9-00d6-48d9-92f2-2b41207e6d53","execution":{"iopub.status.busy":"2022-12-20T22:51:50.194513Z","iopub.execute_input":"2022-12-20T22:51:50.194889Z","iopub.status.idle":"2022-12-20T22:52:04.319070Z","shell.execute_reply.started":"2022-12-20T22:51:50.194852Z","shell.execute_reply":"2022-12-20T22:52:04.317588Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation","metadata":{}},{"cell_type":"code","source":"# Define the validation seed edges:\nedge_label_index = val_data[\"session\", \"rates\", \"aid\"].edge_label_index\nedge_label = val_data[\"session\", \"rates\", \"aid\"].edge_label\n\n\nval_loader = LinkNeighborLoader(\n    data=val_data,\n    num_neighbors=[20,10],\n    edge_label_index=((\"session\", \"rates\", \"aid\"), edge_label_index),\n    edge_label=edge_label,\n    batch_size=512,\n    shuffle=False,\n)\n\nsampled_data = next(iter(val_loader))\n\nprint(\"Sampled mini-batch:\")\nprint(\"===================\")\nprint(sampled_data)\n\n","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:52:11.422330Z","iopub.execute_input":"2022-12-20T22:52:11.422709Z","iopub.status.idle":"2022-12-20T22:52:15.241544Z","shell.execute_reply.started":"2022-12-20T22:52:11.422676Z","shell.execute_reply":"2022-12-20T22:52:15.239064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\n\npreds = []\nground_truths = []\nfor sampled_data in tqdm(val_loader):\n    with torch.no_grad():\n        sampled_data = sampled_data.to(device)\n        ground_truth_ = sampled_data['session','rates','aid'].edge_label\n        pred_ = model(sampled_data)\n        ground_truths.append(ground_truth_)\n        preds.append(pred_)\n        \n\npred = torch.cat(preds, dim=0).cpu().numpy()\nground_truth = torch.cat(ground_truths, dim=0).cpu().numpy()\nauc = roc_auc_score(ground_truth, pred)\nprint()\nprint(f\"Validation AUC: {auc:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:52:19.850558Z","iopub.execute_input":"2022-12-20T22:52:19.850965Z","iopub.status.idle":"2022-12-20T22:52:22.169224Z","shell.execute_reply.started":"2022-12-20T22:52:19.850928Z","shell.execute_reply":"2022-12-20T22:52:22.167927Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del val_loader, train_loader, edge_label_index,edge_label\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:53:19.470801Z","iopub.execute_input":"2022-12-20T22:53:19.471195Z","iopub.status.idle":"2022-12-20T22:53:19.677632Z","shell.execute_reply.started":"2022-12-20T22:53:19.471161Z","shell.execute_reply":"2022-12-20T22:53:19.676634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"user_embed = np.zeros((unique_session_id.shape[0],32)).astype(np.float32)\naid_embed = np.zeros((unique_aid_id.shape[0],32)).astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:52:24.698566Z","iopub.execute_input":"2022-12-20T22:52:24.698942Z","iopub.status.idle":"2022-12-20T22:52:25.199523Z","shell.execute_reply.started":"2022-12-20T22:52:24.698910Z","shell.execute_reply":"2022-12-20T22:52:25.198506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loader = LinkNeighborLoader(\n    data=data,\n    num_neighbors=[20,10],\n    neg_sampling_ratio=0,  # TODO\n    edge_label_index = ((\"session\", \"rates\", \"aid\"), data[\"session\", \"rates\", \"aid\"].edge_index),\n    edge_label = torch.arange(0,data[\"session\", \"rates\", \"aid\"].edge_index.shape[1]),\n    batch_size=10000,\n    shuffle=False,\n)\n\nsampled_data = next(iter(loader))","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:52:30.491241Z","iopub.execute_input":"2022-12-20T22:52:30.491963Z","iopub.status.idle":"2022-12-20T22:52:33.924052Z","shell.execute_reply.started":"2022-12-20T22:52:30.491917Z","shell.execute_reply":"2022-12-20T22:52:33.923031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del data\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for sampled_data in tqdm(loader):\n    with torch.no_grad():\n        sampled_data = sampled_data.to(device)\n        user_id = sampled_data['session'].node_id.detach().cpu()\n        aid_id = sampled_data['aid'].node_id.detach().cpu()\n        user_e, aid_e = model.encode_user_aid(sampled_data)\n\n        user_embed[user_id], aid_embed[aid_id] = user_e.detach().cpu().numpy(), aid_e.detach().cpu().numpy()\n\n","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:56:05.190132Z","iopub.execute_input":"2022-12-20T22:56:05.190508Z","iopub.status.idle":"2022-12-20T22:56:08.853626Z","shell.execute_reply.started":"2022-12-20T22:56:05.190476Z","shell.execute_reply":"2022-12-20T22:56:08.852045Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_session_np = pl.concat([unique_session_id.select(\"session\"),pl.DataFrame(user_embed)],how=\"horizontal\").to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:53:42.848580Z","iopub.execute_input":"2022-12-20T22:53:42.848957Z","iopub.status.idle":"2022-12-20T22:53:45.471268Z","shell.execute_reply.started":"2022-12-20T22:53:42.848924Z","shell.execute_reply":"2022-12-20T22:53:45.470131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_aid_np = pl.concat([unique_aid_id.select(\"aid\"),pl.DataFrame(aid_embed)],how=\"horizontal\").to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:53:49.154456Z","iopub.execute_input":"2022-12-20T22:53:49.154853Z","iopub.status.idle":"2022-12-20T22:53:49.887577Z","shell.execute_reply.started":"2022-12-20T22:53:49.154815Z","shell.execute_reply":"2022-12-20T22:53:49.886461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save(\"item_embeddings\",unique_aid_np)\nnp.save(\"user_embeddings\",unique_session_np)","metadata":{"execution":{"iopub.status.busy":"2022-12-20T22:53:52.018837Z","iopub.execute_input":"2022-12-20T22:53:52.019876Z","iopub.status.idle":"2022-12-20T22:53:53.563422Z","shell.execute_reply.started":"2022-12-20T22:53:52.019835Z","shell.execute_reply":"2022-12-20T22:53:53.561784Z"},"trusted":true},"execution_count":null,"outputs":[]}]}