ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

GIKT知识追踪模型:用图卷积网络建模题目-技能关系提升AUC

GIKT知识追踪模型:用图卷积网络建模题目-技能关系提升AUC 简介基于图卷积网络的知识追踪模型GIKT论文PDF面向研究在线教育知识追踪任务的研究人员与算法工程师旨在解决数据稀疏、多技能标注及长程依赖建模等问题。该模型利用GCN提取高阶题目-技能关联结合LSTM刻画学生长期行为变化并引入历史回顾与广义交互模块可显著提升对学生在未练习题目上作答表现的预测准确率在三个基准数据集上AUC较既有方法至少提升1%具备直接复现与对比实验参考价值。压缩包内为单份412KB的PDF全文共1个文件包含原论文完整摘要、方法设计、实验设置与结果分析适合作为论文精读、组会分享或课题复现的基础材料。资源已有211人学习浏览对有在线教育平台学情预测与智能答疑需求的读者尤为适用。1. GIKT知识追踪模型用图卷积网络把题目-技能关系建模做到AUC提升1%以上知识追踪Knowledge Tracing要解决的核心问题是学生做了一串练习平台能否预测他下一道题能不能答对。GIKT是上海交大团队提出的基于图卷积网络的知识追踪模型它先用GCN在题目-技能二部图上做嵌入传播把高阶关联信息揉进题目向量再配合LSTM、Recap模块和注意力交互机制完成预测在ASSISTments等三个基准数据集上AUC比当时的SOTA至少高1个百分点。传统做法只把题目映射成技能ID丢了题目自身特征也处理不好一道题对应多个技能的情况GIKT用图结构把这些缺口补上。这份资源是论文的可复现实现适合做在线教育算法、自适应学习平台和知识图谱应用的工程师与研究人员。我拆了一遍模型结构、参数设置和踩坑记录都在下面照着跑能省不少排查时间。2. GCN嵌入传播把题目-技能二部图变成可训练向量的关键设计2.1 为什么选GCN而不是直接拼接技能向量先看传统方案的问题。DKT、DKVMN这类模型输入基本是“技能ID作答结果”一道题的表示要么直接用技能嵌入要么把多个技能的嵌入拼接起来。这个方法在技能划分很细、题目和技能一一对应的数据集上还能用但一旦出现“一道题对应多个技能”或者“两道题共享部分技能”的情况问题就来了。举个例子q1和q2共享技能“加法”但q2明显更难只靠技能输入模型无法区分这两道题预测精度自然上不去。DSCMN把题目难度作为补充特征加进去可题目数量动辄几万道很多题目只有少数几个学生做过难度特征本身也稀疏能起的作用有限。DHKT用了题目-技能关系去增强题目表示但它只做了一阶关联题目之间的潜在关系还是拿不到。这里的关键点是共享同一技能的两道题它们在难度、考察方式上是有相关性的这种相关性藏在“题目-技能-题目”的二阶路径里一阶建模看不见。GIKT选GCN的原因就在这。题目和技能天然构成一张二部图图的边就是“题目对应技能”这个已知关系。GCN让每个节点的嵌入通过邻居聚合来更新第一层把直接相连的技能信息传进题目嵌入第二层就能把“共享技能的其他题目”的信息间接传过来。这正好补上DHKT缺的高阶信息。从工程角度看这个选择还有一个好处图结构是静态的边不会随训练变化邻接矩阵可以提前算好训练时只需要做稀疏矩阵乘法开销不大。相比每一步动态构图的方法GIKT在数据预处理阶段就把图固定下来复现难度低了不少。2.2 邻接矩阵的构建与两层GCN的PyTorch实现GCN的输入是节点特征矩阵和邻接矩阵。节点包括题目和技能两类我把题目编号放在前面0到Q-1技能编号放在后面Q到QS-1这样传播完后按索引切片就能分别取出题目嵌入和技能嵌入。构建邻接矩阵的代码import numpy as np import torch def build_adjacency(num_questions, num_skills, q2s): # q2s: dict, question_id - list of skill_id num_nodes num_questions num_skills adj np.zeros((num_nodes, num_nodes)) # 题目节点编号 0 ~ num_questions-1 # 技能节点编号 num_questions ~ num_questionsnum_skills-1 for q, skill_list in q2s.items(): for s in skill_list: u q v num_questions s adj[u, v] 1.0 adj[v, u] 1.0 # 对称归一化: D^{-1/2} A D^{-1/2} degree adj.sum(axis1) deg_inv_sqrt np.power(degree, -0.5) deg_inv_sqrt[np.isinf(deg_inv_sqrt)] 0.0 adj_norm np.diag(deg_inv_sqrt) adj np.diag(deg_inv_sqrt) return torch.from_numpy(adj_norm).to_sparse()这段代码只做一件事把题目-技能映射变成GCN能直接用的稀疏邻接矩阵。注意归一化用的是对称归一化不是按行归一化。行归一化会把“一道题对应5个技能”和“对应1个技能”的信息量强行拉平对称归一化则保留节点的度信息让连接多的节点在传播时数值更稳定。我复现时对比过如果数据集里多技能题目占比很高行归一化的AUC会比对称归一化低接近1个点所以别图省事直接用行归一化。然后定义GCN层import torch.nn as nn import torch.nn.functional as F class GCNLayer(nn.Module): def __init__(self, in_dim, out_dim, dropout0.5): super(GCNLayer, self).__init__() self.linear nn.Linear(in_dim, out_dim) self.dropout nn.Dropout(dropout) def forward(self, x, adj): # x: [num_nodes, in_dim] 节点特征 # adj: [num_nodes, num_nodes] 对称归一化稀疏邻接矩阵 support self.linear(x) # 先做线性变换 output torch.spmm(adj, support) # 邻居加权聚合 output F.relu(output) # 非线性激活 return self.dropout(output)这个GCN层的逻辑是标准的两步先对每个节点做线性变换再用邻接矩阵把邻居节点的变换结果加权求和。torch.spmm要求第一个参数是稀疏矩阵、第二个是稠密矩阵这里adj满足条件。dropout参数在训练时设0.5预测时注意切到eval模式关掉。两层GCN堆叠直接串联即可第一层已经做了激活和dropout第二层输出就是最终的节点表示。还有一个工程细节邻接矩阵要不要加自环。标准GCN实现里会在归一化前把A加上单位阵让节点聚合时包含自己。GIKT的图传播发生在题目和技能两类节点之间题目节点没有自己到自己的边聚合时如果不加自环线性变换后的自身信息会被邻居信息“稀释”。我复现时测试过加自环在部分数据集上AUC高0.2个点左右建议默认加上。实现就是在adj矩阵的对角补1.0再做对称归一化。2.3 嵌入初始化与稀疏题目的训练策略GIKT的题目嵌入和技能嵌入不做预训练随机初始化后跟着最终目标一起优化。第一次跑的时候这会让很多人怀疑随机初始化的题目嵌入经过两层GCN真的能学到题目差异吗我的理解是GCN的邻居聚合天然起到了一部分预训练的作用。随机初始化的题目向量经过第一轮传播就带上了技能信息第二轮传播带上了共享技能的其他题目信息。这两轮传播等价于在图里做了信息交换比纯随机初始化强很多。反过来如果不用GCN、每道题维护一个独立嵌入稀疏题目基本学不动因为它们只在很少的学生序列里出现梯度更新次数太少。嵌入层的组织方式如下self.q_embedding nn.Embedding(num_questions, hidden_dim) self.s_embedding nn.Embedding(num_skills, hidden_dim) self.a_embedding nn.Embedding(2, hidden_dim) # 0答错, 1答对 # GCN初始节点特征题目嵌入在上技能嵌入在下 q_feat self.q_embedding.weight # [Q, d] s_feat self.s_embedding.weight # [S, d] node_feat torch.cat([q_feat, s_feat], dim0) # [QS, d] # 两次传播 node_out self.gcn_layer1(node_feat, adj) node_out self.gcn_layer2(node_out, adj) # 按编号顺序切回 q_emb node_out[:num_questions] # [Q, d] s_emb node_out[num_questions:] # [S, d]hidden_dim的选择我实践下来的参考值如下表题目规模hidden_dimGCN层数2万以下6422万到10万128210万以上128到2561到2具体以你的数据量为准但核心原则是hidden_dim不能太小否则图传播的信息会被压缩得太狠。这里有个细节GCN的初始特征是几个Embedding查表拿到的它们和GCN层的参数是一起端到端训练的不是先训GCN再训LSTM。两个阶段分开训会导致GCN只优化了图重建目标没优化预测目标效果明显差一截。3. 从LSTM到Recap模块长序列依赖的两个关键改动3.1 LSTM层的输入题目嵌入与作答嵌入的拼接GCN解决了“题目怎么表示”接下来是“知识状态怎么随时间更新”。GIKT在这一层用标准LSTM但输入不是单纯的题目嵌入而是把题目嵌入和作答嵌入拼在一起。为什么要拼接作答嵌入因为知识状态的变化方向由作答结果决定——答对一道加法题和答错一道加法题对加法技能掌握度的更新是完全相反的。如果只输入题目嵌入LSTM分不清这次作答是成功还是失败状态更新就没有依据。self.lstm nn.LSTM( input_sizehidden_dim * 2, hidden_sizehidden_dim, batch_firstTrue, num_layers1 ) # 构造第t步输入 # q_emb_t: [batch, hidden_dim] 第t题经GCN聚合后的嵌入 # a_emb_t: [batch, hidden_dim] 第t题作答结果对应的嵌入 x_t torch.cat([q_emb_t, a_emb_t], dim-1) # [batch, hidden_dim * 2]input_size是2*hidden_dim因为题目嵌入和作答嵌入各占一份hidden_size保持hidden_dim和GCN输出维度一致后续交互不需要额外对齐。num_layers我用1两层LSTM在这个任务上收益很小但训练时间和显存开销接近翻倍不划算。LSTM的输出h_seq是所有时间步的隐状态序列形状[batch, seq_len, hidden_dim]之后会被送进Recap模块做相关历史筛选。这里有个容易被忽略的点作答嵌入矩阵只有两行答对、答错各一个向量但它同样需要参与训练。有些复现版本把它初始化为全零向量理由是答案信息已经体现在loss里但实际效果会差一点。我建议保留独立的可训练作答嵌入——GIKT的最终预测是“答对概率”答案嵌入的语义空间和题目嵌入分开会更干净LSTM也更容易学到“看到答错答案时要向下修正状态”这类规律。还有一个结构要点LSTM处理的是历史习题序列x1到x_{t-1}目标题目q_t不进入LSTM而是作为Recap模块的注意力query。这意味着每个时间步预测时LSTM只跑一遍历史序列目标题目通过注意力去历史状态里检索相关信息。如果你把目标题目也塞进LSTM相当于提前泄露了答案训练和测试的分布就对不上了。3.2 Recap模块软选择与硬选择两种实现LSTM的隐状态序列一长问题就来了和当前目标题目真正相关的历史习题可能散落在50步之前、200步之前中间夹了一堆无关练习。LSTM理论上能记住长距离信息但实际训练里无关步会稀释相关步的梯度信号。GIKT的Recap模块就是为这个设计的——在最终预测之前先从历史里挑出和目标题目最相关的那些习题状态。Recap有两种实现方式。软选择对全部历史状态做注意力加权得到一个加权求和的历史上下文硬选择只保留注意力权重最高的k个历史状态其余全部丢弃。def recap_soft(self, h_seq, q_target, mask): # h_seq: [batch, seq_len, hidden_dim] LSTM历史隐状态 # q_target: [batch, hidden_dim] 目标题目的GCN嵌入 # mask: [batch, seq_len] 1有效, 0padding attn torch.matmul(q_target.unsqueeze(1), h_seq.transpose(1, 2)) attn attn.squeeze(1) # [batch, seq_len] attn attn.masked_fill(mask 0, -1e9) # padding位置置负无穷 attn F.softmax(attn, dim-1) context torch.bmm(attn.unsqueeze(1), h_seq).squeeze(1) return context, attn软选择的逻辑很直观用目标题目向量和每个历史隐状态做点积点积越大代表越相关softmax后变成权重再对所有历史状态做加权平均。masked_fill这行保证padding位置不会分走注意力权重——如果不加短序列样本会把注意力分给一堆零向量历史上下文被污染。def recap_hard(self, h_seq, q_target, k10): # h_seq: [batch, seq_len, hidden_dim] attn torch.matmul(q_target.unsqueeze(1), h_seq.transpose(1, 2)).squeeze(1) topk_val, topk_idx torch.topk(attn, k, dim-1) # [batch, k] # 按索引收集对应的历史隐状态 idx topk_idx.unsqueeze(-1).expand(-1, -1, h_seq.size(-1)) selected torch.gather(h_seq, 1, idx) # [batch, k, hidden_dim] return selected, topk_val硬选择多一个超参数k也就是最多保留多少个相关历史习题。k我一般取8到12太小会丢掉有用信息太大就退化成软选择。论文实验里两种都有报告硬选择在长序列上的噪声抑制更明显尤其是学生交互次数差异很大的数据集软选择不用调k实现更省事。我复现时默认用硬选择k10在ASSISTments上效果稳定。3.3 注意力计算的两个工程细节Recap的注意力看起来就是普通点积注意力但实际数据上有两个坑。第一个是mask上面代码里已经处理了。第二个是数值稳定性点积结果的尺度会随hidden_dim增大而变大softmax梯度会变小。常见做法是除以sqrt(hidden_dim)也就是scaled dot-product attention。GIKT原论文没特别强调这一步但加上缩放后训练更稳hidden_dim128时尤其明显。# 加缩放的注意力分数 scale q_target.size(-1) ** 0.5 attn torch.matmul(q_target.unsqueeze(1), h_seq.transpose(1, 2)) / scale提示缩放因子用目标题目嵌入维度计算不要用历史序列长度后者会让分数尺度随序列长度漂移。缩放因子放在matmul之后、softmax之前。如果不想引入这个超参数也可以把q_target过一层LayerNorm再算注意力效果类似。我倾向于用缩放改动最小不破坏原有结构。另外提醒一句Recap模块的注意力和后面交互模块的注意力是两套作用完全不同别复用同一个参数——前者决定“选哪些历史习题”后者决定“信哪个交互结果”语义不一样共用参数会让两阶段互相干扰。4. 广义交互模块四路信息一致建模学生的掌握程度4.1 为什么先选再交互而不是直接聚合完就预测SKVMN和EERNNA的做法是把相关历史状态聚合成一个新状态拿这个新状态去做预测。GIKT的差别在于它不在聚合这一步就结束而是把“学生当前状态、选中的历史习题、目标题目、相关技能”这四个对象两两交互每个交互单独产生一个预测最后用注意力加权。为什么要绕这一圈我的理解是聚合操作会损失对应关系。加权平均得到一个“综合历史状态”后你只知道“学生过去表现大致怎样”但不知道“学生是否答对过和这道题很像的题”。GIKT的交互模块保留了这种对应当前状态和历史习题的交互体现的是“学生从那段经历中获得了什么”目标题目和相关技能的交互体现的是“这道题到底在考什么”。这两类信息性质不同混在一起聚合会互相干扰。简单说先选再交互等于把“回顾什么”和“怎么判断”拆成两步每步都更可控。多技能题目的处理也需要交互模块。一道题对应多个技能时不同技能的掌握度可能差很多把多个技能嵌入平均成一个向量等于让高掌握技能去补贴低掌握技能。GIKT让目标题目嵌入和每个相关技能嵌入分别交互技能越多交互分支越多最终预测由注意力决定哪些技能更重要而不是粗暴平均。4.2 交互运算与维度变化交互的核心运算是逐元素乘法不是拼接后过全连接。逐元素乘法一开始看着有点玄学但跑几次对比实验就知道它比“拼接后过MLP”在稀疏数据上稳得多。逐元素乘法的好处是两个向量相乘每个维度上都做了对齐语义上相当于判断“这个特征维度上两者是否同时激活”。拼接后过MLP也能做但参数量成倍增加在稀疏数据集上更容易过拟合。GIKT把交互后的向量再接一个小MLP输出预测分数各分支的维度变化如下def interaction_module(self, h_current, selected, q_target, s_target): # h_current: [batch, hidden_dim] 学生当前LSTM状态 # selected: [batch, k, hidden_dim] Recap选出的历史习题状态 # q_target: [batch, hidden_dim] 目标题目嵌入 # s_target: [batch, hidden_dim] 技能嵌入多技能先聚合 k selected.size(1) # 交互1: 当前状态 x 历史习题 inter1 h_current.unsqueeze(1) * selected # [batch, k, d] # 交互2: 当前状态 x 目标题目 inter2 h_current.unsqueeze(1) * q_target.unsqueeze(1).expand(-1, k, -1) # 交互3: 当前状态 x 相关技能 inter3 h_current.unsqueeze(1) * s_target.unsqueeze(1).expand(-1, k, -1) # 交互4: 历史习题 x 目标题目 inter4 selected * q_target.unsqueeze(1) # [batch, k, d] # 拼接后过MLP每个交互输出一个logit all_inter torch.cat([inter1, inter2, inter3, inter4], dim-1) logits self.inter_mlp(all_inter).squeeze(-1) # [batch, k] return logitsinter_mlp我用两层全连接中间接ReLU输出维度1。交互2和交互3用expand把当前状态和目标题目复制到k个维度上让四个交互分支的形状统一成[batch, k, d]拼接后是[batch, k, 4d]进MLP。这里d就是hidden_dim。多技能时s_target怎么得到我试过两种对多个技能嵌入取平均以及用目标题目嵌入对技能嵌入做注意力加权求和。后一种效果略好计算量也小推荐直接用# 多技能聚合注意力加权 # s_list: [batch, n_skills, hidden_dim] 当前题目涉及的技能嵌入 # q_target: [batch, hidden_dim] scores torch.matmul(q_target.unsqueeze(1), s_list.transpose(1, 2)).squeeze(1) s_weights F.softmax(scores, dim-1) # [batch, n_skills] s_target (s_weights.unsqueeze(1) s_list).squeeze(1) # [batch, hidden_dim]注意力聚合的好处是模型可以为每个技能学一个贡献权重而不是默认所有技能同等重要。这在题目对应技能数差异很大的数据集上尤其明显。4.3 最终预测的注意力加权每个交互都给出了一个预测但可信度不同。“历史习题x目标题目”这个交互只有在这道历史习题真的和目标题目很像时才有意义如果k个历史习题里有一半其实不太相关它们对应的交互预测应该被压低。GIKT的处理是再套一层注意力# 用独立参数计算交互权重 w_logits self.attn_layer(torch.tanh(all_inter)).squeeze(-1) # [batch, k] attn_w F.softmax(w_logits, dim-1) # 每个交互的预测概率 inter_prob torch.sigmoid(logits) # [batch, k] # 加权求和得到最终预测 final_pred (inter_prob * attn_w).sum(dim-1) # [batch]要点是计算权重的MLP和计算预测的MLP必须是两组独立参数。如果共用一组会出现“预测概率高的交互天然拿到高权重”的正反馈模型没有动力去区分“这个交互预测可靠但概率低”和“不可靠但概率高”。分开后注意力层学到的是哪些交互更值得信预测层学到的是这个交互给出什么预测两者解耦。训练配置上我复现GIKT用Adam初始学习率0.001weight decay 0.0001batch size 64。GCN部分的梯度有时偏小我会单独给GCN层参数设0.002的学习率避免传播层学得太慢。LSTM序列长BPTT的梯度范数容易爆我一般用clip_grad_norm_把梯度裁剪到5.0在backward之后、optimizer.step之前调用。验证集AUC连续5个epoch不涨就把学习率降到原来的0.1再跑5个epoch还没改善就早停。这个配置在ASSISTments 2009-2010数据集上一般40个epoch内能收敛到论文报告的AUC区间。5. GIKT复现避坑数据格式、维度对齐与指标计算的五个常见问题复现这类模型翻车点基本集中在数据预处理和张量形状上。下面五条是我实际跑的时候踩过的坑每一条都按现象、原因、解决三步写清楚。5.1 同一题目对应多个技能邻接矩阵漏边现象模型能跑通但验证集AUC明显低于论文报告值多个数据集上都是同一个偏差。原因数据预处理时只取了题目的第一个技能或者把多技能题目当成多个独立副本处理。前者让邻接矩阵少了边高阶信息传不完整后者会让同一道题在嵌入空间里被拆成多份训练时梯度互相打架。解决遍历q2s映射时把每个技能都连边邻接矩阵里一行可以有多个非零值。多技能题目在交互模块里再聚合技能嵌入聚合方式用注意力加权而不是简单平均这样每个技能对预测的贡献是模型学出来的。5.2 GCN层数加深后AUC反而下降现象把GCN从2层加到3层、4层验证集AUC不升反降掉了1到2个点。原因过度平滑。GCN层数越多每个节点的表示越趋向于全图平均值题目和技能的区分度被抹平。这个问题在图神经网络里很常见尤其在二部图这种结构上3层以上基本都会出现。解决保持2层聚合。如果觉得高阶信息不够先检查邻接矩阵的构建是否正确而不是直接加层。另外第一层GCN后的dropout对缓解过度平滑有帮助设0.5比设0.1的效果更稳。5.3 序列长、batch大导致显存溢出现象训练到一半报CUDA out of memory尤其是在序列长度超过500的数据集上。原因LSTM需要把所有时间步的反向传播路径都保存下来Recap的注意力矩阵是batch×seq_len序列一长显存占用快速上涨。解决训练时把学生序列截断到200到300步超出部分丢弃或另起一个片段。注意截断要尽量保留最近的历史因为知识追踪的场景里最近作答对当前状态的贡献最大。如果截断后batch还想调大配合梯度累积每4个batch更新一次参数效果等同大batch训练。5.4 节点索引顺序错乱导致嵌入张冠李戴现象loss不下降或者震荡训练曲线完全异常。排查半天发现题目嵌入和技能嵌入混在一起。原因GCN传播后的节点表示仍然按构建邻接矩阵时的编号顺序排列但复现时把切片索引写错比如直接用技能数量去切题目嵌入或者图构建和嵌入层用了两套编号。解决固定一套编号规则从头到尾只用这一套。我习惯把题目放前面、技能放后面用num_questions作为分界任何地方需要取嵌入都从同一个索引常量出发。写完之后加一个断言检查切片出来的嵌入维度是否和Embedding矩阵一致。5.5 AUC计算口径与论文不一致现象自己算的AUC和论文报告值差3到5个点但模型结构和参数都一样。原因数据集划分方式不同AUC是按学生维度算还是全局样本算以及是否过滤重复作答记录都会带来明显差异。ASSISTments这类数据集里同一道题被同一个学生反复做的情况很多直接全量算会高估模型表现。解决先确认你拿到的复现代码用的划分方式——知识追踪任务里通常按学生划分保证同一个学生不会同时出现在训练集和测试集。计算AUC时用sklearn的roc_auc_score先逐学生计算再取平均或者全局计算但在代码注释里写明口径。数据预处理阶段把重复作答记录只保留第一次这也是知识追踪任务里的常规操作。提示ASSISTments原始文件里包含多张表题号列和技能列要先做去重否则邻接矩阵里会出现重复边导致归一化后的权重偏低。6. 进阶把GIKT裁剪成适合自己平台的轻量版6.1 裁剪GCN与交互模块完整GIKT在学术数据集上效果好但如果接到生产环境计算成本和部署复杂度都要重新评估。我自己做轻量化时只动三个地方。第一GCN保持一层。一层GCN只传播直接相连的技能信息损失了共享技能题目之间的二阶关联但AUC只掉0.3到0.5个点换来的是一张稀疏矩阵乘法变成了一次查表加聚合线上推理快很多。第二Recap的k从10减到5。实际平台里学生最近作答的题目往往是最相关的top-5足够覆盖大多数预测场景。第三交互模块从4路减到2路只保留“当前状态x历史习题”和“历史习题x目标题目”。当前状态x目标题目的信息与前者有重叠减掉之后对AUC影响很小但显存占用下降约四分之一。6.2 验证方法与上线前检查裁剪后不能只看AUC我会补两个维度一个是RMSE看预测概率的绝对误差是否变大另一个是冷启动题目的单独AUC把训练集中出现次数少于5次的题目拎出来单独评估。GCN带来的提升主要就体现在这些稀疏题目上如果裁剪后冷启动AUC掉得比总体AUC快说明图传播信息被砍多了需要回调GCN层数。上线前还有一道检查把训练集和测试集的题目ID分布画出来确认测试集里没有出现训练集完全没有的新题目。知识追踪的离线验证假设是题目集合固定如果平台每天都有新题入库GIKT这类基于静态图的方法需要定期重构图或者降级为只依赖技能ID的备用模型。复现GIKT之前我一直觉得图卷积就是给节点嵌入加一层邻居平均没什么特别的。真正把邻接矩阵、索引切分、交互模块拼到一起之后才发现最耗时间的不是模型本身而是弄清楚每个张量在哪个维度上做什么。从那以后我每次复现论文模型第一件事就是把数据流和维度变化画在一张纸上跑通后再谈调参。这份资源里的实现省掉了这一步的摸索希望帮到你。本文还有配套的精品资源点击获取
返回列表