Digital Reactor
機械学習

取引の「つながり」は不正検知にどれだけ効くか:GNNとMLPの比較

取引の「つながり」は不正検知にどれだけ効くか:GNNとMLPの比較

はじめに

取引データの不正検知では、口座単体の特徴だけでなく、口座同士のつながり方に手がかりが潜んでいることがあります。金額も頻度も正常の範囲に収まっているのに、つながっている相手が不自然な口座は、単体の特徴量をいくら眺めても見つかりません。グラフニューラルネットワーク(GNN)は、このノード間の関係をそのまま学習に取り込めるモデルです。架空の取引ネットワークを題材に、GNNの代表的な手法であるGraph Convolutional Network(GCN)で異常検知を実装し、エッジ情報(取引関係)の有無で検出性能がどう変わるかを確かめます。先に結果を言うと、同じ特徴量でもつながりの情報を足すだけで、テストデータのAUC(識別性能の指標)が0.10ポイント上がりました。

対象読者:

  • グラフ構造データに対する機械学習に興味がある方
  • 金融取引における不正検知に応用できる技術を探している方
  • GNNの実装を通して、その動作原理を理解したい方

記事のポイント:

  • GNN、特にGraph Convolutional Network (GCN) の基礎と、異常検知への応用を解説
  • 架空の取引ネットワークデータを生成し、GCNによる異常検知モデルを構築・評価
  • エッジ情報(取引関係)の有無で性能を比較(GCN vs MLP)し、エッジ情報の効果を確認
  • モデルの実装にはPyTorch Geometricライブラリを使用

グラフニューラルネットワークの基礎

GNNとは

グラフニューラルネットワーク(Graph Neural Network, GNN)は、グラフ構造を持つデータをそのまま入力できるニューラルネットワークです。通常のニューラルネットワークはノードを1件ずつ独立に扱うため、ノード間のつながりを使いたければ「取引先数」のような集計値に潰して特徴量へ織り込むしかありません。GNNは各ノードの特徴量を近傍ノードの情報で更新するので、つながり方そのものを学習に使えます。ソーシャルネットワーク分析、推薦システム、化学構造の解析などで使われています。

GCNの動作原理

GCN(Graph Convolutional Network)はGNNの代表的な手法です。各ノードがまず、エッジでつながった隣接ノードの特徴量を集めます。次に、集めた情報と自身の特徴量から新しい特徴量を計算し、非線形変換をかけます。この処理を層として重ねると、1層で隣まで、2層で隣の隣までと、層の数だけ遠くのノードの情報が届きます。今回の実装では2層のGCNを使い、2ホップ先までの情報を取り込みます。

1層分の処理を式で書くと次のようになります。

H(l+1)=σ(D~12A~D~12H(l)W(l))H^{(l+1)} = \sigma(\tilde{D}^{-\frac{1}{2}}\tilde{A}\tilde{D}^{-\frac{1}{2}}H^{(l)}W^{(l)})

各記号の意味は次の表の通りです。

記号説明
H(l)H^{(l)}l 層目のノードの特徴量を集めた行列。各行が各ノードの特徴ベクトルに対応します。
A~\tilde{A}自己ループを加えた隣接行列。グラフの接続構造を表します。自己ループは、各ノードが自分自身の情報も考慮することを意味します。
D~\tilde{D}A~\tilde{A} の次数行列。各ノードに接続するエッジの数を表します(自己ループも含む)。
W(l)W^{(l)}学習可能な重み行列。l 層目の学習パラメータであり、特徴量の変換を行います。
σ\sigma非線形活性化関数。ReLUなどが用いられ、ニューラルネットワークに非線形性をもたらします。

つまり、各ノードの新しい特徴量は、隣接ノードと自身の特徴量を混ぜて作られます。D~12A~D~12\tilde{D}^{-\frac{1}{2}}\tilde{A}\tilde{D}^{-\frac{1}{2}} の部分が、隣接ノードからの情報を集めて次数で正規化する役割を担います。

取引ネットワークデータの生成

今回は仮想データとして、次の画像のような構造のグラフを生成します。

ノードの特徴量設計

取引ネットワークの各ノードは口座などの取引主体を表し、次の4次元の特徴量を持たせます。

  1. 取引金額の平均
  2. 取引頻度
  3. 取引先数
  4. 取引の時間的分散

いずれも取引行動を数値化したもので、この4つの組み合わせが各取引主体の行動パターンを表します。

正常ノードと異常ノードの特徴

