原始笔记

二: Training and Inference: Input

一: 数据预处理

A. 第一步:分组(Grouping)

首先,根据 UMLS 的原始数据(如 MRCONSO.RRF),将所有属于同一个 CUI 的术语归为一堆 。

  • 例子
    • CUI: C001 -> {名称 A, 名称 B, 名称 C}
    • CUI: C002 -> {名称 D, 名称 E}

B. 第二步:枚举组合(Enumeration)

对于每一个 CUI,在它的名称集合内部进行两两组合(笛卡尔积),生成格式为 (name1, name2, CUI) 的元组 。

  • 接上面的例子
    • C001 产生的条目:(A, B, C001), (A, C, C001), (B, C, C001)
    • C002 产生的条目:(D, E, C002)

C. 第三步:平衡与裁剪(Trimming)

有些医学概念可能有几百个同义词,会导致生成的对数(Pairs)发生爆炸。为了防止某些高频概念主导训练,作者规定:

  • 裁剪规则:如果一个 CUI 产生的正样本对超过 50 对,则随机剔除,只保留 50 对

D. 第四步:最终数据集

经过这种枚举和裁剪,作者最终得到了一个包含 11,792,953 条配对数据的大表 。

二: Training and Inference: Input

1. 预训练阶段:每一个“词”都是锚点

在预训练过程中,SapBERT 的输入是纯粹的实体名称(Surface Forms),完全脱离了句子上下文 。

  • 数据来源:输入数据是从 UMLS 中提取的成对名称,如 (名称1, 名称2, CUI)
  • 计算逻辑:在一个 Mini-batch 中,系统会将这 512 个名称全部输入 BERT 。
  • 锚点行为:批次内的每一个名称都会轮流担任一次“锚点(Anchor xax_a)” 。
  • Loss 计算:对于每一个当前的锚点,模型都会去寻找它的正样本(同 CUI)和负样本(异 CUI),通过在线挖掘筛选出硬样本对来计算 MS Loss 。
  • 结论:在这个阶段,没有“特定的 mention”,只有“术语名称”作为计算单元 。

2. 下游任务阶段:什么是真正的 Mention?

当你把训练好的 SapBERT 拿去跑测试(如 NCBI 或 COMETA 数据集)或者进行微调时,“Mention”的概念才真正出现 。

  • Mention 的定义:指在实际文本(如 Reddit 帖子或论文摘要)中出现的、需要被识别和链接的字符串 。
  • 任务目标:将这个“Mention”映射到知识库中的标准概念(CUI) 。
  • 操作方式
    • 微调时:将(实体提及 Mention, 对应的标准词)作为训练对 。
    • 推理时:将 Mention 转化为向量,然后在预先算好的医学术语向量库中做最近邻搜索

3. 关键区别总结

阶段 输入是什么? 有没有上下文(Context)? 谁是 Anchor?
预训练 (UMLS) 仅术语名称(如 "HCQ") 没有,只考虑实体本身 批次内的每一个名称
微调/测试 (MEL) 文本中的提及(如 "患者感到头痛") 取决于设置,但 SapBERT 倾向于只看实体本身 文本中被标记的 Mention

三: Training process

第一步:计算全成对相似度矩阵 SS

首先,将 Mini-batch 中 bb 个实体的特征向量(由 BERT 输出的 [CLS] 向量 )表示为矩阵 FRb×dF \in \mathbb{R}^{b \times d}

  1. 计算矩阵 SS:通过 S=FFTS = F \cdot F^T(假设向量已归一化)得到一个 b×bb \times b 的相似度矩阵 。
  2. 矩阵含义SijS_{ij} 代表第 ii 个实体和第 jj 个实体之间的余弦相似度

第二步:构建标签掩码 (Label Masks)

为了区分哪些是“自己人”(同义词),我们需要利用标签(CUI):

  1. 正样本掩码 (MposM_{pos}):如果第 ii 个实体和第 jj 个实体的 CUI 相同,则 Mpos[i,j]=1M_{pos}[i, j] = 1,否则为 0 。
  2. 负样本掩码 (MnegM_{neg}):如果 CUI 不同,则 Mneg[i,j]=1M_{neg}[i, j] = 1,否则为 0 。
  3. 注意:通常会排除对角线(即自己和自己)。

第三步:利用 SS 进行在线硬采样

虽然论文公式使用的是欧几里得距离 2||\cdot||_2 ,但在单位向量下,距离和余弦相似度是可以互换的:uv22=22cos(u,v)||u-v||_2^2 = 2 - 2 \cdot \cos(u, v)实际的筛选逻辑如下

  1. 转换条件:将距离公式转化为相似度公式。违反距离条件等同于: Sap<San+marginS_{ap} < S_{an} + \text{margin} (这里的 margin\text{margin} 对应论文中的 λ\lambda,通常取小负数,如 -0.2 )。
  2. 批量寻找硬三元组
    • 对于每一个锚点 ii
    • S[i,:]S[i, :] 中通过 MposM_{pos} 找出所有的正样本相似度 {Sip}\{S_{ip}\}
    • S[i,:]S[i, :] 中通过 MnegM_{neg} 找出所有的负样本相似度 {Sin}\{S_{in}\}
    • 判定硬样本:如果某对 (p,n)(p, n) 满足 SipSin<marginS_{ip} - S_{in} < \text{margin},则记录索引 pp 为硬正样本,索引 nn 为硬负样本 。

第四步:计算 MS Loss (公式 2 的具体实现)

现在你拥有了每个锚点 ii 对应的硬正样本集 Pi\mathcal{P}_i 和硬负样本集 Ni\mathcal{N}_i 。接下来带入公式 (2) 进行求和 :

  1. 处理负样本项(Push)
    • 计算 nNieα(Sinϵ)\sum_{n \in \mathcal{N}_i} e^{\alpha(S_{in} - \epsilon)}
    • 这会给那些相似度极高(错误地很像锚点)的负样本极大的权重 。
  2. 处理正样本项(Pull)
    • 计算 pPieβ(Sipϵ)\sum_{p \in \mathcal{P}_i} e^{-\beta(S_{ip} - \epsilon)}
    • 这会给那些相似度极低(离锚点太远)的正样本极大的权重 。
  3. 对数求和:对上述两项取 log(1+)\log(1 + \dots) 并乘以温度系数 1/α1/\alpha1/β1/\beta
  4. 最终平均:对批次内所有锚点产生的损失求平均,得到最终 L\mathcal{L} 。 MS Loss 的数学表达式如下 : L=1R0i=0R0(1αlog(1+nNieα(Sinϵ))负样本项 (Push)+1βlog(1+pPieβ(Sipϵ))正样本项 (Pull))\mathcal{L}=\frac{1}{|R_{0}|}\sum_{i=0}^{|R_{0}|} \left( \underbrace{\frac{1}{\alpha}\log(1+\sum_{n\in\mathcal{N}_{i}}e^{\alpha(S_{in}-\epsilon)})}_{\text{负样本项 (Push)}} + \underbrace{\frac{1}{\beta}\log(1+\sum_{p\in\mathcal{P}_{i}}e^{-\beta(S_{ip}-\epsilon)})}_{\text{正样本项 (Pull)}} \right)

Switch to English