You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于GNN的欺诈检测项目:矩阵乘法维度错误排查求助

解决GCN维度错误及边级欺诈检测适配建议

一、维度错误的直接修复

错误根源

你手动初始化线性层参数时形状完全不匹配:
你定义的GCNLayer中,self.projection = nn.Linear(c_in=6, c_out=210)要求权重矩阵形状为(210, 6)(线性层权重的标准形状是(输出维度, 输入维度)),偏置向量形状为(210,),但你手动赋值了2x2的权重和2元素的偏置,导致210x6的节点特征与2x2的权重无法进行矩阵乘法,触发维度错误。

另外,代码中使用torch.bmm(批量矩阵乘法)处理2D张量也是错误的,bmm要求输入为3D张量,2D张量应使用torch.matmul或@运算符。

修正后的代码

import torch
import torch.nn as nn

class GCNLayer(nn.Module):
    def __init__(self, c_in, c_out):
        super().__init__()
        self.projection = nn.Linear(c_in, c_out)

    def forward(self, node_feats, adj_matrix):
        # 统计每个节点的邻居数量,避免后续除以0
        num_neighbours = adj_matrix.sum(dim=-1, keepdims=True)
        num_neighbours = torch.clamp(num_neighbours, min=1)
        
        # 节点特征投影
        node_feats = self.projection(node_feats)
        # 邻接矩阵与投影后节点特征做矩阵乘法
        node_feats = torch.matmul(adj_matrix, node_feats)
        # 邻居特征平均
        node_feats = node_feats / num_neighbours
        return node_feats

# 初始化GCN层,输入特征维度6,输出维度210
layer = GCNLayer(c_in=6, c_out=210)

# 若需自定义初始化参数,保证形状匹配(可选)
# layer.projection.weight.data = torch.randn(210, 6)
# layer.projection.bias.data = torch.zeros(210)

with torch.no_grad():
    # 假设node_features_tensor和adjacency_matrix_tensor已提前定义
    out_feats = layer(node_features_tensor, adjacency_matrix_tensor)

print("Output features shape:", out_feats.shape)  # 输出应为torch.Size([210, 210])

二、适配边级欺诈检测任务

你的标签是边级的(torch.Size([200, 1])),但当前GCN输出的是节点特征,需要额外步骤将节点特征转化为边特征,再进行欺诈预测:

  1. 维护边索引张量edge_index(形状(2, 200)),存储每条边的源节点、目标节点索引
  2. 从GCN输出的节点特征中提取每条边的两个节点特征,通过拼接、点积、差分等方式生成边特征
  3. 结合原始边特征,接入分类头完成欺诈预测

示例代码片段

# 假设edge_index是(2, 200)的张量,存储每条边的源/目标节点索引
# 提取边对应的两个节点特征
src_feats = out_feats[edge_index[0]]  # 形状(200, 210)
dst_feats = out_feats[edge_index[1]]  # 形状(200, 210)

# 生成边特征:以拼接节点特征为例,也可使用点积、差分等方式
edge_feats = torch.cat([src_feats, dst_feats], dim=1)  # 形状(200, 420)
# 结合原始边特征(可选)
edge_feats = torch.cat([edge_feats, edge_features_tensor], dim=1)  # 形状(200, 424)

# 欺诈预测分类头
classifier = nn.Sequential(
    nn.Linear(424, 128),
    nn.ReLU(),
    nn.Linear(128, 1),
    nn.Sigmoid()
)

preds = classifier(edge_feats)  # 形状(200, 1),与标签维度匹配

内容的提问来源于stack exchange,提问作者Marie-Lyne Roustom

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 03:14:58