正常ノードの特徴

正常ノードには、ありふれた取引パターンを持たせます。取引金額は平均的、取引頻度は一定の範囲に収まり、取引先は固定的で、取引タイミングのばらつきも小さい、という設定です。

# 正常ノードの特徴量生成
normal_features = torch.zeros(n_normal, 4)
normal_features[:, 0] = torch.normal(mean=1000.0, std=200.0, size=(n_normal,))  # 取引金額
normal_features[:, 1] = torch.normal(mean=10.0, std=2.0, size=(n_normal,))      # 取引頻度
normal_features[:, 2] = torch.normal(mean=5.0, std=1.0, size=(n_normal,))       # 取引先数
normal_features[:, 3] = torch.normal(mean=2.0, std=0.5, size=(n_normal,))       # 時間的分散

各特徴量は、平均と標準偏差を決めて正規分布から生成します。設定値は次の表にまとめました。

特徴量平均値標準偏差説明
取引金額1000.0200.0一般的な取引規模、適度なばらつき
取引頻度10.02.01日あたりの平均取引回数
取引先数5.01.0定常的な取引関係の数
時間的分散2.00.5取引タイミングの規則性

異常ノードの特徴

異常ノードは3パターン用意しました。それぞれ、実務で想定される不正取引のシナリオに対応させています。

# 異常ノードの特徴量生成(抜粋)
anomaly_features = torch.zeros(n_anomaly, 4)
n_pattern1 = n_anomaly // 3

# パターン1: 大口取引パターン
anomaly_features[:n_pattern1, 0] = torch.normal(mean=5000.0, std=1000.0, size=(n_pattern1,))

各パターンで変える特徴量と、想定する不正シナリオの対応は次の通りです。

パターン特徴量の変更想定シナリオ
大口取引取引金額: 5000.0±1000.0マネーロンダリング、不正な資金移動
多数取引取引頻度: 30.0±5.0 取引先数: 15.0±3.0分散型の不正送金、取引分割による規制回避
不規則取引時間的分散: 8.0±2.0自動化された不正取引、営業時間外取引

エッジの生成ロジック

エッジは取引主体間の取引関係を表します。特徴量だけでなくつながり方にも正常と異常の差が出るよう、エッジの生成ロジックを作り分けます。

正常ノードのエッジ生成

正常ノード間のエッジは、コミュニティ構造ができるように生成します。似た取引パターンを持つノード同士がつながりやすい、という現実の取引ネットワークの性質を模したものです。

# 正常ノード間のエッジ生成(抜粋)
for i in range(n_normal):
    n_edges = int(torch.normal(float(base_edges_per_node), 1.0, size=(1,)).item())
    for _ in range(n_edges):
        j = int(torch.normal(mean=torch.tensor(float(i)), std=torch.tensor(float(n_normal/5))).item()) % n_normal

異常ノードのエッジ生成

異常ノードのエッジは、パターンごとに生成ロジックを変えます。

パターンエッジ数接続先の選択方法
大口取引固定2本正常ノードからランダム
多数取引固定10本正常ノードからランダム
不規則取引1-7本全ノードからランダム

GCNによる異常検知モデル

モデルの実装

実装にはPyTorch Geometricを使います。GCN層を2層重ね、2ホップ先までの近傍情報を集約する構成です。

class GCNAnomalyDetector(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, out_channels)

異常検知の仕組み

このモデルは、ノード自身の特徴量(取引金額や頻度)と、エッジから読み取れる接続の構造(接続数や接続先の性質)を組み合わせて異常を判定します。学習を通じて正常な取引ネットワークのつながり方を覚えておき、そこから外れた接続パターンを持つノードを、特徴量の情報と合わせて拾い上げる仕組みです。

エッジ情報の重要性

エッジ情報の有効性を検証するため、同じデータセットに対してGCN(エッジ情報あり)とMLP(エッジ情報なし)の2つのモデルで比較実験を行いました。

実験設定

項目説明
ノード数300(訓練・テストそれぞれ)
異常ノード比率10%

実験結果

評価指標GCN(エッジ情報あり)MLP(エッジ情報なし)
テストデータAUC0.80480.7004

AUCの差はどこから生まれたか

今回の仮想データでは、GCNがMLPをテストAUCで0.10ポイント上回りました(0.8048対0.7004)。テストデータで差がついている点が要点で、エッジ情報は未知データへの汎化に寄与しています。

