Nearly every post on this blog is about transformers, because that is where the money and the hype are. But a large share of the enterprise PyTorch work we get called into has no language model in it at all: fraud rings, money-laundering typologies, account-takeover detection, supply-chain risk, insurance abuse. These problems share a shape. The signal is not in any single row, it is in the relationships between rows — the shared device, the recycled phone number, the merchant two hops away that six flagged accounts all touched last Tuesday.
Gradient-boosted trees on flat features miss that, and they miss it in a specific way: they see each transaction as an island. Graph neural networks (GNNs) do not. This tutorial is a practical build: a heterogeneous GNN for fraud detection in PyTorch, trained with neighbour sampling so it works on a graph far bigger than GPU memory, and served with latency low enough to sit in a payment authorisation path.
When a graph model is worth it
Be honest about this before you start, because GNN projects fail more often on data plumbing than on modelling.
A GNN is likely to beat your tabular baseline when:
- Entities are genuinely linked and the links are adversarially interesting (shared devices, IPs, bank accounts, addresses, referral chains).
- Labels are sparse — a few thousand confirmed fraud cases against hundreds of millions of transactions. Message passing propagates signal from labelled nodes to their unlabelled neighbourhoods, which is exactly what semi-supervised learning is for.
- Fraud is coordinated. Single-actor fraud is a tabular problem. Rings, mule networks and bot farms are graph problems.
A GNN is probably not worth it when your entities are isolated, your features already encode the aggregates you care about, or when nobody can produce a reliable entity-resolution layer. If "is this the same device" is a coin flip in your data, the graph is noise with extra steps.
The usual honest answer in production is both: a GNN that produces node embeddings plus a GBDT that consumes those embeddings alongside the existing tabular features. That ensemble is what tends to ship.
Modelling the graph
Use a heterogeneous graph; collapsing everything into one node type throws away the structure that makes the model work. A workable schema for card fraud:
| Node type | Features |
|---|---|
transaction | amount, currency, hour-of-day, channel, MCC, risk flags |
card | tenure, issuing country, historical chargeback rate |
device | OS, browser fingerprint hash buckets, first-seen age |
merchant | category, country, settlement lag, volume bucket |
| Edge type | Meaning |
|---|---|
(card, used_in, transaction) | the card that paid |
(device, initiated, transaction) | the device at the point of sale |
(transaction, at, merchant) | where it happened |
Labels live on transaction nodes: confirmed fraud, confirmed legitimate, or unknown. Most nodes are unknown, which is fine.
In PyTorch Geometric that is a HeteroData object:
import torch
from torch_geometric.data import HeteroData
data = HeteroData()
data["transaction"].x = tx_features # [N_tx, F_tx] float32
data["transaction"].y = tx_labels # [N_tx] long, -1 for unknown
data["transaction"].time = tx_timestamps # [N_tx] int64 epoch seconds
data["card"].x = card_features
data["device"].x = device_features
data["merchant"].x = merchant_features
data["card", "used_in", "transaction"].edge_index = card_tx_edges # [2, E]
data["device", "initiated", "transaction"].edge_index = device_tx_edges
data["transaction", "at", "merchant"].edge_index = tx_merchant_edges
# message passing needs both directions
import torch_geometric.transforms as T
data = T.ToUndirected()(data)
Two rules that will save you a wasted quarter:
- Put a timestamp on every node and every edge. You need it for the split and for inference-time neighbour filtering.
- Never put a feature on a transaction node that was computed after the fraud decision. Chargeback flags, manual review outcomes, settlement status — all of these leak. In graph land leakage is worse than in tabular land, because a leaky feature on one node contaminates every node within
khops of it.
The temporal split
Random splits on fraud graphs produce beautiful, meaningless AUCs. Split by time, and filter edges by time:
train_mask = data["transaction"].time < t_train_end
val_mask = (data["transaction"].time >= t_train_end) & (data["transaction"].time < t_val_end)
test_mask = data["transaction"].time >= t_val_end
When you sample neighbourhoods for a training node, only neighbours older than that node may be included. PyG supports this directly through time_attr on the loader, and it is the single most important correctness knob in the whole pipeline. If your offline AUC is 0.99 and your shadow-mode AUC is 0.78, you skipped this.
Neighbour sampling: training on a graph that does not fit
Full-batch message passing over a billion-edge graph is not happening on one GPU. NeighborLoader samples a small multi-hop subgraph per mini-batch of seed nodes and moves only that to the device:
from torch_geometric.loader import NeighborLoader
train_loader = NeighborLoader(
data,
num_neighbors={key: [15, 10] for key in data.edge_types}, # 2 hops
input_nodes=("transaction", train_mask),
time_attr="time", # temporal sampling: no peeking into the future
batch_size=1024,
shuffle=True,
num_workers=8,
persistent_workers=True,
)
Fan-out [15, 10] means up to 15 neighbours at hop one and 10 per those at hop two. Keep it small. Two hops with modest fan-out is almost always the right starting point: hop three rarely adds AUC and it multiplies subgraph size, and on a graph with hub nodes (a popular merchant, a shared data-centre IP) it will blow up your batches. Cap or drop hub edges above a degree threshold — they carry almost no information and dominate sampling cost.
The model
Start with heterogeneous GraphSAGE. It is the boring choice and it wins a lot.
import torch.nn.functional as F
from torch import nn
from torch_geometric.nn import SAGEConv, HeteroConv
class HeteroFraudGNN(nn.Module):
def __init__(self, metadata, in_dims, hidden=128, out=64, layers=2):
super().__init__()
node_types, edge_types = metadata
# project each node type's raw features into a shared space
self.proj = nn.ModuleDict(
{nt: nn.Linear(in_dims[nt], hidden) for nt in node_types}
)
self.convs = nn.ModuleList([
HeteroConv(
{et: SAGEConv((-1, -1), hidden) for et in edge_types},
aggr="sum",
)
for _ in range(layers)
])
self.norms = nn.ModuleList([
nn.ModuleDict({nt: nn.LayerNorm(hidden) for nt in node_types})
for _ in range(layers)
])
self.head = nn.Sequential(
nn.Linear(hidden, out), nn.ReLU(), nn.Linear(out, 1)
)
def forward(self, x_dict, edge_index_dict):
h = {nt: self.proj[nt](x) for nt, x in x_dict.items()}
for conv, norm in zip(self.convs, self.norms):
h = conv(h, edge_index_dict)
h = {nt: F.relu(norm[nt](v)) for nt, v in h.items()}
return self.head(h["transaction"]).squeeze(-1), h # logits, embeddings
Returning the embedding dict as well as the logits is deliberate: those embeddings are what you hand to the GBDT later.
Training loop
Fraud is a 0.1%-positive problem, so the loss needs to say so. pos_weight on BCEWithLogitsLoss is the simplest effective fix.
model = HeteroFraudGNN(data.metadata(), in_dims).cuda()
opt = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4, fused=True)
pos_weight = torch.tensor([ (y_train == 0).sum() / (y_train == 1).sum() ]).cuda()
loss_fn = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
for epoch in range(20):
model.train()
for batch in train_loader:
batch = batch.to("cuda", non_blocking=True)
with torch.autocast("cuda", dtype=torch.bfloat16):
logits, _ = model(batch.x_dict, batch.edge_index_dict)
# only the seed nodes of this batch carry supervision
n = batch["transaction"].batch_size
seed_logits = logits[:n]
seed_y = batch["transaction"].y[:n].float()
keep = seed_y >= 0 # drop unlabelled seeds
loss = loss_fn(seed_logits[keep], seed_y[keep])
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
opt.zero_grad(set_to_none=True)
The [:n] slice is the detail newcomers get wrong. A sampled subgraph contains thousands of support nodes that exist only to supply messages; PyG orders the seed nodes first, and only those should contribute to the loss. Score the support nodes too and you train on duplicated, time-leaking targets.
A few things that reliably move the metric:
- Lazy
(-1, -1)input dims inSAGEConvlet PyG infer per-edge-type feature sizes. Run one dummy batch before wrapping in DDP or FSDP so the lazy modules are materialised. - Degree and recency features on nodes (
transactions_last_1h,distinct_cards_on_device_7d) help a lot. A GNN learns structure, but handing it cheap structural summaries still pays. - Edge dropout (
DropEdge, 10-20%) is the most effective regulariser we see on fraud graphs; it also makes the model less brittle when an entity-resolution rule changes upstream.
Evaluation that means something
AUC-ROC on a 0.1% positive rate is a vanity metric. Report:
- PR-AUC (average precision) as the headline number.
- Recall at a fixed review budget — "catch rate in the top 500 alerts per day" is what a fraud ops team actually buys.
- Dollars at risk caught, weighting each case by amount.
- Ring-level recall: of known fraud clusters in the test window, how many did you flag at least one member of? This is the metric a GNN should dominate a GBDT on, and if it does not, your graph is probably not carrying the signal you assumed.
Always run a GBDT-on-tabular-features baseline in the same temporal split. Without it you cannot tell whether the graph is earning its operational cost.
Serving inside a 50ms budget
This is where GNN projects die, so design for it from day one. Two architectures:
Batch embeddings + online tabular model (recommended for most). Run the GNN as a nightly or hourly job over the whole graph, write card, device and merchant embeddings to a feature store (Redis, DynamoDB, Feast), and have the online scorer look up those embeddings, concatenate them with live transaction features, and run a small MLP or GBDT. Latency is a key-value lookup plus a tiny model: single-digit milliseconds. The entity embeddings are hours stale, which is acceptable because card and device behaviour does not change in minutes. This gets you most of the lift for a fraction of the operational risk.
For the offline pass, use layer-wise inference rather than per-node sampling — it computes each layer once over all nodes instead of re-expanding overlapping neighbourhoods:
@torch.inference_mode()
def nightly_embeddings(model, data, device="cuda"):
model.eval()
# PyG's HGTLoader/NeighborLoader in inference mode, or a layer-wise pass
# over the full node set, writing hidden states back to CPU per layer.
...
Online subgraph fetch (only when you need it). Pull the 1-2 hop neighbourhood for the incoming transaction from a graph store at request time and run the GNN on it. Fresh, and much harder: you need a sub-10ms neighbour fetch, aggressive fan-out caps, hub-node guards, and a fallback path for when the graph store is slow. Reserve this for high-value decisions where minutes-fresh structure genuinely matters, such as account-takeover at login.
Either way, export the model with torch.export and AOTInductor, or at least torch.compile it with dynamic shapes enabled — sampled subgraphs vary in size every call, and recompilation on every request will destroy your tail latency.
Operating it
- Monitor the graph, not just the model. Edge counts per type, mean degree, new-node rate, entity-resolution match rate. A silent upstream change in device fingerprinting will change your graph's topology and quietly degrade the model with every input feature still looking perfectly normal.
- Retrain on a schedule and measure drift in ring structure, not only in feature distributions. Adversaries adapt faster here than in almost any other ML domain.
- Keep an explanation path. Fraud analysts will not action a score they cannot interpret. PyG's
ExplainerwithGNNExplainerorCaptumExplainercan surface the handful of edges that drove a score, which is enough to render "this card shares a device with three accounts charged back last week" in the review UI. Budget for this; it is frequently the difference between a deployed model and a shelved one. - Version the graph schema alongside the model weights. A checkpoint trained on a different edge-type set is not loadable in any useful sense, and the failure is often a silent shape mismatch rather than an exception.
Where this is heading
Two developments are worth tracking. First, graph structure is increasingly being fed to language models rather than consumed only by classifiers: a GNN produces the candidate ring, an LLM writes the narrative for the suspicious-activity report. That hybrid shortens investigation time more than any AUC improvement does. Second, graph foundation models — pretrained on many graphs, applied zero- or few-shot to a new one — are starting to be usable for cold-start cases, in much the same way time-series foundation models now are. Neither changes the engineering above. The hard parts remain entity resolution, temporal correctness, and a serving path that fits the latency budget.
If you are weighing a graph approach for fraud, AML, abuse or risk and want a second opinion on whether the data supports it before committing a quarter to it, get in touch — a short scoping conversation usually settles it either way.