资讯动态

GIKT:图卷积知识追踪模型的PyTorch实现与调参指南

发布时间:2026/9/17 18:48:43 来源:尧图企业网站定制
简介围绕基于图卷积网络的知识追踪模型GIKT这份PDF提供了完整的论文内容与实验细节面向在线教育知识追踪方向的研究人员和技术人员旨在解决数据稀疏与多技能关联带来的学生掌握度预测难题。文中详细阐述了GIKT的完整设计嵌入层通过GCN提取高阶问题-技能关系LSTM层捕捉学生的长期行为变化历史回顾模块选择与新题目相关的历史练习广义交互模块综合多种因素进行最终预测在三个公开数据集上均取得AUC至少1个百分点的提升。压缩包内为1个PDF文件大小约412KB涵盖模型结构、数学公式、实验对比与结论推导便于直接阅读和引用。目前已有209人学习浏览适合需要复现或借鉴GIKT设计思路、进一步探索图神经网络与知识追踪结合的研究者参考。1. GIKT图卷积网络进入知识追踪的第一站做在线教育的数据建模绕不开一个问题给定学生的历史答题序列怎么预测他能不能答对下一道题知识追踪Knowledge Tracing就是专门解决这个问题的任务传统 IRT 假设学生能力是固定标量DKT 把学生当作一个 RNN 的隐藏状态但这些方法默认知识点之间是独立的比如“一元二次方程”和“二次函数”在模型里是两个互不相干的下标。实际上教学大纲里有明确的前驱、后继关联GIKTGraph-based Interaction-aware Knowledge Tracing就是把这些关联显式建模成图再用图卷积网络GCN把知识点和习题的表示放进一个结构化的空间里。它不算复杂在常见知识追踪代码上增加一个图卷积分支但换来的是概念表示有语义、冷启动知识点的嵌入有邻居信息。适合已经跑通过 DKT/SAKT、想进一步提高可解释性和长尾概念效果的工程师阅读也适合刚接触教育数据挖掘、想知道 GCN 为什么被引入的学习者。2. 从GCN到知识追踪GIKT建模的四个关键设计2.1 为什么要把知识状态建在图而不是序列上DKT 这类序列模型处理的是“连续交互”的时间依赖学生的能力被编码成一个不断更新的向量但它对题目和知识点的表示是从 one-hot 随机初始化的相互之间没有结构关系。实际学习数据里一个学生做错“求二次函数顶点坐标”极可能也做错“利用判别式判断根的情况”因为它们共享部分概念如果模型只看序列这两个题的知识点被分开编码模型需要大量数据才能自己学出它们相关这在数据稀疏时几乎做不到。GIKT 的出发点很直接把知识点和习题放进一张图让表示学习不依赖大量样本。图上的每个节点是一个实体边是关系——知识点与知识点的先修关系、习题与知识点的考察关系或者学生与知识点的交互关系。图卷积网络的作用是把邻居节点的信息聚合到中心节点上经过一两层卷积后相近的概念在语义空间中自然接近。序列模型负责捕捉学生的瞬时状态图模型负责固化稳定概念结构两个角色分工明确这也是 GIKT 相比 DKT 在 AUC 上有稳定提升的原因。2.2 构建知识点-习题二部图与邻接矩阵常见做法是用“习题-知识点”二部图左边是习题节点右边是知识点节点边表示“这道题考察了该知识点”。一个知识点被多道题考察就与多道题相连一道题通常关联 1 到 3 个知识点。这种图的好处是习题和概念处在同一个图卷积层中习题的表示能与它覆盖的概念相互加强且不需要额外维护一个知识图谱文件。建图方式节点构成边含义适用场景概念-概念图只有知识点先修/包含/相关有教学大纲需要体现概念层级习题-概念二部图习题知识点习题考察知识点开放题库题目标签较规范习题-习题图只有习题相似习题/同卷关系没有知识点标注只有做题行为我一般会优先选二部图因为习题节点的存在让卷积消息能沿着“习题-概念-习题”路径传播等于间接建模了相似题。构建时需要把二部图转成邻接矩阵 (A \in \mathbb{R}^{(QC)\times(QC)})其中 (Q) 是习题数(C) 是知识点数。如果 (A_{ij}1) 表示节点 i 和 j 有边则对称归一化后的传播矩阵是[ \hat{A} D^{-\frac{1}{2}} A D^{-\frac{1}{2}} ]其中 (D) 是度矩阵。归一化是为了防止高连接度节点比如被很多习题引用的大概念主导聚合。实际工程里也可以只对习题节点做消息传播但为了后续代码简单建一个完整的邻接矩阵更省心。2.3 图卷积层如何更新知识点表示一层图卷积做的事情很朴素把自身表示和邻居表示加权求和再过一次线性变换和激活函数即[ H^{(l1)} \sigma(\hat{A} H^{(l)} W^{(l)}) ]其中 (H^{(l)} \in \mathbb{R}^{(QC)\times d}) 是第 (l) 层的所有节点嵌入(W^{(l)} \in \mathbb{R}^{d\times d}) 是线性变换矩阵(\sigma) 常用 ReLU。堆叠两层后每个知识点的表示包含了两跳邻居的信息即“与它关联习题相似的另一道题考察的其他知识点”也会影响它。这个特性让知识追踪进入模型里的概念嵌入不再是孤立的。下面给出一个 PyTorch 实现的最小图卷积层这个模块可以直接嵌入到知识追踪模型里import torch import torch.nn as nn import torch.nn.functional as F class GCNLayer(nn.Module): def __init__(self, in_dim, out_dim, dropout0.2): super().__init__() self.linear nn.Linear(in_dim, out_dim) self.dropout nn.Dropout(dropout) self.norm nn.LayerNorm(out_dim) def forward(self, H, adj_norm): # H: [num_nodes, in_dim] 当前所有节点表示 # adj_norm: [num_nodes, num_nodes] 对称归一化邻接矩阵 ctx torch.sparse.mm(adj_norm, H) # 聚合邻居信息 out self.linear(ctx) # 线性变换 out F.relu(out) out self.dropout(out) return self.norm(out) # LayerNorm稳定训练这里我把归一化后的邻接矩阵作为稀疏矩阵传入避免大图上的全矩阵乘法。torch.sparse.mm直接做邻域聚合效率比密集矩阵乘法高。注意我把 LayerNorm 放在激活和 dropout 之后这是 GCN 里常见但容易被人忽略的细节知识追踪图中节点度数差异大某些概念节点几乎只有一条边聚合后数值尺度不一致LayerNorm 比 BatchNorm 更适合这种需要保持每个节点独立尺度的问题。2.4 交互编码与预测层图卷积得到的是知识点和习题的静态表示学生的动态状态依然需要序列模型来追踪。常见做法是把每个交互步 (t) 对应的习题向量 (q_{emb}) 和该习题关联的知识点向量 (c_{emb}) 拼接起来再与学生上一轮正确性编码 (r_{prev})答对/答错向量拼接送入 LSTM得到隐藏状态 (h_t) 作为学生当前知识掌握程度。预测第 (t1) 步答题正确概率时不能直接用 (h_t) 对所有习题输出 sigmoid因为 DKT 那种对所有习题做内积的做法在习题量大时计算成本高。GIKT 常见的做法是用目标习题向量与当前时刻学生状态做点积logit (h_t target_q_emb.T) (h_t target_c_emb.T) p_correct torch.sigmoid(logit)其中target_q_emb是被预测的那道习题的图卷积输出嵌入。这里同时用了知识点嵌入和习题嵌入等于让“会不会这个知识点”和“做过这类题的经验”共同参与判断。损失函数用二元交叉熵[ \mathcal{L} -\frac{1}{T}\sum_{t1}^{T} [r_t \log p_t (1-r_t)\log(1-p_t)] ]在实际代码里为了避免上一轮正确性标签泄露到下一轮训练时会对r_prev做掩码或错位处理一般把第 (t-1) 步的答案作为第 (t) 步输入的一部分预测第 (t) 步时不使用当前真实答案。这是 DKT 系列实现中最容易出现数据泄露的地方后面训练一节会专门说检查点。3. 用PyTorch从零实现一个可运行的GIKT训练管线3.1 训练数据格式与 DataLoader 设计假设你已经有一份交互日志字段为student_id, question_id, concept_id, is_correct, timestamp。一条交互对应一道习题和一个知识点如果一道题包含多个知识点需要拆成多行或者用question_id到concept_id的映射表在模型内做聚合。为了方便演示这里按每个交互只关联一个知识点处理这是很多公开数据集如 ASSISTments2009 的预训练版本使用的格式。我一般会先把每个学生的交互记录按时间排序再切成长度为 (L) 的固定窗口。对于超过窗口长度的学生用滑动窗截成多个样本不足 (L) 的序列在前面填充 0。注意填充的交互不能参与损失计算需要额外准备一个mask张量。下面是一个精简的 Datasetclass KTDataset(torch.utils.data.Dataset): def __init__(self, samples, num_q, num_c): self.samples samples # list of (seq_q, seq_c, seq_r, target_q, target_r, target_mask) self.num_q num_q self.num_c num_c def __len__(self): return len(self.samples) def __getitem__(self, idx): seq_q, seq_c, seq_r, t_q, t_r, mask self.samples[idx] return { seq_q: torch.tensor(seq_q, dtypetorch.long), seq_c: torch.tensor(seq_c, dtypetorch.long), seq_r: torch.tensor(seq_r, dtypetorch.float), target_q: torch.tensor(t_q, dtypetorch.long), target_r: torch.tensor(t_r, dtypetorch.float), mask: torch.tensor(mask, dtypetorch.bool), }这里target_q是序列每一步要预测的下一条交互的习题 ID所以长度为 (L) 而不仅是最后一个时间点这样每条样本可以提供 (L) 个训练信号数据利用率更高。mask对应每个时间步是否有效填充位置为False。3.2 GIKT 模型主体代码模型结构分成三部分图卷积模块、交互编码模块、序列状态与预测模块。核心代码如下import torch import torch.nn as nn import torch.nn.functional as F class GIKT(nn.Module): def __init__(self, num_q, num_c, emb_dim128, dropout0.3, gcn_layers2, max_len200): super().__init__() self.q_emb nn.Embedding(num_q, emb_dim, padding_idx0) self.c_emb nn.Embedding(num_c, emb_dim, padding_idx0) self.adj_dense None # 由外部设置 self.gcn nn.ModuleList() for _ in range(gcn_layers): self.gcn.append(GCNLayer(emb_dim, emb_dim, dropout)) self.lstm nn.LSTM(emb_dim * 3, emb_dim, batch_firstTrue) self.dropout nn.Dropout(dropout) self.out_layer nn.Linear(emb_dim * 2, emb_dim) def set_graph(self, adj): # adj: [num_qnum_c, num_qnum_c] 归一化后的邻接矩阵 self.adj adj def forward(self, seq_q, seq_c, seq_r, target_q): # 所有节点嵌入 node_emb torch.cat([self.q_emb.weight, self.c_emb.weight], dim0) for layer in self.gcn: node_emb layer(node_emb, self.adj) # 提取当前交互涉及的习题和概念表示 q_emb node_emb[:self.q_emb.num_embeddings] c_emb node_emb[self.q_emb.num_embeddings:] cur_q q_emb[seq_q] # [B, L, d] cur_c c_emb[seq_c] # [B, L, d] r_emb seq_r.unsqueeze(-1) * cur_q # 正确性调制 inp torch.cat([cur_q, cur_c, r_emb], dim-1) h, _ self.lstm(inp) # [B, L, d] # 预测每一步之后的下一题正确率 t_q q_emb[target_q] # [B, L, d] h_proj self.out_layer(self.dropout(torch.cat([h, t_q], dim-1))) logit (h_proj * t_q).sum(-1, keepdimTrue) return torch.sigmoid(logit).squeeze(-1)逻辑说明node_emb把习题与知识点拼成一张图的全部节点并用 GCN 层更新。r_emb用seq_r乘上当前习题嵌入相当于把“上次答对/答错”这个事实作为当前交互的特征输入到 LSTM 时模型能感知到结果反馈。预测部分把 LSTM 当前状态与目标习题嵌入拼接后投影回嵌入维度再与目标习题做点积数值上等同于一个协同过滤的打分。需要注意代码中的self.adj是归一化后的邻接矩阵要传入 GPU 并与node_emb设备一致。也可以用稀疏矩阵torch.sparse_coo_tensor构建以节省显存但每一轮训练前要调用adj adj.to(device)。3.3 训练循环与 AUC 评估GIKT 的评估通常看 AUC 和 ACC因为答对概率本身是数值而不是硬类别。训练循环直接用 BCEWithLogitsLoss但为了避免在填充位置算损失要手动加权from sklearn.metrics import roc_auc_score, accuracy_score def evaluate(model, loader, device): model.eval() preds, labels, masks [], [], [] with torch.no_grad(): for batch in loader: seq_q batch[seq_q].to(device) seq_c batch[seq_c].to(device) seq_r batch[seq_r].to(device) t_q batch[target_q].to(device) p model(seq_q, seq_c, seq_r, t_q) preds.append(p[batch[mask]].cpu().numpy()) labels.append(batch[target_r][batch[mask]].numpy()) preds np.concatenate(preds) labels np.concatenate(labels) return roc_auc_score(labels, preds), accuracy_score(labels, preds 0.5)训练时每个 batch 对seq_r的错位处理已经体现在调用方式里调用model时传入的是上一时刻的答案序列标签是target_r两者来自同一条交互日志但模型内r_emb使用的是当前时刻之前的行为不存在未来信息泄露。还要注意在构建训练样本时target_q通常取seq_q向右移一位后的结果最后一个位置需要填充 0对应 mask 为 False不参与计算。3.4 用小型数据先跑通的 3 个检查点第一次训练我一般选择公开数据集中最少的子集比如只选 20 个学生、200 道题、50 个知识点。检查三点一是维度能跑通显存不炸二是 loss 在前 5 个 epoch 明显下降如果不降说明学习率太大或 GCN 层后嵌入被规范化到同一尺度三是评估时 AUC 不低于 0.55低于这个值说明图结构或输入特征有问题需要回到数据构建检查。4. 训练GIKT必调的5个参数与常见失败模式4.1 参数表与推荐范围GIKT 比 DKT 多出的参数主要在 GCN 层其余参数继承自动编码器的常见经验。下面这张表是我在不同教育数据集上调参后得到的稳定区间参数推荐范围默认经验值影响嵌入维度 emb_dim64 ~ 256128太小学不到概念关系太大容易过拟合且训练慢GCN 层数1 ~ 32超过 2 层容易过平滑节点表示趋同Dropout0.2 ~ 0.50.3图卷积后加入 dropout防止邻居信息过强学习率3e-4 ~ 1e-35e-4Adam 优化器下建议从 1e-3 开始衰减序列长度15 ~ 20050长度影响 LSTM 展开步数太长训练慢且梯度不稳定Batch size32 ~ 256128受显存限制与序列长度共同决定占用第一轮调参时先固定 GCN 层数为 2把学习率设为 5e-4只调嵌入维度和 dropout。观察验证集 AUC 曲线如果出现大幅震荡说明学习率需要降如果 AUC 在第 3 个 epoch 就封顶且不再上升说明模型容量不足以表达图结构应增大嵌入维度。4.2 图卷积层数不要盲目加到3以上很多人第一次接触 GIKT 时会觉得“图卷积层数越多概念聚合越充分”实际结果正相反。图中节点度分布极不均匀像“代数”“函数”这种上层概念可能连接了几百个题目子节点而“祖暅原理”这种冷门知识点只有几条边。GCN 每多一层每个节点都会吸收更大范围内的信息两层之后高阶邻居的噪声远大于信号所有概念嵌入逐渐变得相似这就是过平滑。判断是否过平滑的方法是训练完成后取出知识点嵌入计算两两余弦相似度的标准差如果标准差小于 0.1说明嵌入几乎全部挤在一起需要减少层数或增加残差连接。如果业务上确实需要捕捉多跳关系正确做法不是加层数而是在 GCN 层外加上一层注意力聚合或者在每层输出上拼接原始嵌入做残差。后者实现成本最低out relu(adj H W) H这一行能有效缓解过平滑同时保持两层甚至三层卷积的语义聚合能力。4.3 测试期冷启动知识点被“淹没”的问题线上推理时经常出现新知识点或新习题它们在图里没有边。GCN 的聚合公式决定了没有邻居的节点经过卷积后只会保留自身变换如果该节点参数是随机初始化且没有参与训练预测质量就会很差。常见做法是给所有节点额外加一个“虚拟根节点”让冷启动节点与虚拟根节点建边根节点有自己的可训练向量这样新节点至少能聚合到一个有意义的全局表示。另一种更轻量的方案是直接把 cold-start 节点按文本描述映射到最近邻已知知识点例如用课程名称的 TF-IDF 向量做余弦相似度将相似度大于 0.8 的邻居提前固定到邻接矩阵中。这样做不会让 GIKT 完全失效只是要记住图结构一旦更新需要重新做归一化并重新训练或者微调最后一层。4.4 评估指标的选择和时间划分知识追踪论文里常见 AUC 和 ACC 混用但实际业务里这两个指标经常不一致当数据中正负样本比例为 7:3 时分类阈值调到 0.5 往往能获得 70% 以上的 ACC但 AUC 可能只有 0.6。我觉得只看 ACC 会被先验分布欺骗至少同时报告 AUC 和 RMSE更严格的做法是报告每个知识点分组的 AUC。另外数据划分不能随机洗牌应该按时间切分——前 60% 时间内的交互作为训练中间 20% 作为验证最后 20% 作为测试。随机划分会把同一学生的学习习惯分割到训练和测试集里导致测试集看到“未来的答题”而虚高评估结果。下面的划分代码可以复用到你的数据流水线中import pandas as pd from sklearn.model_selection import TimeSeriesSplit df df.sort_values(timestamp) split_idx len(df) * 0.6 train, val_test df.iloc[:int(split_idx)], df.iloc[int(split_idx):] val_idx len(val_test) * 0.5 val, test val_test.iloc[:int(val_idx)], val_test.iloc[int(val_idx):]注意这里按时间整体切分而不是按学生切分因为 GIKT 的学生状态在推理时会延续同一学生在训练和测试中出现不算作弊但同一时间窗内被随机打乱才会破坏时间依赖。5. GIKT落地中的3个技巧特征融合、早停验证、多任务输出最后一个应用阶段我一般关注的不再是模型结构而是如何让 GIKT 在真实题库里更稳定地工作。第一个技巧是在预测层融合答题耗时。GIKT 的交互编码目前只使用了正确性而实际做题过程中一个花了 10 分钟才答对的题目和一个 2 秒答对的题目知识掌握程度完全不同。可以在r_emb处把耗时归一化后作为一个连续特征并乘以习题嵌入time_feat normalized_time.unsqueeze(-1) * cur_q然后替换掉原来的r_emb部分。这样模型既能区分“快速正确”和“勉强正确”也为后续个性化推荐提供更多信号。注意耗时数据的清洗超过 30 分钟或低于 3 秒的通常视为异常值应截断。第二个技巧是早停验证时不要只盯着整体 AUC。在训练过程中每 3 个 epoch 记录一次每组知识点人数大于 10 的知识点 AUC并计算这些分组 AUC 的最小值。如果整体 AUC 还在提升但小样本知识点 AUC 开始掉说明图卷积层输出的概念表示分布变得偏向高频知识点早停应该以“小知识点 AUC 不再下降”为条件而不是整体 AUC。我一般会在验证代码中维护一个最小 AUC 队列连续 5 次验证不提升就停止训练这样能显著减少过拟合。第三个技巧是使用多任务输出。GIKT 的预测层现在只输出一个正确率但真实场景中还需要预测学生放弃作答的概率或者预测下一题知识点转移。可以在 LSTM 隐藏状态上增加两个并行分类头一个输出正确率一个输出“会不会跳过本题”的概率。训练时两个任务共享 LSTM 主参数损失为两个 BCE 的加权和权重可以设为 1:0.3。这样做的好处是当正确率预测的梯度不稳定时回答时间的辅助任务能提供额外的监督信号在冷启动学生的早期交互阶段反而能提升正确率 AUC 约 1~2 个百分点。实现时只需要在模型 forward 中增加一个skip_head nn.Linear(emb_dim, 1)在输出阶段与正确率并列输出即可。落到工程上GIKT 并不是一个花哨的大模型它的价值在于把知识结构显式注入到学习序列中。如果你正在做的场景有可靠的知识点标注、习题量大于 5000、且学生行为序列足够长GIKT 值得替换 DKT 作为基线。如果上游没有知识点标签构建一个概念图反而会变成负担先考虑用自编码器从文本里抽知识概念再跑 GIKT 的图卷积分支也不迟。最后一个小建议把训练好的知识点嵌入导出成文件在访谈业务老师时用余弦相似度展示“哪些概念被模型视为相近”这种可视化通常比 AUC 提升更能说服团队接受新模型。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价