差の出どころはデータの作り方から説明できます。異常ノードの特徴量は正常ノードに近づけて生成しているため、特徴量だけを見るMLPには判別しにくいノードが残ります。一方、異常ノードのつながり方は正常ノードと変えてあります。複数コミュニティにまたがる接続や、1つのコミュニティへの集中接続は、GCNが近傍集約を通じて拾える手がかりです。正常側にコミュニティ構造を持たせたことも効いています。コミュニティ内では近傍の特徴が似通うため、そこに紛れ込んだ異常ノードは、集約後の表現で周囲から浮きやすくなります。

まとめ

同じ特徴量を使っても、取引関係のエッジ情報を入れたGCNは、特徴量だけのMLPよりテストAUCで0.10ポイント高くなりました。単体の指標では正常に見える口座でも、つながり方まで見れば検出の余地がある、というのが今回の実験から言えることです。実データに適用するなら、手元の取引データを「どの主体とどの主体の間に取引があったか」のエッジリストに整形し、今回のコードのデータ生成部分を差し替えるところから始めると小さく試せます。実データでは不正の正解ラベルが少ないことが多いので、教師ありのまま進めるか、半教師あり・教師なしの設定に切り替えるかは、確保できるラベルの量を見て決めてください。

コード

import torch
import torch.nn.functional as F
from torch_geometric.nn import GCNConv
from torch_geometric.data import Data
import networkx as nx
import numpy as np
import matplotlib.pyplot as plt
import japanize_matplotlib
from sklearn.metrics import confusion_matrix, accuracy_score, precision_score, recall_score, roc_auc_score
import seaborn as sns
import sklearn.metrics
import argparse

# 引数パーサーの設定
parser = argparse.ArgumentParser(description='GCNを用いた異常検知')
parser.add_argument('--no-edge-info', action='store_true',
                   help='エッジ情報を使用しない(MLPモードで実行)')
args = parser.parse_args()

# モード名の設定(ファイル名用)
mode = "mlp" if args.no_edge_info else "gcn"

# 乱数シードを固定
torch.manual_seed(102)
np.random.seed(102)

