Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 50 additions & 44 deletions GNN-tutorial-solution.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -156,94 +156,99 @@
"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",
"\n",
" 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"
]
},
Expand Down Expand Up @@ -450,7 +455,7 @@
" outs['{}_loss'.format(key)] = loss\n",
" outs['{}_acc'.format(key)] = acc\n",
"\n",
" return outs"
" return outs\n"
]
},
{
Expand All @@ -474,6 +479,7 @@
}
],
"source": [
"\n",
"runs = 10\n",
"epochs = 200\n",
"lr = 0.01\n",
Expand Down Expand Up @@ -1310,4 +1316,4 @@
},
"nbformat": 4,
"nbformat_minor": 2
}
}