StellarGraph中使用HinSAGE做链接预测输出NaN的问题求助
HinSAGE链接预测输出NaN排查方向
- 优先检查输入特征合法性
先打印graph.node_features("person")和graph.node_features("products"),检查是否存在NaN、无穷大值,特征输入含非法值会直接导致前向传播输出异常。 - 修正图输入与边切分逻辑的对应关系
你当前代码中初始化HinSAGELinkGenerator时传入的是原始全量图G = graph,而非切分掉测试集、训练集正边后的graph_train。这会导致严重的数据泄露,且切边逻辑完全失效,训练时模型可直接获取待预测的边信息,极易出现梯度异常输出NaN。需将生成器初始化改为传入graph_train。 - 调整优化器学习率
Adam默认学习率1e-3对于异质图场景可能过高,易引发梯度爆炸输出NaN,可尝试将学习率下调至1e-4或1e-5:optimizer=optimizers.Adam(learning_rate=1e-4) - 校验边的顺序与头节点类型匹配性
你设置的head_node_types=["person", "product"]要求输入的边列表中每一条边的第一个节点是person类型、第二个是product类型。需检查edges_train、edges_test的边顺序是否符合要求,若存在顺序颠倒的边,会导致特征维度匹配错误,计算输出NaN。 - 检查切边后是否存在孤立节点
调用graph_train.isolates()检查切分边后是否存在完全没有连接的person或product节点,若存在孤立节点,HinSAGE采样邻居时无法获取有效特征,会导致输出异常。 - 验证负边采样有效性
EdgeSplitter生成负边时若出现类型不匹配的负边(比如生成了person-person类型的负边而非person-product类型),也会引发计算错误,可打印部分负边验证边类型是否符合要求。
内容的提问来源于stack exchange,提问作者Bjørn Øst Hansen
相关产品推荐
相关产品推荐