def generate_transaction_network(n_nodes=100, base_edges_per_node=3, anomaly_ratio=0.1, use_edge_info=True):
    """取引ネットワークを生成する関数

    Parameters:
    -----------
    n_nodes : int
        ノードの総数
    base_edges_per_node : int
        1ノードあたりの基本エッジ数
    anomaly_ratio : float
        異常ノードの割合
    use_edge_info : bool
        エッジ情報を使用するかどうか
    """
    n_normal = int(n_nodes * (1 - anomaly_ratio))
    n_anomaly = n_nodes - n_normal

    # 正常ノードの特徴量生成(取引パターンを表現)
    # 特徴量の差を小さくするため、標準偏差を大きくする
    normal_features = torch.zeros(n_normal, 4)
    normal_features[:, 0] = torch.normal(mean=1000.0, std=500.0, size=(n_normal,))  # 取引金額
    normal_features[:, 1] = torch.normal(mean=10.0, std=5.0, size=(n_normal,))      # 取引頻度
    normal_features[:, 2] = torch.normal(mean=5.0, std=3.0, size=(n_normal,))       # 取引先数
    normal_features[:, 3] = torch.normal(mean=2.0, std=1.0, size=(n_normal,))       # 時間的分散

    # 異常ノードの特徴量生成(異常パターンを表現)
    # 特徴量を正常ノードに近づける
    anomaly_features = torch.zeros(n_anomaly, 4)
    # パターン1: やや大きな取引金額
    n_pattern1 = n_anomaly // 3
    anomaly_features[:n_pattern1, 0] = torch.normal(mean=2000.0, std=500.0, size=(n_pattern1,))
    anomaly_features[:n_pattern1, 1:] = normal_features[:n_pattern1, 1:]

    # パターン2: やや多い取引頻度と取引先
    n_pattern2 = n_anomaly // 3
    anomaly_features[n_pattern1:n_pattern1+n_pattern2, 0] = normal_features[:n_pattern2, 0]
    anomaly_features[n_pattern1:n_pattern1+n_pattern2, 1] = torch.normal(mean=15.0, std=5.0, size=(n_pattern2,))
    anomaly_features[n_pattern1:n_pattern1+n_pattern2, 2] = torch.normal(mean=8.0, std=3.0, size=(n_pattern2,))
    anomaly_features[n_pattern1:n_pattern1+n_pattern2, 3] = normal_features[:n_pattern2, 3]

    # パターン3: やや大きな時間的分散
    n_pattern3 = n_anomaly - n_pattern1 - n_pattern2
    anomaly_features[n_pattern1+n_pattern2:, 0:3] = normal_features[:n_pattern3, 0:3]
    anomaly_features[n_pattern1+n_pattern2:, 3] = torch.normal(mean=4.0, std=1.0, size=(n_pattern3,))

    # 特徴量を結合
    x = torch.cat([normal_features, anomaly_features], dim=0)

    # 特徴量の正規化
    mean = x.mean(dim=0, keepdim=True)
    std = x.std(dim=0, keepdim=True)
    x = (x - mean) / std

    # エッジの生成(取引関係を表現)
    if use_edge_info:
        edge_list = []

        # 正常ノード間のエッジ生成(コミュニティ構造を強化)
        communities = [[] for _ in range(5)]  # 5つのコミュニティを作成
        nodes_per_community = n_normal // 5

        # ノードをコミュニティに割り当て
        for i in range(n_normal):
            community_idx = i // nodes_per_community
            if community_idx < 5:  # 余りのノードは最後のコミュニティに
                communities[community_idx].append(i)
            else:
                communities[4].append(i)

        # コミュニティ内のエッジを生成(密な接続)
        for community in communities:
            for i in community:
                # コミュニティ内で密な接続
                n_edges = int(torch.normal(float(base_edges_per_node * 2), 1.0, size=(1,)).item())
                for _ in range(n_edges):
                    j = np.random.choice(community)
                    if i != j:
                        edge_list.append([i, j])
                        edge_list.append([j, i])

                # コミュニティ間の疎な接続
                if torch.rand(1).item() < 0.3:  # 30%の確率で他のコミュニティとも接続
                    other_community = communities[np.random.randint(0, 5)]
                    j = np.random.choice(other_community)
                    if i != j:
                        edge_list.append([i, j])
                        edge_list.append([j, i])

        # 異常ノード間のエッジ生成(より特徴的なパターン)
        for i in range(n_normal, n_nodes):
            pattern_type = (i - n_normal) // (n_anomaly // 3)
            if pattern_type == 0:  # 大きな取引金額のパターン
                # 複数のコミュニティと接続
                target_communities = np.random.choice(5, 2, replace=False)  # 2つのコミュニティを選択
                for comm_idx in target_communities:
                    for _ in range(2):  # 各コミュニティから2つのノードと接続
                        j = np.random.choice(communities[comm_idx])
                        edge_list.append([i, j])
                        edge_list.append([j, i])

            elif pattern_type == 1:  # 多数の取引先パターン
                # 1つのコミュニティに集中して接続
                target_community = np.random.randint(0, 5)
                n_edges = 8
                connected_nodes = set()
                for _ in range(n_edges):
                    j = np.random.choice(communities[target_community])
                    if j not in connected_nodes:
                        connected_nodes.add(j)
                        edge_list.append([i, j])
                        edge_list.append([j, i])

            else:  # 不規則な取引パターン
                # ランダムなコミュニティと接続
                n_edges = np.random.randint(1, 4)  # より少ない接続数
                for _ in range(n_edges):
                    target_community = np.random.randint(0, 5)
                    j = np.random.choice(communities[target_community])
                    edge_list.append([i, j])
                    edge_list.append([j, i])

        edge_index = torch.tensor(edge_list).t()
    else:
        # エッジ情報を使用しない場合は、自己ループのみ設定
        edge_index = torch.stack([torch.arange(n_nodes), torch.arange(n_nodes)], dim=0)

    # ラベルの生成
    y = torch.zeros(n_nodes)
    y[n_normal:] = 1

    return Data(x=x, edge_index=edge_index, y=y)

class GCNAnomalyDetector(torch.nn.Module):
    """GCNを用いた異常検知モデル"""
    def __init__(self, in_channels, hidden_channels, out_channels, use_edge_info=True):
        super().__init__()
        if use_edge_info:
            # GCNモード:エッジ情報を使用
            self.conv1 = GCNConv(in_channels, hidden_channels)
            self.conv2 = GCNConv(hidden_channels, out_channels)
        else:
            # MLPモード:エッジ情報を使用しない
            self.conv1 = torch.nn.Linear(in_channels, hidden_channels)
            self.conv2 = torch.nn.Linear(hidden_channels, out_channels)
        self.use_edge_info = use_edge_info

    def forward(self, x, edge_index):
        if self.use_edge_info:
            # GCNモード
            x = self.conv1(x, edge_index)
            x = F.relu(x)
            x = F.dropout(x, p=0.5, training=self.training)
            x = self.conv2(x, edge_index)
        else:
            # MLPモード
            x = self.conv1(x)
            x = F.relu(x)
            x = F.dropout(x, p=0.5, training=self.training)
            x = self.conv2(x)
        return torch.sigmoid(x)

