目录

GCN Node Classification on Cora (Jittor)

赛道一热身赛:基于图卷积网络(GCN)的 Cora 半监督节点分类

本仓库实现了一个两层 GCN 模型,在 Cora 引用网络数据集上完成半监督节点分类任务。代码基于 JittorJittorGeometric


目录结构

.
├── configs/            # 训练/评测配置文件(YAML)
│   └── default.yaml
├── src/                # 核心源码(模型 / 数据 / 训练 / 评测)
│   ├── config.py
│   ├── dataset.py
│   ├── model.py
│   ├── train.py
│   ├── eval.py
│   └── main.py         # 入口脚本
├── scripts/            # 运行脚本
│   ├── train.sh
│   └── eval.sh
├── data/               # 数据说明(原始数据文件不提交)
│   └── README.md
├── outputs/            # 日志、权重、预测结果(默认 git 忽略)
├── requirements.txt
├── .gitignore
└── LICENSE

1. 环境安装

  • Python: 3.8+
  • 依赖: Jittor、JittorGeometric、PyYAML、numpy

安装步骤:

# 1) 安装 Jittor(参考官方文档:https://github.com/Jittor/jittor)
pip install jittor

# 2) 安装 JittorGeometric 及其依赖
#    详见:https://github.com/AlgRUC/JittorGeometric?tab=readme-ov-file#installation
pip install jittor-geometric

# 3) 安装其他 Python 依赖
pip install -r requirements.txt

2. 数据准备

数据集文件 cora.pkl(pickle 格式,约 15 MB)需要手动放入 data/ 目录:

data/
  cora.pkl    # <-- 将文件放在这里

cora.pkl 包含以下字段:

字段 类型 / Shape 说明
x numpy.ndarray (2708, 1433) 节点特征矩阵(词袋)
y numpy.ndarray (2708,) 节点标签(测试集为 -1
edge_index numpy.ndarray (2, E) COO 格式的边列表
train_mask numpy.ndarray (2708,) bool 训练集掩码
val_mask numpy.ndarray (2708,) bool 验证集掩码
test_mask numpy.ndarray (2708,) bool 测试集掩码
num_classes int 类别数(7)
num_features int 特征维度(1433)

数据根目录通过 --data-path 参数配置(默认 data/cora.pkl)。


3. 训练

使用默认配置直接训练:

bash scripts/train.sh

或直接用 Python 并自定义参数:

python -m src.main \
    --config configs/default.yaml \
    --data-path data/cora.pkl \
    --hidden-dim 256 \
    --dropout 0.8 \
    --epochs 200 \
    --lr 0.01 \
    --weight-decay 5e-4 \
    --seed 42

训练结束后,产物保存在 outputs/

  • best_ckpt.pkl — 最佳验证准确率对应的模型权重
  • config.yaml — 本次运行实际使用的完整配置
  • command.txt — 启动命令
  • train.log — 训练日志

4. 评测 / 推理

使用训练好的 checkpoint 生成测试集预测结果:

bash scripts/eval.sh --ckpt outputs/best_ckpt.pkl

或:

python -m src.main --ckpt outputs/best_ckpt.pkl

预测结果保存为 outputs/result.json,格式为 {节点编号: 预测类别}


5. 结果说明

  • 指标: 节点分类准确率(Accuracy),即预测正确的节点数占对应集合节点总数的比例。
  • 计算方式: argmax(logits) 得到预测类别,与真实标签逐元素比较。
  • 参考结果: 在默认配置(seed=42,200 epochs)下,验证集准确率约为 0.80 左右(具体数值与运行环境有关)。
  • 与线上提交的差异: 本地评测使用固定的 test_mask,与比赛线上提交可能存在微小差异,原因在于随机种子、硬件精度以及数据版本。如需完全复现线上结果,请以官方提交系统的输出为准。

6. 可复现性

  • 所有随机种子通过 --seed 统一设置(默认 42),并在代码中调用 jt.misc.set_global_seed()
  • 每次运行会自动保存:
    • 实际使用的配置 → outputs/config.yaml
    • 启动命令 → outputs/command.txt
    • 训练日志 → outputs/train.log
  • 命令行参数优先级高于配置文件(configs/default.yaml)。
  • 若数据文件缺失或路径错误,脚本会打印明确的错误提示与修复方法。

7. 模型架构

  • 模型: 两层 GCN(GCNConv
  • 隐藏层维度: 256
  • Dropout: 0.8(作用于第一层之后)
  • 激活函数: ReLU
  • 优化器: Adam(lr=0.01, weight_decay=5e-4
  • 训练轮数: 200
  • 特征预处理: 行归一化(使每行特征和为 1)
  • 边预处理: GCN 标准归一化(加自环,对称归一化 D1/2AD1/2D^{-1/2} A D^{-1/2}

8. 第三方引用与声明

  • Jittor — 清华大学动态编译深度学习框架(Apache-2.0)
  • JittorGeometric — Jittor 图神经网络库
  • Cora 数据集 — 经典引用网络基准数据集(LDC 许可,由 planspace.org 整理)

License

本项目基于 MIT License 开源。

关于

第六届计图人工智能挑战赛热身赛赛道一代码实现

349.0 KB
邀请码
    Gitlink(确实开源)
  • 加入我们
  • 官网邮箱:gitlink@ccf.org.cn
  • QQ群
  • QQ群
  • 公众号
  • 公众号

版权所有:中国计算机学会技术支持:开源发展技术委员会
京ICP备13000930号-9 京公网安备 11010802047560号