AI Fraud Detection with Graph Neural Networks

Using graph neural networks to detect fraud by modeling transaction relationships, account networks, and behavioral patterns in financial systems

#fraud-detection#gnn#graph-neural-networks#fintech
Cover image for the article: AI Fraud Detection with Graph Neural Networks

Traditional fraud detection relies on rule-based systems and tabular ML models that evaluate each transaction in isolation. But fraud is inherently a graph problem - fraudulent actors create networks of accounts, transactions, and devices that form distinctive patterns invisible to point-wise classification. Graph Neural Networks (GNNs) capture these relational patterns, delivering 30-50% improvement in fraud detection over tabular approaches.

This article covers the architecture for deploying GNN-based fraud detection in production financial systems.

Why Graphs for Fraud Detection

Fraudsters exploit relationships that tabular models cannot see:

Fraud PatternGraph SignalTabular Signal
Account takeover ringDense connections between new accountsIndividual transaction anomalies
Money laundering chainSequential transfers through intermediariesHigh-value transactions
Identity theft networkShared devices, IPs, addresses across accountsUnusual login location
Collusion fraudReciprocal transactions between partiesHigh transaction frequency
Synthetic identityIsolated subgraph with few external connectionsThin credit file

Chart

Graph Construction

The first challenge is building a meaningful graph from transactional data:

import torch
import numpy as np
from torch_geometric.data import Data, HeteroData
from typing import Dict, List, Tuple

class FraudGraphBuilder:
    """Build heterogeneous graph from transaction data."""

    NODE_TYPES = ["account", "device", "ip_address", "merchant", "phone"]
    EDGE_TYPES = [
        ("account", "transacts_with", "account"),
        ("account", "uses", "device"),
        ("account", "connects_from", "ip_address"),
        ("account", "pays", "merchant"),
        ("account", "has_phone", "phone"),
    ]

    def build_graph(self, transactions: List[dict],
                   accounts: List[dict]) -> HeteroData:
        """Construct heterogeneous graph from raw data."""
        data = HeteroData()

        # Build node features
        account_features = self._build_account_features(accounts)
        data["account"].x = torch.tensor(account_features, dtype=torch.float32)

        # Build edges for each relationship type
        for edge_type in self.EDGE_TYPES:
            src_type, rel, dst_type = edge_type
            edges = self._extract_edges(transactions, edge_type)
            if edges:
                src_ids, dst_ids = zip(*edges)
                data[src_type, rel, dst_type].edge_index = torch.tensor(
                    [list(src_ids), list(dst_ids)], dtype=torch.long
                )

        # Add edge features (transaction amount, time, etc.)
        tx_edges = self._extract_transaction_edges(transactions)
        data["account", "transacts_with", "account"].edge_attr = torch.tensor(
            tx_edges["features"], dtype=torch.float32
        )

        return data

    def _build_account_features(self, accounts: List[dict]) -> np.ndarray:
        """Compute node-level features for accounts."""
        features = []
        for account in accounts:
            feature_vec = [
                account["account_age_days"] / 365.0,
                account["total_transactions"] / 1000.0,
                account["avg_transaction_amount"] / 10000.0,
                account["unique_merchants"] / 100.0,
                account["unique_devices"] / 10.0,
                account["avg_transactions_per_day"] / 50.0,
                account["max_transaction_amount"] / 50000.0,
                float(account["has_verified_email"]),
                float(account["has_verified_phone"]),
                account["days_since_last_password_change"] / 180.0,
            ]
            features.append(feature_vec)
        return np.array(features)

    def _extract_edges(self, transactions: List[dict],
                      edge_type: Tuple[str, str, str]) -> List[Tuple[int, int]]:
        """Extract edges of a specific type from transactions."""
        src_type, rel, dst_type = edge_type
        edges = set()

        for tx in transactions:
            if rel == "transacts_with":
                edges.add((tx["sender_id"], tx["receiver_id"]))
            elif rel == "uses":
                edges.add((tx["account_id"], tx["device_id"]))
            elif rel == "connects_from":
                edges.add((tx["account_id"], tx["ip_id"]))

        return list(edges)

GNN Model Architecture

A heterogeneous GNN that handles multiple node and edge types:

import torch
import torch.nn as nn
from torch_geometric.nn import HeteroConv, SAGEConv, GATConv, Linear

class FraudDetectionGNN(nn.Module):
    """Heterogeneous GNN for fraud detection."""

    def __init__(self, in_channels: int, hidden_channels: int = 128,
                 out_channels: int = 2, num_layers: int = 3):
        super().__init__()
        self.num_layers = num_layers

        # Heterogeneous convolution layers
        self.convs = nn.ModuleList()
        for i in range(num_layers):
            in_ch = in_channels if i == 0 else hidden_channels
            conv = HeteroConv({
                ("account", "transacts_with", "account"): GATConv(
                    in_ch, hidden_channels, heads=4, concat=False
                ),
                ("account", "uses", "device"): SAGEConv(in_ch, hidden_channels),
                ("account", "connects_from", "ip_address"): SAGEConv(
                    in_ch, hidden_channels
                ),
            }, aggr="sum")
            self.convs.append(conv)

        # Batch normalization per layer
        self.bns = nn.ModuleList([
            nn.BatchNorm1d(hidden_channels) for _ in range(num_layers)
        ])

        # Classification head
        self.classifier = nn.Sequential(
            nn.Linear(hidden_channels, 64),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(64, out_channels),
        )

    def forward(self, x_dict, edge_index_dict):
        for i, conv in enumerate(self.convs):
            x_dict = conv(x_dict, edge_index_dict)
            # Apply batch norm and activation to account nodes
            x_dict = {
                key: self.bns[i](x.relu()) if key == "account" else x.relu()
                for key, x in x_dict.items()
            }

        # Classify account nodes
        out = self.classifier(x_dict["account"])
        return out

