TransE是2013年提出的知识图谱嵌入模型,因其结构简单、训练高效,至今仍是许多图嵌入任务的基线。它假设一个三元组(头实体,关系,尾实体)满足h + r ≈ t的平移关系,即把关系看作从头实体指向尾实体的向量平移。大多数教程用Python实现,其实用Node.js同样可以完整跑通整个流程,而且不依赖任何深度学习框架,反而更能看清算法本质。本文将详细讲解原理和实现细节。

TransE的核心原理与损失函数
TransE的目标是为每个实体和每个关系学习一个d维向量。对于一个正确三元组(h, r, t),希望h + r的结果尽量接近t;对于构造出来的错误三元组,则希望距离越远越好。模型使用L1或L2范数来度量距离,通常L2效果更稳定,实现也更简单。
训练采用margin-based ranking loss,即带间隔的排序损失。公式为L = Σ max(0, γ + d(h+r, t) - d(h'+r, t')),其中γ是间隔超参数,h'和t'是负样本。直观理解:正确三元组的距离加上γ之后仍然要小于错误三元组的距离,否则产生损失。这个设计让模型不追求绝对距离小,而是追求相对排序正确,属于典型的对比学习思想。
此外,TransE要求实体向量归一化,即每个训练步之后把实体向量缩放为单位长度,防止向量通过无限变大来作弊降低距离。这一点在实现中容易被忽略,却是复现效果的关键,很多自制实现效果差就是漏掉了这一步。
在Node.js中搭建数据结构与负采样
首先定义向量运算工具函数。由于嵌入维度通常在50到100之间,用普通数组配合手写循环即可,性能完全够用,不必引入额外的矩阵运算库。
// 初始化一个随机向量,范围[-bound, bound]
function randomVector(dim, bound) {
const v = new Array(dim);
for (let i = 0; i < dim; i++) {
v[i] = (Math.random() * 2 - 1) * bound;
}
return v;
}
// L2距离的平方
function distSq(a, b) {
let s = 0;
for (let i = 0; i < a.length; i++) {
const d = a[i] - b[i];
s += d * d;
}
return s;
}
// 向量归一化为单位长度
function normalize(v) {
let s = 0;
for (const x of v) s += x * x;
const n = Math.sqrt(s) || 1;
for (let i = 0; i < v.length; i++) v[i] /= n;
}
接着构建实体表和关系表。用Map把字符串形式的实体名映射到索引,向量和索引分别存储在数组中,访问速度是O(1)。负采样策略是:随机替换头实体或尾实体,如果替换后恰好是真实存在的三元组,则重新采样。为了快速判断三元组是否存在,把正确三元组序列化为字符串存入Set,查询复杂度同样为O(1)。
const tripleSet = new Set();
triples.forEach(([h, r, t]) => tripleSet.add(h + '|' + r + '|' + t));
// 对一个正样本生成一个负样本
function corrupt(h, r, t, entities) {
while (true) {
const e = entities[Math.floor(Math.random() * entities.length)];
if (Math.random() < 0.5) {
if (!tripleSet.has(e + '|' + r + '|' + t)) return [e, r, t];
} else {
if (!tripleSet.has(h + '|' + r + '|' + e)) return [h, r, t, e][0] === h ? [h, r, e] : [h, r, e];
}
}
}
编写训练循环与梯度更新
训练时对每个三元组计算正负样本的距离差,若损失大于零则进行梯度更新。以L2距离为例,令D = Σ(h_i + r_i - t_i)^2,则对h_i的梯度是2(h_i + r_i - t_i),对r_i和t_i同理。更新方向是让正样本距离变小、负样本距离变大,正负样本的梯度符号恰好相反。学习率建议设在0.01附近,margin取1,迭代轮数视数据规模而定。
function add(a, b) {
const c = new Array(a.length);
for (let i = 0; i < a.length; i++) c[i] = a[i] + b[i];
return c;
}
function updateTowards(h, r, t, sign) {
for (let i = 0; i < h.length; i++) {
const g = 2 * (h[i] + r[i] - t[i]) * lr * sign;
h[i] -= g;
r[i] -= g;
t[i] += g;
}
}
function trainStep(hIdx, rIdx, tIdx, nhIdx, ntIdx) {
const h = entVec[hIdx], r = relVec[rIdx], t = entVec[tIdx];
const nh = entVec[nhIdx], nt = entVec[ntIdx];
const dp = distSq(add(h, r), t);
const dn = distSq(add(nh, r), nt);
if (dp + margin - dn > 0) {
updateTowards(h, r, t, -1); // 正样本:缩小距离
updateTowards(nh, r, nt, 1); // 负样本:拉大距离
normalize(h); normalize(t); normalize(nh); normalize(nt);
return margin + dp - dn;
}
return 0;
}
上面的updateTowards函数根据距离公式逐维更新三个向量:头实体和关系向量朝缩小误差的方向移动,尾实体朝相反方向移动。注意关系向量不做归一化,只有实体向量需要。每轮训练结束后打乱数据顺序,能明显提升收敛稳定性,避免模型记住数据的固定排列。
评估效果与工程优化建议
训练完成后可以做链接预测:给定(h, r)预测t,把所有实体作为候选,计算d(h + r, e),按距离升序排序,查看真实尾实体的排名。常用指标是Hits@10和Mean Rank。用JavaScript的sort配合map即可完成,代码非常直观。
function predictTail(hIdx, rIdx) {
const target = add(entVec[hIdx], relVec[rIdx]);
return allEntityIdx
.map(e => ({ e, dist: distSq(target, entVec[e]) }))
.sort((a, b) => a.dist - b.dist)
.map(x => x.e);
}
// 计算Hits@10
function hitsAt10(testTriples) {
let hit = 0;
for (const [h, r, t] of testTriples) {
const ranking = predictTail(entityIndex[h], relIndex[r]);
if (ranking.indexOf(entityIndex[t]) < 10) hit++;
}
return hit / testTriples.length;
}
性能方面有几个可行优化:一是把向量从普通数组换成Float32Array存储,内存更紧凑且缓存友好,运算速度也有提升;二是用worker_threads模块并行采样负样本,主线程专注梯度更新;三是实体归一化可以每k步批量做一次而非每步都做,在大规模数据上能节省不少时间。经过这些调整,在数万三元组规模的数据集上,Node.js实现的训练速度完全能满足学习和实验需求。
整体来看,TransE的数学结构决定了它非常适合手工实现:没有复杂的网络层,梯度和模型本身几乎是一体的。用Node.js走一遍完整流程,不仅得到一个可用的嵌入模型,更能加深对对比学习损失、负采样和向量归一化这些通用技术的理解,这些概念在后续学习TransH、TransR等改进模型时同样适用。