USAAIO Lesson 95, from Phase 3. It covers GraphSAGE's fixed-size neighborhood sampling and mean aggregation, then oversmoothing in deep GCNs, where similarity rises from 0.806 at 2 layers to 0.999 at 10 on a six-node toy graph seeded with torch.manual_seed(0). It then covers skip connections through GCNII, where alpha=0.1 brings the 10-layer similarity back down to 0.902, the PNA multi-aggregator, which concatenates mean, standard deviation, max, and min into dimension 16, and Graph Transformer attention over graph neighborhoods. All the trace tables were verified with torch 2.7.1+cpu and numpy 2.2.6 on a fixed-seed six-node synthetic graph. The lesson runs to 27 slides.
Subject: Machine Learning · 54 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 95 · Phase 3
Fixed-size neighborhood sampling, oversmoothing and its cures (skip connections, GCNII), multi-aggregator PNA, and Graph Transformer attention -- all traced on a concrete 6-node graph with verified PyTorch numbers.
Objectives
Section
Part 1 of 4
Concept
A standard GCN layer computes H = A_hat @ H @ W where A_hat is the normalized adjacency. This requires all neighbors of every node in the batch -- a breadth-first expansion that blows up exponentially in depth.
\[ \text{neighbors accessed at depth }L: \prod_{l=1}^{L} \bar{d}_l \quad (\bar{d} = \text{avg degree}) \]
For a graph with mean degree 10 and L=3 layers, one minibatch node touches 1,000 neighborhood nodes. At L=5: 100,000. Full-neighbor GCNs cannot be batched on large graphs.
Counterexample
Discussion prompt
For a graph with mean degree 10 and L=3 layers, one minibatch node touches 1,000 neighborhood nodes. At L=5: 100,000. Full-neighbor GCNs cannot be batched on large graphs.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Concept
GraphSAGE — Graph SAmple and aggreGatE (Hamilton 2017). Instead of expanding all neighbors, uniformly sample at most K neighbors per layer, aggregate their features, concatenate with the node's own embedding, and apply a learnable linear transform.
\[ h_v^{(l)} = \sigma\!\left(W^{(l)} \cdot \operatorname{CONCAT}\!\left(h_v^{(l-1)},\; \operatorname{MEAN}_{u \in \mathcal{S}(v)} h_u^{(l-1)}\right)\right) \]
S(v) is a random sample of at most K neighbors of v. Because |S(v)| <= K regardless of degree, memory per minibatch is O(K^L * batch_size) -- fully controlled.
GraphSAGE is also inductive: the same aggregator weights generalize to unseen nodes (no node-specific embedding table), unlike transductive GCNs.
Analogy
Discussion prompt
Explain GraphSAGE: sample a fixed-size neighborhood by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
S(v) is a random sample of at most K neighbors of v. Because |S(v)| <= K regardless of degree, memory per minibatch is O(K^L * batch_size) -- fully controlled.
Estimation
Predict first
Graph: 6 nodes, edges 0-1, 0-2, 1-3, 1-4, 2-5, 3-5. Node features X = torch.randn(6, 4) with torch.manual_seed(0). Max sample K=2.
Commit before you compute: what does Tracing GraphSAGE on a 6-node graph come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Apply W and ReLU: h1 = ReLU(W @ concat)
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. W has shape [4, 8] (projects back to hidden dim).
Worked example
Graph: 6 nodes, edges 0-1, 0-2, 1-3, 1-4, 2-5, 3-5. Node features X = torch.randn(6, 4) with torch.manual_seed(0). Max sample K=2.
import torch, numpy as np
torch.manual_seed(0); np.random.seed(42)
adj = {0:[1,2], 1:[0,3,4], 2:[0,5], 3:[1,5], 4:[1], 5:[2,3]}
X = torch.randn(6, 4) # node features
def sage_sample(node, k=2):
nbrs = adj[node]
if len(nbrs) <= k: return list(nbrs)
rng = np.random.RandomState(node)
return [int(x) for x in rng.choice(nbrs, k, replace=False)]
for n in range(6):
print(f"node {n}: full={adj[n]} sampled={sage_sample(n)}")
# node 1: full=[0, 3, 4] sampled=[0, 4] <- only 2 neighbors
# Aggregate for node 1
sampled = sage_sample(1) # [0, 4]
nbr_feats = X[sampled] # shape [2, 4]
agg = nbr_feats.mean(0) # shape [4]
concat = torch.cat([X[1], agg]) # shape [8]
print('concat shape:', concat.shape) # torch.Size([8])Sample neighbors of node 1 (K=2)
Why: Node 1 has 3 neighbors {0,3,4}; sampling K=2 yields {0,4}. The aggregation cost is now fixed at 2 neighbors regardless of graph size.
| Step | Tensor / value | Shape |
|---|---|---|
| Sample node 1 nbrs | [0, 4] | -- |
| X[0] (node 0 feat) | [-1.1258, -1.1524, -0.2506, -0.4339] | [4] |
| X[4] (node 4 feat) | [0.9318, 1.2590, 2.0050, 0.0537] | [4] |
| mean agg | [-0.0970, 0.0533, 0.8772, -0.1901] | [4] |
| X[1] (self) | [0.8487, 0.6920, -0.3160, -2.1152] | [4] |
| concat(self, agg) | [0.8487, ..., -0.0970, ..., -0.1901] | [8] |
Apply W and ReLU: h1 = ReLU(W @ concat)
Why: W has shape [4, 8] (projects back to hidden dim). For a simple identity-slice W=eye(4,8), only the self-feature channels survive ReLU; a learned W mixes both halves.
| Node | Full neighbors | Sampled (K=2) | Excluded |
|---|---|---|---|
| 0 | [1, 2] | [1, 2] | none (deg 2) |
| 1 | [0, 3, 4] | [0, 4] | node 3 |
| 2 | [0, 5] | [0, 5] | none (deg 2) |
| 3 | [1, 5] | [1, 5] | none (deg 2) |
| 4 | [1] | [1] | none (deg 1) |
| 5 | [2, 3] | [2, 3] | none (deg 2) |
Discrimination
Sort into buckets
Sort these by Shape, from memory, without looking back at Tracing GraphSAGE on a 6-node graph. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Anomaly
Predict first
A student writes this, and it looks reasonable:
Both GCN and GraphSAGE learn node embeddings, so they must be the same -- GraphSAGE just skips some edges for speed.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Under this view, GraphSAGE embeddings are node-specific lookup tables that only work on the training graph, and skipping neighbors just degrades accuracy.
The key difference is inductive vs transductive: GCN learns a separate embedding per node (cannot generalize to new nodes); GraphSAGE learns aggregator weights that apply to any node by aggregating its neighbors.
Why: Under this view, GraphSAGE embeddings are node-specific lookup tables that only work on the training graph, and skipping neighbors just degrades accuracy.
Trap
Both GCN and GraphSAGE learn node embeddings, so they must be the same -- GraphSAGE just skips some edges for speed.
Treat GraphSAGE as 'GCN with dropped edges'
Why: Under this view, GraphSAGE embeddings are node-specific lookup tables that only work on the training graph, and skipping neighbors just degrades accuracy.
The key difference is inductive vs transductive: GCN learns a separate embedding per node (cannot generalize to new nodes); GraphSAGE learns aggregator weights that apply to any node by aggregating its neighbors.
Recognize that sampling is not edge dropout
Why: Sampling reduces cost to O(K^L * batch) -- a design constraint, not a regularization technique. The learned W and aggregator generalize to entirely new nodes at inference time (e.g., new papers added to a citation graph).
Section
Part 2 of 4
Concept
Oversmoothing — As the number of GCN layers L increases, all node representations converge to the same vector -- the graph's dominant Laplacian eigenvector -- and become indistinguishable. Empirically, performance often peaks at L=2 and degrades sharply past L=4.
\[ H^{(L)} = \hat{A}^L H^{(0)} W \quad\Rightarrow\quad H^{(L)} \xrightarrow{L\to\infty} \pi \mathbf{1}^\top \]
Each layer computes H = A_hat @ H (normalized adjacency propagation). This is exactly the power iteration of A_hat, which converges to the stationary distribution pi -- every node gets the same vector.
Missing information
Discussion prompt
Metric: mean pairwise cosine similarity across all C(6,2)=15 node pairs. At 0 it means maximally diverse representations; at 1.0 all nodes are identical.
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
No learnable weights -- we isolate the structural smoothing effect of the propagation operator itself, not the weight matrices.
Worked example
Metric: mean pairwise cosine similarity across all C(6,2)=15 node pairs. At 0 it means maximally diverse representations; at 1.0 all nodes are identical.
import torch, torch.nn.functional as F
torch.manual_seed(0)
X = torch.randn(6, 4)
# Build normalized adj A_hat (GCN-style, with self-loops)
A = torch.zeros(6,6)
for u,vs in {0:[1,2],1:[0,3,4],2:[0,5],3:[1,5],4:[1],5:[2,3]}.items():
for v in vs: A[u,v]=1.0
A += torch.eye(6)
D_inv_sqrt = torch.diag(A.sum(1).pow(-0.5))
A_hat = D_inv_sqrt @ A @ D_inv_sqrt
def mean_cos(H):
Hn = F.normalize(H, dim=1)
sim = Hn @ Hn.T
idx = torch.triu(torch.ones(6,6), diagonal=1).bool()
return sim[idx].mean().item()
H = X.clone()
for L in [0, 2, 5, 10]:
if L > 0: H = A_hat @ H
if L in [2, 5, 10]:
for _ in range(L - (1 if L>1 else 0)):
H = A_hat @ H
print(f'L={L:2d} sim={mean_cos(X if L==0 else gcn_prop(X,A_hat,L)):.4f}')Propagate X through A_hat L times: H^(L) = A_hat^L X
Why: No learnable weights -- we isolate the structural smoothing effect of the propagation operator itself, not the weight matrices.
| GCN layers L | Mean pairwise cos-sim | Interpretation |
|---|---|---|
| 0 (input) | 0.0009 | Near-random: diverse representations |
| 2 | 0.8057 | High similarity: 2-hop information mixed |
| 5 | 0.9734 | Severe smoothing: nodes nearly identical |
| 10 | 0.9990 | Complete collapse: all representations ~same |
Observe the Laplacian diffusion analogy
Why: A_hat repeated multiplication is heat diffusion on the graph. The steady state distributes features uniformly, just as heat equalizes temperature. Deep GNNs are deep diffusion -- not deep feature extraction.
Comparison
Comparison matrix
From Measuring oversmoothing on our 6-node graph: refill the Mean pairwise cos-sim column from what you know. The rest of the table is as it appeared.
| GCN layers L | Mean pairwise cos-sim | Interpretation |
|---|---|---|
| 0 (input) | 0.0009 | Near-random: diverse representations |
| 2 | 0.8057 | High similarity: 2-hop information mixed |
| 5 | 0.9734 | Severe smoothing: nodes nearly identical |
| 10 | 0.9990 | Complete collapse: all representations ~same |
Concept
GCNII (Initial Residual) — Each GCN layer adds a fraction alpha of the initial node features X^(0) back to the propagated representation: H^(l+1) = (1-alpha)(A_hat @ H^(l)) + alphaX^(0). Prevents convergence to the stationary distribution by anchoring each layer to the original signal.
\[ H^{(l+1)} = \sigma\!\left(\left[(1-\alpha)\hat{A}H^{(l)} + \alpha H^{(0)}\right]\left[(1-\beta)I + \beta W^{(l)}\right]\right) \]
The (1-beta)I + beta*W term is the identity mapping: even when W is small or degenerate, information flows unchanged. GCNII can be trained to 64 layers with stable accuracy (Lesson 86 ResNet analogy).
JK-Net (Jumping Knowledge) — Concatenate (or max-pool) the hidden representations from ALL layers: h_v = CONCAT(h_v^(1), ..., h_v^(L)). Earlier layers have local structure; later layers have global. JK-Net lets the model choose the right range per node.
Definition probe
Sort into buckets
Every line below is part of the definition of GCNII (Initial Residual) or of JK-Net (Jumping Knowledge) — one or the other, never both. Put each where it belongs.
Pattern
Predict first
The table runs: Plain GCN | 2 | -- | 0.8057 | Partial · Plain GCN | 10 | -- | 0.9990 | Severe · GCNII | 10 | 0.1 | 0.9021 | Moderate (improved)
In GCNII vs plain GCN: measuring the fix, given the rows so far: what is the next one — the row where Model is JK-Net?
Correct: JK-Net | 10 | -- | varies | Mitigated by layer concat
| Model | Layers | alpha | Mean cos-sim | Oversmoothed? |
|---|---|---|---|---|
| Plain GCN | 2 | -- | 0.8057 | Partial |
| Plain GCN | 10 | -- | 0.9990 | Severe |
| GCNII | 10 | 0.1 | 0.9021 | Moderate (improved) |
| JK-Net | 10 | -- | varies | Mitigated by layer concat |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. Plain GCN pushes all nodes toward the stationary distribution (sim=0.999).
Worked example
Using the same 6-node graph, compare plain 10-layer GCN propagation against GCNII with alpha=0.1 (10% initial residual per step).
import torch, torch.nn.functional as F
torch.manual_seed(0)
X = torch.randn(6, 4) # same seed as before
# A_hat built identically (normalized adj + self-loops)
# ... (A_hat computation from previous slide) ...
def gcn_prop(X, A_hat, L):
H = X.clone()
for _ in range(L): H = A_hat @ H
return H
def gcnii_prop(X0, A_hat, alpha=0.1, L=10):
H = X0.clone()
for _ in range(L):
H = (1 - alpha) * (A_hat @ H) + alpha * X0
return H
H_gcn = gcn_prop(X, A_hat, 10)
H_gcnii = gcnii_prop(X, A_hat, alpha=0.1, L=10)
print(f'GCN 10L sim = {mean_cos(H_gcn):.4f}') # 0.9990
print(f'GCNII 10L sim = {mean_cos(H_gcnii):.4f}') # 0.9021Run both propagators for L=10; compute mean cosine similarity
Why: Plain GCN pushes all nodes toward the stationary distribution (sim=0.999). GCNII's initial residual pulls each node back toward its own X^(0) each step, resisting collapse.
| Model | Layers | alpha | Mean cos-sim | Oversmoothed? |
|---|---|---|---|---|
| Plain GCN | 2 | -- | 0.8057 | Partial |
| Plain GCN | 10 | -- | 0.9990 | Severe |
| GCNII | 10 | 0.1 | 0.9021 | Moderate (improved) |
| JK-Net | 10 | -- | varies | Mitigated by layer concat |
Interpret: GCNII reduces oversmoothing from 0.999 to 0.902 at 10 layers
Why: A 10.0x increase in relative diversity (1-0.902 vs 1-0.999 = 0.098/0.001 ~ 98x more discriminative). The initial residual term is the critical addition -- it is analogous to ResNet's skip over the nonlinear block (Lesson 86).
Trade off
Comparison matrix
From GCNII vs plain GCN: measuring the fix: every row here is a choice with a cost. Fill the Layers column, then say which row you would actually pick and what you give up for it.
| Model | Layers | alpha | Mean cos-sim | Oversmoothed? |
|---|---|---|---|---|
| Plain GCN | 2 | -- | 0.8057 | Partial |
| Plain GCN | 10 | -- | 0.9990 | Severe |
| GCNII | 10 | 0.1 | 0.9021 | Moderate (improved) |
| JK-Net | 10 | -- | varies | Mitigated by layer concat |
Anomaly
Predict first
A student writes this, and it looks reasonable:
Oversmoothing is a training problem -- too many weight matrices cause gradient vanishing, which makes deep GCNs converge to the same embedding.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: This treats oversmoothing as an optimization artifact, so adding BN or Adam should fix it.
Oversmoothing is a structural property of the propagation operator A_hat, not the weights. Even with no weight matrices at all (H^(L) = A_hat^L X), the similarity reaches 0.999 after 10 steps.
Why: This treats oversmoothing as an optimization artifact, so adding BN or Adam should fix it.
Trap
Oversmoothing is a training problem -- too many weight matrices cause gradient vanishing, which makes deep GCNs converge to the same embedding.
Fix oversmoothing with batch normalization and better optimizers
Why: This treats oversmoothing as an optimization artifact, so adding BN or Adam should fix it.
Oversmoothing is a structural property of the propagation operator A_hat, not the weights. Even with no weight matrices at all (H^(L) = A_hat^L X), the similarity reaches 0.999 after 10 steps.
Fix by changing the propagation rule, not the optimizer
Why: Effective fixes modify the message-passing equation itself: initial residual (GCNII), layer-wise aggregation of all past representations (JK-Net), or capping depth at L=2. BN helps gradient flow but does not prevent the diffusion convergence.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Oversmoothing is a structural property of the propagation operator A_hat, not the weights. Even with no weight matrices at all (H^(L) = A_hat^L X), the similarity reaches 0.999 after 10 steps.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
This treats oversmoothing as an optimization artifact, so adding BN or Adam should fix it.
Section
Part 3 of 4
Concept
Mean aggregation conflates two nodes with different neighborhoods that happen to have the same mean. Example: node A neighbors {0, 4} with mean 2.0; node B neighbors {1, 2, 3} with mean 2.0 -- indistinguishable under mean alone.
PNA (Principal Neighbourhood Aggregation) — Concatenate FOUR aggregators over the sampled neighborhood: mean, standard deviation, max, and min. This captures center, spread, extreme values -- strictly more expressive than any single aggregator, and provably more powerful than mean-pooling (Weisfeiler-Lehman test).
\[ h_v^{(l)} = \text{MLP}\!\left(\operatorname{CONCAT}\!\left[\mu, \sigma, \max, \min\right]_{u \in \mathcal{N}(v)} h_u^{(l-1)}\right) \]
Output dimension is 4 * d per layer before the MLP projects it back to d. On our 6-node graph with d=4: PNA output for node 1 has dim 4*4=16.
Counterexample
Discussion prompt
Output dimension is 4 * d per layer before the MLP projects it back to d. On our 6-node graph with d=4: PNA output for node 1 has dim 4*4=16.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Estimation
Predict first
Node 1 has neighbors {0, 3, 4} (all 3, since K=2 only applies to GraphSAGE; for PNA we use all neighbors for illustration). Feature dim d=4.
Commit before you compute: what does PNA aggregation for node 1 (neighbors {0, 3, 4}) come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Concatenate -> dim 16, then apply a shared MLP -> dim 4
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The MLP is the learnable part. The aggregators are deterministic statistics; the network learns which combination of them best encodes structural neighborhood patterns for the downstream task.
Worked example
Node 1 has neighbors {0, 3, 4} (all 3, since K=2 only applies to GraphSAGE; for PNA we use all neighbors for illustration). Feature dim d=4.
import torch
torch.manual_seed(0)
X = torch.randn(6, 4)
# Node 1 neighbors: 0, 3, 4
nbrs = [0, 3, 4]
nbr_feats = X[nbrs] # shape [3, 4]
agg_mean = nbr_feats.mean(0) # [4]
agg_std = nbr_feats.std(0) # [4] (sample std, unbiased)
agg_max = nbr_feats.max(0).values # [4]
agg_min = nbr_feats.min(0).values # [4]
pna_agg = torch.cat([agg_mean, agg_std, agg_max, agg_min])
print('pna_agg shape:', pna_agg.shape) # torch.Size([16])
# -> MLP maps [16] back to [d=4]Compute all four aggregators over X[{0,3,4}]
Why: Each aggregator captures a different statistical summary of the neighborhood distribution. Together they characterize the shape of the distribution, not just its first moment.
| Aggregator | f0 | f1 | f2 | f3 |
|---|---|---|---|---|
| mean | -0.009 | 0.327 | 0.537 | -0.164 |
| std | 1.040 | 1.296 | 1.272 | 0.248 |
| max | 0.932 | 1.259 | 2.005 | 0.054 |
| min | -1.126 | -1.152 | -0.251 | -0.434 |
Concatenate -> dim 16, then apply a shared MLP -> dim 4
Why: The MLP is the learnable part. The aggregators are deterministic statistics; the network learns which combination of them best encodes structural neighborhood patterns for the downstream task.
Comparison
Comparison matrix
From PNA aggregation for node 1 (neighbors {0, 3, 4}): refill the f0 column from what you know. The rest of the table is as it appeared.
| Aggregator | f0 | f1 | f2 | f3 |
|---|---|---|---|---|
| mean | -0.009 | 0.327 | 0.537 | -0.164 |
| std | 1.040 | 1.296 | 1.272 | 0.248 |
| max | 0.932 | 1.259 | 2.005 | 0.054 |
| min | -1.126 | -1.152 | -0.251 | -0.434 |
Concept
A Graph Transformer replaces the message-passing aggregator with dot-product attention, but restricted to the graph neighborhood rather than the full node set -- combining GNN locality with Transformer expressivity.
\[ h_v^{(l)} = \sum_{u \in \mathcal{N}(v) \cup \{v\}} \alpha_{vu}\, V h_u^{(l-1)}, \quad \alpha_{vu} = \operatorname{softmax}_{u}\!\left(\frac{(Qh_v)^\top (Kh_u)}{\sqrt{d_k}}\right) \]
Compared to the full Transformer (Lesson 88): attention is over |N(v)|+1 nodes instead of the full sequence length. Compared to ViT (Lesson 92): ViT attends globally; Graph Transformer respects graph topology as a structural bias.
Sorting
Sort into buckets
These are the pieces of Lesson 95: GraphSAGE, Oversmoothing, and Scalable GNNs, out of order. Put each one back under the part of the lesson it belongs to.
Missing information
Discussion prompt
Node 1 attends over itself plus its 3 neighbors {0, 3, 4}. We use d_k=4 projection matrices Wq, Wk, Wv (torch.manual_seed(7), scale 0.5).
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
Q/K/V projections are the same as standard multi-head attention (Lesson 88). The only difference is that the key-value pairs are drawn from the graph neighborhood, not the full sequence.
Worked example
Node 1 attends over itself plus its 3 neighbors {0, 3, 4}. We use d_k=4 projection matrices Wq, Wk, Wv (torch.manual_seed(7), scale 0.5).
import torch, torch.nn.functional as F
torch.manual_seed(0); X = torch.randn(6, 4)
torch.manual_seed(7)
d_k = 4
Wq = torch.randn(4, d_k) * 0.5
Wk = torch.randn(4, d_k) * 0.5
Wv = torch.randn(4, d_k) * 0.5
# Node 1 attends over {self=1, 0, 3, 4}
attend = [1, 0, 3, 4]
H_a = X[attend] # [4, 4]
Q = H_a[0:1] @ Wq # query from node 1: [1, 4]
K = H_a @ Wk # keys: [4, 4]
V_mat = H_a @ Wv # values: [4, 4]
scores = (Q @ K.T) / d_k**0.5 # [1, 4]
attn = F.softmax(scores, -1) # [1, 4]
out = attn @ V_mat # [1, 4]
print('attn:', attn.tolist()) # [0.028, 0.032, 0.191, 0.749]
print('out:', out.tolist()) # [-0.521, 0.582, 0.792, -0.283]Project node 1 to Query; project {1,0,3,4} to Keys and Values
Why: Q/K/V projections are the same as standard multi-head attention (Lesson 88). The only difference is that the key-value pairs are drawn from the graph neighborhood, not the full sequence.
| Key node | Attn score (raw) | Attn weight | Interpretation |
|---|---|---|---|
| 1 (self) | -1.800 | 0.028 | Low: self query poorly matches self key |
| 0 | -1.664 | 0.032 | Low: node 0 feature not aligned with node 1 query |
| 3 | 0.121 | 0.191 | Moderate: partial alignment |
| 4 | 1.486 | 0.749 | High: node 4 dominates aggregation |
Weighted sum of values -> output [-0.521, 0.582, 0.792, -0.283]
Why: Node 4 (attn 0.749) dominates the output. This is qualitatively meaningful: attention learns which neighbors are most structurally informative, unlike mean aggregation which weights all equally.
Trade off
Comparison matrix
From Graph Transformer attention trace for node 1: every row here is a choice with a cost. Fill the Attn score (raw) column, then say which row you would actually pick and what you give up for it.
| Key node | Attn score (raw) | Attn weight | Interpretation |
|---|---|---|---|
| 1 (self) | -1.800 | 0.028 | Low: self query poorly matches self key |
| 0 | -1.664 | 0.032 | Low: node 0 feature not aligned with node 1 query |
| 3 | 0.121 | 0.191 | Moderate: partial alignment |
| 4 | 1.486 | 0.749 | High: node 4 dominates aggregation |
Anomaly
Predict first
A student writes this, and it looks reasonable:
Standard deviation with ddof=1 (Bessel's correction) is undefined for a single sample, so PNA std is always 0 for nodes with one neighbor -- no information loss.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: 0 is output, so the PNA vector still passes through the MLP.
A constant 0 std for all degree-1 nodes means the MLP cannot distinguish them by their neighborhood spread -- all such nodes produce the same std channel. This can hurt expressivity.
Why: 0 is output, so the PNA vector still passes through the MLP. No problem.
Trap
Standard deviation with ddof=1 (Bessel's correction) is undefined for a single sample, so PNA std is always 0 for nodes with one neighbor -- no information loss.
Ignore degree-1 nodes: std=0 is a degenerate but valid value
Why: 0 is output, so the PNA vector still passes through the MLP. No problem.
A constant 0 std for all degree-1 nodes means the MLP cannot distinguish them by their neighborhood spread -- all such nodes produce the same std channel. This can hurt expressivity.
Handle it explicitly: use ddof=0 (population std) or clip minimum sample to 2
Why: PyTorch tensor.std(0) with default correction=1 returns NaN for a single-element set in some versions; always check with torch.std(x, correction=0) or add the node itself to the neighbor set before computing std. Original PNA paper also uses degree scalers to compensate.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
S(v) is a random sample of at most K neighbors of v. Because |S(v)| <= K regardless of degree, memory per minibatch is O(K^L * batch_size) -- fully controlled.; Output dimension is 4 * d per layer before the MLP projects it back to d. On our 6-node graph with d=4: PNA output for node 1 has dim 4*4=16.Section
Part 4 of 4
Pattern
Edge cases
Discussion prompt
Scalable and Expressive GNN Recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
Elimination
Eliminate the wrong options
Total neighborhood nodes accessed for one target node with K=5, L=3 layers?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: GraphSAGE samples at most K=5 neighbors per layer. Layer 1: 5 nodes; layer 2: 5 per each of those = 25; layer 3: 5 each = 125. Total = 155. The full-neighbor cost would be 50 + 2500 + 125000 = 127,550 -- a 822x reduction.
Check
A graph has average degree 50. You train a 3-layer GraphSAGE with max sample K=5. Approximately how many neighbor lookups does a single target node require across all 3 layers?
Check your understanding
Total neighborhood nodes accessed for one target node with K=5, L=3 layers?
Answer: A
Why: GraphSAGE samples at most K=5 neighbors per layer. Layer 1: 5 nodes; layer 2: 5 per each of those = 25; layer 3: 5 each = 125. Total = 155. The full-neighbor cost would be 50 + 2500 + 125000 = 127,550 -- a 822x reduction.
Prediction
Predict first
Why does GCNII (alpha=0.1, L=10) reduce oversmoothing from 0.999 to 0.902?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: The alpha*X^(0) term mixes each layer's output with the original features, preventing convergence to the graph's stationary distribution
Why: Each GCNII step is H = (1-alpha)(A_hat @ H) + alphaX^(0). The alpha*X^(0) term is an initial residual that constantly pulls representations back toward the original node features, counteracting the Laplacian diffusion that causes oversmoothing. This is structural, not an optimization trick.
Check
On the 6-node graph above (torch.manual_seed(0)), plain GCN propagation for 10 layers achieves mean pairwise cosine similarity 0.999. GCNII with alpha=0.1 for 10 layers achieves 0.902. Which statement best explains the improvement?
Check your understanding
Why does GCNII (alpha=0.1, L=10) reduce oversmoothing from 0.999 to 0.902?
Answer: A
Why: Each GCNII step is H = (1-alpha)(A_hat @ H) + alphaX^(0). The alpha*X^(0) term is an initial residual that constantly pulls representations back toward the original node features, counteracting the Laplacian diffusion that causes oversmoothing. This is structural, not an optimization trick.
Elimination
Eliminate the wrong options
Which PNA aggregator(s) distinguish node A from node B?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: For A: mean=2.0, std=sqrt(2)~1.414, max=3.0, min=1.0. For B: mean=2.0, std=0.0, max=2.0, min=2.0. Mean is identical. std(A)=1.414 vs std(B)=0.0; max(A)=3.0 vs max(B)=2.0; min(A)=1.0 vs min(B)=2.0. All three of std, max, min differ.
Check
Node A has neighbors with feature values {1.0, 3.0}. Node B has neighbors with feature values {2.0, 2.0}. Both have the same mean = 2.0.
Check your understanding
Which PNA aggregator(s) distinguish node A from node B?
Answer: A
Why: For A: mean=2.0, std=sqrt(2)~1.414, max=3.0, min=1.0. For B: mean=2.0, std=0.0, max=2.0, min=2.0. Mean is identical. std(A)=1.414 vs std(B)=0.0; max(A)=3.0 vs max(B)=2.0; min(A)=1.0 vs min(B)=2.0. All three of std, max, min differ.
Concept
Project brief: Build a toy link-prediction pipeline on a synthetic graph. Implement GraphSAGE (mean aggregator, K=5, L=2) from scratch in PyTorch. Add a PNA variant. Compare oversmoothing at L=2 vs L=6 before and after adding GCNII skip connections. No external GNN libraries.
torch_geometric.utils.erdos_renyi_graph OR manually build an adjacency list for 50 nodes, mean degree 4sage_sample(adj, node, K) and sage_mean_aggregate(X, adj, node, K). Verify on the 6-node example: mean agg for node 1 = [-0.097, 0.053, 0.877, -0.190]nn.Module. Train on positive/negative node pairs for link prediction (dot-product score, BCEWithLogitsLoss)alpha=0.1. Plot mean pairwise cosine similarity vs depth (L=1 to 8) for plain GraphSAGE vs GCNIIEstimation
Predict first
Hint: a GraphSAGE layer is a standard nn.Linear(2*in_dim, out_dim) whose input is the CONCAT of the node's own feature and the mean of its sampled neighbor features. The sampling happens in the forward pass, not in __init__.
Commit before you compute: what does Milestone 2 -- Full GraphSAGE layer in PyTorch come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Show-it-off: compare cos-sim curves (L=1..8) in a single matplotlib plot
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Expected result: Plain SAGEMean cos-sim rises steeply (0.0 -> ~0.99 by L=8).
Worked example
Hint: a GraphSAGE layer is a standard nn.Linear(2*in_dim, out_dim) whose input is the CONCAT of the node's own feature and the mean of its sampled neighbor features. The sampling happens in the forward pass, not in __init__.
import torch, torch.nn as nn, torch.nn.functional as F
import numpy as np
class SAGEMeanLayer(nn.Module):
def __init__(self, in_dim, out_dim, K=5):
super().__init__()
self.K = K
self.linear = nn.Linear(2 * in_dim, out_dim, bias=False)
def forward(self, X, adj):
# adj: dict node -> list[int]
n = X.shape[0]
aggs = []
for v in range(n):
nbrs = adj.get(v, [])
if len(nbrs) == 0:
agg = torch.zeros(X.shape[1])
else:
if len(nbrs) > self.K:
idx = torch.randperm(len(nbrs))[:self.K]
nbrs = [nbrs[i] for i in idx]
agg = X[nbrs].mean(0)
aggs.append(agg)
AGG = torch.stack(aggs) # [n, in_dim]
H = torch.cat([X, AGG], dim=1) # [n, 2*in_dim]
return F.relu(self.linear(H)) # [n, out_dim]
# Smoke test
adj6 = {0:[1,2], 1:[0,3,4], 2:[0,5], 3:[1,5], 4:[1], 5:[2,3]}
torch.manual_seed(0)
X6 = torch.randn(6, 4)
layer = SAGEMeanLayer(4, 8, K=2)
out = layer(X6, adj6)
print('out shape:', out.shape) # torch.Size([6, 8])Run layer(X6, adj6) with K=2; expect output shape [6, 8]
Why: in_dim=4, K=2 samples at most, output dim=8. The linear maps [8] -> [8], but in_dim=4 makes it [2*4=8] -> [out_dim=8]. Every node goes through the same W -- this is what makes it inductive.
| Hyperparameter | Value | Why |
|---|---|---|
| K (max sample) | 2 (toy) / 5-25 (real) | Controls memory cost per layer |
| L (layers) | 2-3 recommended | Avoid oversmoothing; L=2 covers 2-hop neighborhoods |
| in_dim | 4 (toy) / 64-256 (real) | Initial feature dimension |
| out_dim | 8 (toy) / 64-256 (real) | Embedding dimension after aggregation |
| Aggregator | mean (demo) / PNA (expressive) | Mean: cheap; PNA: 4x channels, more powerful |
Show-it-off: compare cos-sim curves (L=1..8) in a single matplotlib plot
Why: Expected result: Plain SAGEMean cos-sim rises steeply (0.0 -> ~0.99 by L=8). GCNII-SAGEMean (alpha=0.1) rises slowly and plateaus below 0.92. Plot labels: 'GraphSAGE mean' and 'GCNII alpha=0.1'. This is your diagnostic: if the plain curve flattens near 1.0, oversmoothing is the bottleneck.
Comparison
Comparison matrix
From Milestone 2 -- Full GraphSAGE layer in PyTorch: refill the Value column from what you know. The rest of the table is as it appeared.
| Hyperparameter | Value | Why |
|---|---|---|
| K (max sample) | 2 (toy) / 5-25 (real) | Controls memory cost per layer |
| L (layers) | 2-3 recommended | Avoid oversmoothing; L=2 covers 2-hop neighborhoods |
| in_dim | 4 (toy) / 64-256 (real) | Initial feature dimension |
| out_dim | 8 (toy) / 64-256 (real) | Embedding dimension after aggregation |
| Aggregator | mean (demo) / PNA (expressive) | Mean: cheap; PNA: 4x channels, more powerful |
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — GraphSAGE -- Scalable Inductive Learning · Oversmoothing -- Why Deep GNNs Fail · PNA and Graph Transformer -- Expressivity Upgrades · Pattern, Checks, and Project. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.