def evaluate_model(model, data):
    """モデルの評価を行う関数"""
    model.eval()
    with torch.no_grad():
        pred = model(data.x, data.edge_index)
        pred_probs = pred.squeeze()  # 予測確率
        pred_labels = (pred_probs > 0.5).float()  # 予測ラベル

        # 評価指標の計算
        accuracy = accuracy_score(data.y.numpy(), pred_labels.numpy())
        precision = precision_score(data.y.numpy(), pred_labels.numpy())
        recall = recall_score(data.y.numpy(), pred_labels.numpy())
        auc = roc_auc_score(data.y.numpy(), pred_probs.numpy())
        conf_matrix = confusion_matrix(data.y.numpy(), pred_labels.numpy())

        return accuracy, precision, recall, auc, conf_matrix, pred_labels

# トレーニングデータとテストデータの生成
train_data = generate_transaction_network(n_nodes=300, base_edges_per_node=3, use_edge_info=not args.no_edge_info)
test_data = generate_transaction_network(n_nodes=300, base_edges_per_node=3, use_edge_info=not args.no_edge_info)

# モデルの初期化と学習
model = GCNAnomalyDetector(in_channels=4, hidden_channels=16, out_channels=1, use_edge_info=not args.no_edge_info)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

print(f"\nモデルモード: {'MLP(エッジ情報なし)' if args.no_edge_info else 'GCN(エッジ情報あり)'}")
print("モデルの学習を開始します...")

# 学習ループ
train_losses = []
test_losses = []
for epoch in range(200):
    # 訓練
    model.train()
    optimizer.zero_grad()
    out = model(train_data.x, train_data.edge_index)
    train_loss = F.binary_cross_entropy(out.squeeze(), train_data.y)
    train_loss.backward()
    optimizer.step()
    train_losses.append(train_loss.item())

    # テストデータでの損失計算
    model.eval()
    with torch.no_grad():
        test_out = model(test_data.x, test_data.edge_index)
        test_loss = F.binary_cross_entropy(test_out.squeeze(), test_data.y)
        test_losses.append(test_loss.item())

    if (epoch + 1) % 20 == 0:
        print(f"Epoch {epoch+1}/200: Train Loss = {train_loss.item():.4f}, Test Loss = {test_loss.item():.4f}")

print("\n学習が完了しました。評価を開始します...")

# 学習曲線の可視化
plt.figure(figsize=(10, 5))
plt.plot(train_losses, label='Training Loss')
plt.plot(test_losses, label='Test Loss')
plt.title(f'Training and Test Loss ({mode.upper()})')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.grid(True)
plt.legend()
plt.savefig(f'learning_curves_{mode}.png')
plt.close()

# モデルの評価
train_metrics = evaluate_model(model, train_data)
test_metrics = evaluate_model(model, test_data)

# 評価結果の表示と保存
def plot_confusion_matrix(conf_matrix, title, filename):
    plt.figure(figsize=(8, 6))
    if conf_matrix.size > 0:  # 混同行列が空でないことを確認
        sns.heatmap(conf_matrix, annot=True, fmt='d', cmap='Blues')
        plt.title(f'{title} ({mode.upper()})')
        plt.ylabel('True Label')
        plt.xlabel('Predicted Label')
        plt.savefig(f'generated/gnn-anomaly-detection/{filename}_{mode}.png')
    plt.close()

# 訓練データの混同行列
plot_confusion_matrix(train_metrics[4], 'Confusion Matrix (Training Data)', 'train_confusion_matrix')

# テストデータの混同行列
plot_confusion_matrix(test_metrics[4], 'Confusion Matrix (Test Data)', 'test_confusion_matrix')

