diff --git a/GNN-tutorial-solution.ipynb b/GNN-tutorial-solution.ipynb index f57c03d..01a6c30 100644 --- a/GNN-tutorial-solution.ipynb +++ b/GNN-tutorial-solution.ipynb @@ -156,40 +156,65 @@ "import torch\n", "from torch_geometric.nn import MessagePassing\n", "import math\n", + "from typing import Optional\n", "\n", "def glorot(tensor):\n", " if tensor is not None:\n", " stdv = math.sqrt(6.0 / (tensor.size(-2) + tensor.size(-1)))\n", " tensor.data.uniform_(-stdv, stdv)\n", "\n", + "def maybe_num_nodes(index, num_nodes=None):\n", + " return index.max().item() + 1 if num_nodes is None else num_nodes\n", "\n", "def zeros(tensor):\n", " if tensor is not None:\n", " tensor.data.fill_(0)\n", "\n", " \n", - "def add_self_loops(edge_index, num_nodes=None):\n", - " loop_index = torch.arange(0, num_nodes, dtype=torch.long,\n", - " device=edge_index.device)\n", + "def add_self_loops(edge_index, edge_weight: Optional[torch.Tensor] = None,\n", + " fill_value: float = 1., num_nodes: Optional[int] = None):\n", + " r\"\"\"Adds a self-loop :math:`(i,i) \\in \\mathcal{E}` to every node\n", + " :math:`i \\in \\mathcal{V}` in the graph given by :attr:`edge_index`.\n", + " In case the graph is weighted, self-loops will be added with edge weights\n", + " denoted by :obj:`fill_value`.\n", + "\n", + " Args:\n", + " edge_index (LongTensor): The edge indices.\n", + " edge_weight (Tensor, optional): One-dimensional edge weights.\n", + " (default: :obj:`None`)\n", + " fill_value (float, optional): If :obj:`edge_weight` is not :obj:`None`,\n", + " will add self-loops with edge weights of :obj:`fill_value` to the\n", + " graph. (default: :obj:`1.`)\n", + " num_nodes (int, optional): The number of nodes, *i.e.*\n", + " :obj:`max_val + 1` of :attr:`edge_index`. (default: :obj:`None`)\n", + "\n", + " :rtype: (:class:`LongTensor`, :class:`Tensor`)\n", + " \"\"\"\n", + " N = maybe_num_nodes(edge_index, num_nodes)\n", + "\n", + " loop_index = torch.arange(0, N, dtype=torch.long, device=edge_index.device)\n", " loop_index = loop_index.unsqueeze(0).repeat(2, 1)\n", "\n", + " if edge_weight is not None:\n", + " assert edge_weight.numel() == edge_index.size(1)\n", + " loop_weight = edge_weight.new_full((N, ), fill_value)\n", + " edge_weight = torch.cat([edge_weight, loop_weight], dim=0)\n", + "\n", " edge_index = torch.cat([edge_index, loop_index], dim=1)\n", "\n", - " return edge_index\n", + " return edge_index, edge_weight\n", "\n", "\n", "def degree(index, num_nodes=None, dtype=None):\n", " out = torch.zeros((num_nodes), dtype=dtype, device=index.device)\n", - " return out.scatter_add_(0, index, out.new_ones((index.size(0))))\n", + " ret = out.scatter_add_(0, index, out.new_ones((index.size(0))))\n", + " return ret\n", " \n", - "\n", "class GCNConv(MessagePassing):\n", " def __init__(self, in_channels, out_channels):\n", " super(GCNConv, self).__init__(aggr='add') # \"Add\" aggregation.\n", " self.lin = torch.nn.Linear(in_channels, out_channels)\n", - " \n", - " self.reset_parameters()\n", - " \n", + "\n", " def reset_parameters(self):\n", " glorot(self.lin.weight)\n", " zeros(self.lin.bias)\n", @@ -197,53 +222,33 @@ " def forward(self, x, edge_index):\n", " # x has shape [N, in_channels]\n", " # edge_index has shape [2, E]\n", - " \n", - " ########################################################################\n", - " # START OF YOUR CODE (DO NOT DELETE/MODIFY THIS LINE) #\n", - " ########################################################################\n", + "\n", " # Step 1: Add self-loops to the adjacency matrix.\n", - " \n", - " edge_index = add_self_loops(edge_index, num_nodes=x.size(0))\n", + " edge_index, _ = add_self_loops(edge_index, num_nodes=x.size(0))\n", "\n", " # Step 2: Linearly transform node feature matrix.\n", " x = self.lin(x)\n", "\n", - " # Step 3-5: Start propagating messages.\n", - "\n", - " return self.propagate(edge_index, x=x)\n", - " ########################################################################\n", - " # END OF YOUR CODE #\n", - " ######################################################################## \n", - "\n", - "\n", - "\n", - "\n", - " def message(self, x_j, edge_index, size):\n", - " # x_j has shape [E, out_channels]\n", - "\n", - " ########################################################################\n", - " # START OF YOUR CODE (DO NOT DELETE/MODIFY THIS LINE) #\n", - " ########################################################################\n", - "\n", - " # Step 3: Normalize node features.\n", + " # Step 3: Compute normalization\n", " row, col = edge_index\n", - " deg = degree(row, size[0], dtype=x_j.dtype)\n", + " deg = degree(row, x.size(0), dtype=x.dtype)\n", " deg_inv_sqrt = deg.pow(-0.5)\n", - " deg_inv_sqrt[deg_inv_sqrt == float('inf')] = 0\n", " norm = deg_inv_sqrt[row] * deg_inv_sqrt[col]\n", "\n", - " return norm.view(-1, 1) * x_j \n", - " \n", - " ########################################################################\n", - " # END OF YOUR CODE #\n", - " ######################################################################## \n", - " \n", + " # Step 4-6: Start propagating messages.\n", + " return self.propagate(edge_index, size=(x.size(0), x.size(0)), x=x,\n", + " norm=norm)\n", + "\n", + " def message(self, x_j, norm):\n", + " # x_j has shape [E, out_channels]\n", "\n", + " # Step 4: Normalize node features.\n", + " return norm.view(-1, 1) * x_j\n", "\n", " def update(self, aggr_out):\n", " # aggr_out has shape [N, out_channels]\n", "\n", - " # Step 5: Return new node embeddings.\n", + " # Step 6: Return new node embeddings.\n", " return aggr_out" ] }, @@ -450,7 +455,7 @@ " outs['{}_loss'.format(key)] = loss\n", " outs['{}_acc'.format(key)] = acc\n", "\n", - " return outs" + " return outs\n" ] }, { @@ -474,6 +479,7 @@ } ], "source": [ + "\n", "runs = 10\n", "epochs = 200\n", "lr = 0.01\n", @@ -1310,4 +1316,4 @@ }, "nbformat": 4, "nbformat_minor": 2 -} +} \ No newline at end of file