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

StellarGraph无监督GraphSage示例运行报错求助

问题排查与解决方案

错误原因分析

这个错误本质是GraphSAGE无监督模型的输入结构与生成器输出不匹配:

  • 无监督GraphSAGE训练完成后,用于提取embedding的模型(embedding_model)需要接收多个输入张量(对应中心节点特征、各层邻居采样的特征),但你的node_gen生成器只输出了单个张量,导致模型输入不匹配。
  • 大概率是StellarGraph与TensorFlow版本不兼容,或者生成器配置、模型构建的步骤出现偏差。

具体解决方案

1. 安装兼容版本的StellarGraph和TensorFlow

版本不兼容是这类输入不匹配问题的常见诱因,建议锁定经过验证的版本组合:

pip install stellargraph==1.2.1 tensorflow==2.4.1

安装完成后重启运行环境,确保版本生效。

2. 修正生成器与模型构建的关键步骤

确保整个流程的输入输出结构完全匹配:

  • 创建UnsupervisedSampler:
    from stellargraph.data import UnsupervisedSampler
    
    # 采样参数需与后续GraphSAGE的层数对应
    sampler = UnsupervisedSampler(
        G, 
        nodes=G.nodes(), 
        length=5, 
        number_of_walks=2
    )
    
  • 创建GraphSAGENodeGenerator:
    from stellargraph.mapper import GraphSAGENodeGenerator
    
    # num_samples的长度要和GraphSAGE的layer_sizes长度一致
    generator = GraphSAGENodeGenerator(
        G, 
        batch_size=50, 
        num_samples=[10, 5]  # 对应2层GraphSAGE,每层采样10、5个邻居
    )
    
  • 构建无监督GraphSAGE模型并提取embedding模型:
    from stellargraph.layer import GraphSAGE, link_classification
    from tensorflow.keras.models import Model
    
    graphsage = GraphSAGE(
        layer_sizes=[128, 128],  # 层数与num_samples长度一致
        generator=generator,
        bias=True,
        dropout=0.5,
    )
    x_inp, x_out = graphsage.in_out_tensors()
    
    # 构建训练用的链接分类模型
    prediction = link_classification(
        output_dim=1,
        output_act="sigmoid",
        edge_embedding_method="ip"
    )(x_out)
    train_model = Model(inputs=x_inp, outputs=prediction)
    
    # 构建用于提取embedding的模型(关键:复用x_inp和x_out)
    embedding_model = Model(inputs=x_inp, outputs=x_out)
    
  • 创建节点embedding的生成器:
    node_gen = generator.flow(G.nodes())
    
    这里的generator.flow(G.nodes())会自动生成模型需要的多输入张量结构(中心节点+各层邻居特征)。

3. 调整predict调用参数

避免因多线程导致的输入解析问题,先尝试单worker模式:

node_embeddings = embedding_model.predict(node_gen, workers=1, verbose=1)

运行正常后再逐步调高workers数值。

验证方法

运行以下代码检查生成器输出结构是否符合要求:

sample = next(iter(node_gen))
print(f"生成器输出的张量数量: {len(sample)}")
print(f"每个张量的形状: {[t.shape for t in sample]}")

正常情况下,输出的张量数量应该等于len(num_samples) + 1(中心节点 + 每层邻居),比如num_samples=[10,5]时,应该有3个输入张量,与模型的输入层数量对应。

内容的提问来源于stack exchange,提问作者Shradhit Subudhi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 03:16:08