Training with Class Imbalance

Fraud is extremely rare (0.1-0.5% of transactions). Handle this through sampling and loss design:

from torch_geometric.loader import NeighborLoader
from sklearn.metrics import precision_recall_curve, average_precision_score
import torch.nn.functional as F

class FraudModelTrainer:
    def __init__(self, model: FraudDetectionGNN, graph: HeteroData):
        self.model = model
        self.graph = graph

        # Focal loss for class imbalance
        self.criterion = FocalLoss(alpha=0.75, gamma=2.0)
        self.optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

    def train_epoch(self):
        self.model.train()

        # Use neighbor sampling for mini-batch training
        loader = NeighborLoader(
            self.graph,
            num_neighbors=[15, 10, 5],  # Neighbors per layer
            batch_size=512,
            input_nodes=("account", self.graph["account"].train_mask),
            shuffle=True,
        )

        total_loss = 0
        for batch in loader:
            self.optimizer.zero_grad()
            out = self.model(batch.x_dict, batch.edge_index_dict)

            # Only compute loss on target (account) nodes
            target_mask = batch["account"].train_mask
            loss = self.criterion(
                out[target_mask],
                batch["account"].y[target_mask],
            )

            loss.backward()
            self.optimizer.step()
            total_loss += loss.item()

        return total_loss / len(loader)


class FocalLoss(nn.Module):
    """Focal loss for handling extreme class imbalance."""

    def __init__(self, alpha: float = 0.75, gamma: float = 2.0):
        super().__init__()
        self.alpha = alpha
        self.gamma = gamma

    def forward(self, inputs, targets):
        ce_loss = F.cross_entropy(inputs, targets, reduction="none")
        pt = torch.exp(-ce_loss)
        focal_weight = self.alpha * (1 - pt) ** self.gamma
        return (focal_weight * ce_loss).mean()

Feature Engineering for Graphs

Feature CategoryExamplesComputation
Node structuralDegree, clustering coefficient, PageRankGraph algorithms
Node behavioralTransaction velocity, amount distributionAggregation
Edge temporalTime between transactions, periodicityTime series
SubgraphConnected component size, densityCommunity detection
Cross-typeShared device count, IP overlapBipartite analysis

Production Inference Pipeline

Real-time fraud scoring requires sub-second inference on dynamic graphs:

class RealTimeFraudScorer:
    """Score transactions in real-time using pre-computed embeddings."""

    def __init__(self, model: FraudDetectionGNN, embedding_cache):
        self.model = model.eval()
        self.cache = embedding_cache

    @torch.inference_mode()
    def score_transaction(self, transaction: dict) -> dict:
        """Score a single transaction for fraud risk."""
        # Get pre-computed node embeddings for involved accounts
        sender_emb = self.cache.get(f"emb:{transaction['sender_id']}")
        receiver_emb = self.cache.get(f"emb:{transaction['receiver_id']}")

        # Construct local subgraph for fresh context
        subgraph = self._extract_local_subgraph(transaction)

        # Run inference on subgraph
        predictions = self.model(subgraph.x_dict, subgraph.edge_index_dict)
        fraud_prob = torch.softmax(predictions, dim=-1)[:, 1]

        sender_risk = fraud_prob[0].item()
        receiver_risk = fraud_prob[1].item() if len(fraud_prob) > 1 else 0.0

        # Combined risk score
        transaction_risk = max(sender_risk, receiver_risk * 0.7)

        return {
            "transaction_id": transaction["id"],
            "fraud_probability": transaction_risk,
            "sender_risk": sender_risk,
            "receiver_risk": receiver_risk,
            "decision": self._make_decision(transaction_risk),
        }

    def _make_decision(self, risk_score: float) -> str:
        if risk_score > 0.85:
            return "block"
        elif risk_score > 0.6:
            return "review"
        elif risk_score > 0.3:
            return "challenge"  # Step-up authentication
        return "allow"

Model Performance Comparison

ModelPrecision@95% RecallAUC-PRDetection Rate (>$1K)False Positive Rate
Rules-based12.4%0.1872%4.2%
XGBoost (tabular)34.2%0.4284%1.8%
GNN (homogeneous)48.6%0.5689%1.2%
GNN (heterogeneous)62.3%0.6893%0.8%
GNN + temporal67.8%0.7295%0.7%

The heterogeneous GNN with temporal features achieves 2x the precision of tabular models at the same recall level.

Key Takeaways

  • Graph structure reveals fraud patterns invisible to tabular models. Account networks, shared devices, and transaction chains are the strongest fraud signals.
  • Heterogeneous GNNs outperform homogeneous ones by 25-30%. Modeling distinct node and edge types captures richer relational patterns.
  • Class imbalance requires focal loss and careful evaluation. Standard accuracy is meaningless at 0.1% fraud rates. Optimize for precision at fixed recall.
  • Pre-computed embeddings enable real-time scoring. Update embeddings in batch (hourly), score individual transactions against cached embeddings in milliseconds.
  • Temporal patterns matter as much as structural ones. The sequence and timing of transactions are powerful features that complement graph topology.

GNN-based fraud detection is the highest-impact application of graph ML in production today. The relational patterns it captures represent a genuine capability gap that simpler models cannot bridge.

Comments

    No comments yet. Be the first to share your thoughts.