
简介这是一份TransE模型的Python实现配套FB15k知识图谱数据集面向知识图谱表示学习入门者、算法工程师及需要完成链接预测、知识图谱补全等任务的研究人员。资源共21个文件包含14个txt数据文件例如训练集、验证集与测试集的划分、4个py源码文件、2个md说明文档以及1个png示意图压缩包整体仅5.85MB。已有401人学习下载。源码完整实现了TransE的核心训练流程随机初始化实体与关系向量通过负采样构造错误三元组采用L1或L2距离作为损失函数利用梯度下降进行优化使正确三元组满足头实体向量加关系向量约等于尾实体向量这一约束。基于FB15k真实数据可直接运行训练、验证与测试并输出MRR、HITSk等常用评估指标。配套的README与示意图能辅助理解模型结构、数据格式和训练参数适合作为课程设计、论文复现或实际知识图谱项目的基线也便于后续调整向量维度、学习率、批大小等超参数开展对比实验。1. 用 Python 把 TransE 跑在 FB15k 上先绕开这 3 个坑TransE 是知识图谱嵌入里最朴素也最常被当作基线模型的算法FB15k 则是 Freebase 的一个 15k 子集train.txt里存的就是它的标准三元组训练数据。标题里的train.txt分_TransE.zip拆开看其实已经点明了完整工作流把 FB15k 的三元组数据按实体、关系、训练/验证/测试切分好再用 Python 实现 TransE 在train.txt上完成训练和评估。一个反直觉的结论是TransE 在 FB15k 上收敛很快但指标涨得很慢。多数人第一次跑出来的 Hits10 和公开基线差一截问题往往不在模型本身而在数据划分方式、负采样策略和归一化时机这三个细节上。这篇文章面向刚接触知识图谱嵌入的工程师也适合已经跑通代码但调不动指标的人。你会看到理论怎么落到数据上以及哪些参数值得调、哪些坑必须绕开。2. TransE 的训练目标与 FB15k 数据格式先搞清楚 train.txt 里每一行是什么2.1 TransE 的得分函数与 margin lossh r ≈ t 怎么变成可优化目标TransE 的核心假设只有一句话在一个嵌入空间里头实体向量加上关系向量应该约等于尾实体向量也就是h r ≈ t。这个假设简单到近乎粗暴但它直接把知识图谱里的符号三元组变成了向量空间里的平移关系。得分函数最常见的写法是L1 距离||h r - t||_1L2 距离||h r - t||_2得分越低说明这个三元组越合理。训练目标不是让正样本得分变成 0而是让正样本得分比负样本低出一个 margin。这就是 margin-based ranking loss形式是loss max(0, margin d(h r, t) - d(h r, t))其中(h, r, t)是train.txt里的正样本三元组(h, r, t)是把头实体或尾实体随机替换掉后生成的负样本。这个 loss 的意思是正样本的距离要比负样本的距离小至少margin。如果差距已经够大loss 就是 0梯度也归零所以 TransE 的训练过程天然是稀疏更新的。选择 margin 时要留意一个权衡margin 太小正负样本的区分度不够嵌入向量会挤成一团margin 太大模型会过度关注那些难分的负样本训练不稳定。FB15k 上常见的取值在 0.5 到 2.0 之间但这和 embedding 维度、距离函数、学习率都耦合在一起后面第 4 章会具体说怎么组合调参。2.2 FB15k 的 train.txt / valid.txt / test.txt 结构与 ID 映射FB15k 标准发布包里通常有三个文件train.txt、valid.txt、test.txt。每个文件一行一个三元组字段之间用 tab 分隔格式是head\trelation\ttail。注意这里不是空格也不是逗号。拿到train.txt分_TransE.zip之后第一步先解压看结构unzip train.txt分_TransE.zip wc -l train.txt valid.txt test.txt head -5 train.txtwc -l统计行数能快速确认三个文件的规模比例。FB15k 常见的划分是训练集约 47 万条、验证集约 5 万条、测试集约 5.9 万条但不同版本的预处理可能略有出入以你实际解压出来的行数为准。head -5看前五行长什么样确认是head\trelation\ttail而不是别的顺序。实体和关系在原始文件里是字符串名字比如/m/02_m8jq这种 Freebase 风格的 ID。不能直接把字符串喂给神经网络需要先做 ID 映射。常见做法是遍历所有文件把出现过的实体和关系各编一个从 0 开始的整数 IDentities, relations set(), set() for fname in [train.txt, valid.txt, test.txt]: with open(fname, r, encodingutf-8) as f: for line in f: h, r, t line.strip().split(\t) entities.add(h) entities.add(t) relations.add(r) entity2id {e: i for i, e in enumerate(entities)} relation2id {r: i for i, r in enumerate(relations)}这段代码把所有三元组的头实体、尾实体收进entities集合关系收进relations集合然后分别建映射。enumerate从 0 开始编号ID 连续且紧凑。这样做的原因是 embedding 层本质是一个查找表ID 必须连续否则nn.Embedding会报索引越界或者浪费大量内存。2.3 数据划分train.txt 分出来之后怎么生成训练和评估要用的张量第 2.2 节建好了 ID 映射接下来要把train.txt转成 PyTorch 能直接读取的 tensor。这里有一个很多人会忽略的坑验证集和测试集的负样本必须排除掉训练集里出现过的三元组。FB15k 的标准做法是 filtered 评估也就是在计算排名时把训练集里已有的其他正确三元组从候选里去掉否则 Hits10 会被严重高估。这个逻辑在数据准备阶段就要留下训练集三元组的集合后面评估要用。import torch def load_triples(fname, entity2id, relation2id): triples [] with open(fname, r, encodingutf-8) as f: for line in f: h, r, t line.strip().split(\t) triples.append((entity2id[h], relation2id[r], entity2id[t])) return torch.tensor(triples, dtypetorch.long)load_triples输出的 tensor 形状是(num_triples, 3)每一行是(head_id, relation_id, tail_id)。训练时按 batch 切分这个 tensor评估时再单独处理。到这一步train.txt已经变成了模型可以消费的数字形式后面的实现都建立在这套 ID 体系上。3. Python 实现 TransE 的最小可训练版本用 PyTorch 从零写训练闭环3.1 数据加载器把 train.txt 按 batch 喂给模型并生成负样本数据准备阶段拿到的是完整训练三元组 tensor训练时要随机打乱并按 batch 切分。PyTorch 的DataLoader配合TensorDataset可以做这件事但负样本生成需要自己写。常见的负采样做法是对 batch 里的每个正样本随机替换头实体或尾实体替换的实体从全量实体集合里均匀采样。下面这个NegativeSampler是我常用的写法class NegativeSampler: def __init__(self, num_entities): self.num_entities num_entities def corrupt(self, heads, tails, bern_prob0.5): heads heads.clone() tails tails.clone() mask torch.rand(heads.size(0)) bern_prob heads[mask] torch.randint(0, self.num_entities, (mask.sum(),)) tails[~mask] torch.randint(0, self.num_entities, ((~mask).sum(),)) return heads, tailscorrupt方法接收正样本的heads和tails向量按bern_prob的概率决定是破坏头还是破坏尾。mask是布尔张量True的位置替换头实体False的位置替换尾实体。torch.randint在[0, num_entities)区间均匀采样这个范围必须覆盖全部实体 ID否则会出现采样不到某些实体的情况。3.2 模型定义与参数初始化nn.Embedding 和 L2 归一化的边界TransE 的模型结构是一个双 embedding 表一个给实体一个给关系。实体 embedding 的形状是(num_entities, dim)关系 embedding 的形状是(num_relations, dim)。PyTorch 实现如下import torch.nn as nn class TransE(nn.Module): def __init__(self, num_entities, num_relations, dim): super().__init__() self.entities nn.Embedding(num_entities, dim) self.relations nn.Embedding(num_relations, dim) nn.init.uniform_(self.entities.weight, -1.0, 1.0) nn.init.uniform_(self.relations.weight, -1.0, 1.0) def forward(self, heads, relations, tails, normalizeTrue): h self.entities(heads) r self.relations(relations) t self.entities(tails) if normalize: h torch.nn.functional.normalize(h, p2, dim1) t torch.nn.functional.normalize(t, p2, dim1) return h, r, t初始化用均匀分布[-1, 1]这是 TransE 原论文里的做法。关键点在normalizeTrue头实体和尾实体的 embedding 每步都要做 L2 归一化但关系向量不归一化。这样做的原因是实体的模长差异如果不被约束模型会倾向于让某些实体的模长变得特别大从而让距离度量失效。关系不归一化是因为关系的模长本身携带了语义信息归一化会把这种信息抹掉。这个边界很微妙后面第 5 章还会展开。3.3 训练循环与损失计算SGD margin loss 的完整代码模型定义好了负样本也有生成器了接下来是训练主循环。损失函数用第 2.1 节的 margin ranking loss优化器一般用SGD而不是 Adam。Adam 对 TransE 效果不稳定因为它的自适应学习率会放大稀疏梯度的影响而 TransE 的梯度本来就是稀疏的def train(model, train_triples, sampler, optimizer, margin, batch_size1024, epochs1000): model.train() num_train train_triples.size(0) for epoch in range(epochs): perm torch.randperm(num_train) train_triples train_triples[perm] total_loss 0.0 for batch_start in range(0, num_train, batch_size): batch train_triples[batch_start:batch_start batch_size] heads, relations, tails batch[:, 0], batch[:, 1], batch[:, 2] neg_heads, neg_tails sampler.corrupt(heads, tails) h, r, t model(heads, relations, tails) h_neg, r_neg, t_neg model(neg_heads, relations, neg_tails) pos_dist (h r - t).norm(p2, dim1) neg_dist (h_neg r_neg - t_neg).norm(p2, dim1) loss torch.clamp(margin pos_dist - neg_dist, min0).mean() optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() if epoch % 100 0: print(fepoch {epoch}, loss {total_loss:.4f})每个参数都要说清楚含义。perm是torch.randperm生成的随机排列索引用它重排训练三元组保证每个 epoch 看到的 batch 组成不同。batch_size1024是兼顾内存和梯度稳定性的常见取值。pos_dist计算正样本的 L2 距离neg_dist计算负样本的 L2 距离torch.clamp(..., min0)实现max(0, margin pos - neg)的合页损失。.mean()取 batch 内平均让 loss 数值和 batch_size 无关。optimizer.zero_grad()是必须的否则梯度会跨 batch 累积。优化器的初始学习率一般取0.01到0.1之间比常规神经网络的默认值高一个量级。原因是 TransE 的目标函数相对简单损失曲面没有那么多局部最优大学习率反而能更快找到一个合理的嵌入空间optimizer torch.optim.SGD(model.parameters(), lr0.05)4. 在 FB15k 上调参embedding 维数、margin、batch size 怎么影响指标4.1 三个最影响结果的参数dim、margin、学习率FB15k 上的表现很大程度上被三个超参数决定embedding 维度dim、损失函数里的margin、优化器的学习率。它们之间的耦合关系比大多数人以为的更强。先看一个实际参考区间参数推荐范围对结果的影响备注dim50 ~ 200维度太低表达力不足维度太高容易过拟合且训练变慢FB15k 上 100 是常用起点margin0.5 ~ 2.0太小区分度差太大训练不稳定与距离函数强相关L2 用大 marginlearning rate0.01 ~ 0.1太大收敛震荡太小涨点极慢配合 SGD 使用Adam 慎用batch size512 ~ 4096影响梯度方差和内存占用1024 是保守起点epochs500 ~ 2000过拟合时 valid 指标先升后降必须配 early stoppingdim的选择要看实体和关系的数量。FB15k 实体约 1.5 万、关系约 1300 个100 维在表达力和计算开销之间比较平衡。如果你把 dim 提到 200训练时间接近翻倍但 Hits10 往往只提升一两个点甚至可能因为过拟合而下降。margin的设置和距离函数强耦合用 L2 距离时距离值的量级比 L1 大margin 通常取 1.0 以上用 L1 距离时margin 取 0.5 到 1.0 就够了。4.2 用 valid.txt 做 early stopping并算 MRR 和 Hits10训练 TransE 不能只盯着训练 loss。很多时候训练 loss 一路下降但验证集指标早就开始掉头了这就是典型的过拟合。要避免这种状况需要在每个 epoch 结束后用valid.txt评估一次指标并且只保存验证指标最好的那一次模型权重。评估要同时算 MRR 和 Hits10这两个指标在 FB15k 上是标准配置def evaluate(model, valid_triples, all_triples, num_entities): model.eval() ranks [] with torch.no_grad(): for h, r, t in valid_triples: h, r, t h.item(), r.item(), t.item() h_emb model.entities(torch.tensor([h])) r_emb model.relations(torch.tensor([r])) t_emb model.entities.weight scores (h_emb r_emb - t_emb).norm(p2, dim1) # filtered: 排除训练集中存在的其他正确三元组 for e in range(num_entities): if e t: continue if (h, r, e) in all_triples: scores[e] float(inf) rank (scores scores[t]).sum().item() 1 ranks.append(rank) mrr (1.0 / torch.tensor(ranks)).mean().item() hits10 sum(1 for r in ranks if r 10) / len(ranks) return mrr, hits10这里我在评估时把尾实体全部替换为候选实体计算每个候选的得分然后找到正确尾实体得分排第几。filtered 的部分很关键if (h, r, e) in all_triples时把这个候选的分数设成inf意思是跳过那些在训练集里真实存在但不是当前尾实体的三元组。如果不做这一步MRR 会虚高。all_triples要在开始训练之前就把train.txt的全部三元组转成集合传进来。4.3 FB15k 上的常见坑训练 loss 下降但验证指标不涨怎么排查症状一训练 loss 下降但 MRR 和 Hits10 完全不动。先看负采样是不是有问题。如果corrupt时替换了头实体但新头实体恰好和原头实体是同一个实体这个负样本就是假的。虽然概率只有1/num_entities但 FB15k 实体数不多假负样本会把模型往错误方向推。症状二验证集 Hits10 和训练集 Hits10 差距很大训练集上接近 0.9、验证集只有 0.3这是过拟合。此时优先调小dim或提前停止训练而不是加数据。症状三训练 loss 一开始就震荡很可能是学习率太大把学习率从 0.1 降到 0.02 再试。还有一个在 FB15k 上特别容易踩的坑实体 embedding 归一化之后embedding 向量全部落在单位球面上距离的数值范围被压得很小。这时如果 margin 设置得比最大可能距离还大loss 会恒大于零模型会一直尝试拉开距离但永远做不到。排查方法是打印几轮pos_dist和neg_dist的均值如果正样本距离已经在 0.5 以下但 margin 是 2.0说明 margin 设大了。5. 在 train.txt 上把数据价值榨干bern 负采样、归一化时机与两阶段训练5.1 为什么均匀负采样在 FB15k 上不够改成 bern 采样第 3.1 节的均匀负采样在数学上是无偏的但对 FB15k 这种真实知识图谱来说不够高效。原因在于不同关系的头实体和尾实体分布差异极大。举一个具体例子/location/location/contains这种关系头实体通常是国家或城市尾实体通常是景点或区域实体类型差异明显。如果用均匀采样替换头实体很可能采到一个本质上不可能作为头实体的实体这种负样本太简单模型学不到东西。Bernoulli 负采样策略会统计每个关系下头实体被替换和尾实体被替换的频率然后按这个频率分配替换概率。具体做法是训练前统计每个关系r下正样本三元组(h, r, t)中每个头实体h平均对应多少个尾实体以及每个尾实体t平均对应多少个头实体然后算出两个替换概率。实现上要额外维护两个计数表def count_bern_probs(train_triples, num_relations): head_count torch.zeros(num_relations) tail_count torch.zeros(num_relations) rel_count torch.zeros(num_relations) for h, r, t in train_triples: head_count[r] 1 tail_count[r] 1 rel_count[r] 1 head_prob head_count / (head_count tail_count) return head_prob把head_prob传到负采样器里替换头实体的概率就不再是固定的 0.5而是由数据分布决定的动态值。FB15k 上用 bern 采样通常能让 Hits10 提高 2 到 5 个点提升幅度甚至比调 embedding 维度还明显。5.2 L2 归一化的时机训练完成后再归一化不会提升指标实体 embedding 的 L2 归一化不止是一个「要不要做」的问题更关键的是「什么时候做」。常见做法是在 forward 里对实体的 embedding 做归一化也就是第 3.2 节代码里的normalizeTrue。如果改成训练时不归一化、训练完再统一归一化指标不会变好。原因在于距离的数值分布已经被训练过程塑形了事后归一化只是把所有向量等比例缩放到单位球面相对距离排序一点不变。另外需要明确只有实体需要归一化关系永远不归一化。验证方法是训练完成后打印实体 embedding 和关系 embedding 的模长分布你会看到实体模长全为 1而关系模长分布很广。如果关系也被归一化了得分函数的表达能力会被限制在一个球面上无法表达模长差异带来的语义丰富性。5.3 一个实用技巧先粗训再精调用早停的模型做冷启动FB15k 上我一般会跑两阶段训练。第一阶段用较粗的参数快速找到大致合理的区域dim50, margin1.0, lr0.1, batch_size2048训练到验证集 Hits10 不再上升就停通常 200 到 400 个 epoch 就够。这个阶段不求指标高只求找到一个不会太差的初始点。第二阶段做精调把 embedding 维度提高到目标值初始化用第一阶段学到的向量插值到高维空间学习率降到0.01margin 改成1.5继续训练。两阶段训练的好处是第一阶段能快速鉴别数据预处理和代码有没有 bug。如果阶段一在几十个 epoch 内 Hits10 一直小于 0.1别急着调参先回头看负采样有没有问题看 ID 映射有没有错位。等模型能稳定跑到 0.2 以上再进入第二阶段精调。这时候每个 epoch 的训练速度已经快了很多因为权重的初始化已经处于一个较优区域梯度更新的步长可以更小收敛也更平稳。最后保存验证集 MRR 最高的 checkpoint用同一份权重在test.txt上跑一次最终指标作为对外报告的结果。检查test.txt里的三元组是否有和训练集完全重复的行如果有评估前直接去掉否则 Hits10 会被异常拉高。本文还有配套的精品资源点击获取