# ネットワークの可視化(テストデータ)
def visualize_network(data, pred_labels, title, filename):
    if not args.no_edge_info:  # エッジ情報がある場合のみ可視化
        G = nx.Graph()

        # すべてのノードを追加
        for i in range(len(data.y)):
            G.add_node(i)

        # エッジを追加
        edge_list = data.edge_index.t().tolist()
        G.add_edges_from(edge_list)

        plt.figure(figsize=(12, 8))
        pos = nx.spring_layout(G)

        # データをNumPy配列に変換
        true_labels = data.y.numpy()
        pred_labels = (pred_labels > 0.5).numpy()  # 予測確率を二値ラベルに変換

        # 正常ノードと異常ノードを別々に描画(予測と実際のラベルを比較)
        true_normal = np.where((true_labels == 0) & (pred_labels == 0))[0]  # 正常を正常と予測
        true_anomaly = np.where((true_labels == 1) & (pred_labels == 1))[0]  # 異常を異常と予測
        false_normal = np.where((true_labels == 1) & (pred_labels == 0))[0]  # 異常を正常と予測
        false_anomaly = np.where((true_labels == 0) & (pred_labels == 1))[0]  # 正常を異常と予測

        nx.draw_networkx_nodes(G, pos, nodelist=true_normal, 
                              node_color='lightblue', node_size=100, label='True Normal')
        nx.draw_networkx_nodes(G, pos, nodelist=true_anomaly,
                              node_color='red', node_size=100, label='True Anomaly')
        nx.draw_networkx_nodes(G, pos, nodelist=false_normal,
                              node_color='orange', node_size=100, label='Missed Anomaly')
        nx.draw_networkx_nodes(G, pos, nodelist=false_anomaly,
                              node_color='purple', node_size=100, label='False Anomaly')
        nx.draw_networkx_edges(G, pos, alpha=0.2)

        plt.title(f'{title} ({mode.upper()})')
        plt.legend()
        plt.savefig(f'{filename}_{mode}.png')
        plt.close()

# トレーニングデータとテストデータのネットワーク可視化
visualize_network(train_data, train_metrics[5], 'Training Data Network', 'train_network')
visualize_network(test_data, test_metrics[5], 'Test Data Network', 'test_network')

# 評価結果をファイルに保存
with open(f'evaluation_results_{mode}.txt', 'w') as f:
    f.write(f'Model: {mode.upper()}\n\n')
    f.write('Training Data Metrics:\n')
    f.write(f'Accuracy: {train_metrics[0]:.4f}\n')
    f.write(f'Precision: {train_metrics[1]:.4f}\n')
    f.write(f'Recall: {train_metrics[2]:.4f}\n')
    f.write(f'AUC: {train_metrics[3]:.4f}\n\n')

    f.write('Test Data Metrics:\n')
    f.write(f'Accuracy: {test_metrics[0]:.4f}\n')
    f.write(f'Precision: {test_metrics[1]:.4f}\n')
    f.write(f'Recall: {test_metrics[2]:.4f}\n')
    f.write(f'AUC: {test_metrics[3]:.4f}\n')

# ROC曲線の描画
def plot_roc_curve(model, data, title, filename):
    model.eval()
    with torch.no_grad():
        pred = model(data.x, data.edge_index)
        pred_probs = pred.squeeze().numpy()
        true_labels = data.y.numpy()

        fpr, tpr, _ = sklearn.metrics.roc_curve(true_labels, pred_probs)

        plt.figure(figsize=(8, 6))
        plt.plot(fpr, tpr, label=f'ROC curve (AUC = {roc_auc_score(true_labels, pred_probs):.4f})')
        plt.plot([0, 1], [0, 1], 'k--')
        plt.xlim([0.0, 1.0])
        plt.ylim([0.0, 1.05])
        plt.xlabel('False Positive Rate')
        plt.ylabel('True Positive Rate')
        plt.title(f'{title} ({mode.upper()})')
        plt.legend(loc="lower right")
        plt.grid(True)
        plt.savefig(f'{filename}_{mode}.png')
        plt.close()

# 訓練データとテストデータのROC曲線を描画
plot_roc_curve(model, train_data, 'ROC Curve (Training Data)', 'train_roc_curve')
plot_roc_curve(model, test_data, 'ROC Curve (Test Data)', 'test_roc_curve')

print("\n評価結果:")
print(f"モデル: {mode.upper()}")
print(f"訓練データ - 精度: {train_metrics[0]:.4f}, 適合率: {train_metrics[1]:.4f}, 再現率: {train_metrics[2]:.4f}, AUC: {train_metrics[3]:.4f}")
print(f"テストデータ - 精度: {test_metrics[0]:.4f}, 適合率: {test_metrics[1]:.4f}, 再現率: {test_metrics[2]:.4f}, AUC: {test_metrics[3]:.4f}")

関連記事

← 技術ブログ一覧へ