SQLova架构详解(二):Seq2SQL_v1模型的六大预测子模块与列注意力机制逐层拆解
【免费下载链接】sqlova项目地址: https://gitcode.com/gh_mirrors/sq/sqlova
SQLova 是一个开源的神经语义解析模型,核心功能是把自然语言问题翻译成 SQL 查询(NL2SQL),在 WikiSQL 基准上取得了 83.6% 的逻辑形式准确率。本文逐层拆解 SQLova 核心 Seq2SQL_v1 模型的六大预测子模块(SCP、SAP、WNP、WCP、WOP、WVP)与贯穿其中的列注意力机制,并说明执行引导解码如何进一步提升准确率,帮你快速读懂这套"序列到 SQL"架构的设计思路。
🧭 整体架构回顾:从自然语言问题到 SQL 查询
SQLova 的推理链路可以概括为三步:
- 表感知词嵌入:用 BERT 把"问题 + 表头"拼成一条序列统一编码,得到上下文相关的词向量;
- 序列到 SQL(Seq2SQL):由 wikisql_models.py 中的
Seq2SQL_v1把 SQL 拆成 6 个可独立预测的部件,由 6 个子模块依次打分; - 执行引导解码(SQLova-EG):把候选 SQL 放进真实数据库试执行,只保留"跑得通"的候选。
下图是 WikiSQL 任务中一个典型的"问题 + 数据表"输入示例(来自人评界面):
面对上图这类问题("哪名球员的背号是 31?"),模型并不直接生成一串 SQL 文本,而是把查询拆成结构化部件:SELECT 列、聚合操作、WHERE 列 操作符 值× N——这正是六大子模块的分工。
📋 六大预测子模块一览
SQLova 借鉴了 SQLNet 的"序列到集合"(sequence-to-set)结构,用 6 个轻量模块分别预测 SQL 的一个部件。Seq2SQL_v1的初始化代码把分工写得很直白:
| 子模块 | 类名 | 预测目标 | 输出形态 | 对应 SQL 部件 |
|---|---|---|---|---|
| 1️⃣ 选择列 | SCP | SELECT 哪一列 | 列上的分数向量 | SELECT col |
| 2️⃣ 聚合操作 | SAP | MAX/MIN/COUNT/SUM/AVG/无 | 6 类分类 | agg(col) |
| 3️⃣ 条件数量 | WNP | 0–4 个 WHERE 条件 | 5 类分类 | 条件个数 |
| 4️⃣ 条件列 | WCP | 每列是否被条件引用 | 逐列 sigmoid 分数 | WHERE col |
| 5️⃣ 条件操作符 | WOP | =、>、<、其他 | 每条条件 4 类 | WHERE col op |
| 6️⃣ 条件值 | WVP_se | 问题中的起止 token 区间 | 每个 token 的 (start, end) 分数 | WHERE col op value |
前向传播中六步是级联的(见 forward):SCP 先选出列,SAP 以该列为条件预测聚合;WNP/WCP 决定条件骨架,WOP 基于已选条件列预测操作符,WVP_se 再基于列+操作符抽取值区间。训练时各部件的损失函数在 Loss_sw_se 中统一汇总。
🔎 列注意力机制逐层拆解
六个子模块虽然任务不同,但共享同一套"双 LSTM 编码 + 注意力"的骨架,理解它一次就全懂了:
- 两个编码器:
enc_n对问题 token 做双向 LSTM 编码,enc_h把每个列头当作一句"伪问题"(pseudo-utterance)单独编码,实现见 encode / encode_hpu; - 注意力打分:
torch.bmm(wenc_hs, self.W_att(wenc_n).transpose(1, 2)),即"列向量 × 问题向量的转置"得到打分矩阵 [bS, 列数, 问题长度]; - 填充惩罚:对 padding 位置置
-1e9,保证 softmax 后权重为 0; - 上下文向量:注意力权重加权求和问题向量得到
c_n,再与列向量拼接经线性层输出最终分数。
各模块的注意力方向各有巧思:
- SCP / WCP(列 → 问题):每一列各自"看"整个问题,回答"这个问题跟我这列有多相关",见 SCP 的 forward 与 WCP 的 forward。WCP 输出逐列独立打分(sigmoid + BCE 损失),因此天然支持多条件、无需排列组合;
- SAP / WOP / WVP_se(问题 → 已选列):固定住已选列,反向对问题词求注意力,得到"与这列最相关的语义上下文",再分类;
- WNP(列自加权):先对列头自身做注意力加权得到表级摘要
c_hs,用它初始化问题编码器的隐藏状态,让 LSTM 带着"表结构先验"去读问题,见 WNP 的 forward。
💡 小细节:代码里每个子模块都保留了
show_p_*可视化开关(如show_p_sc),运行时会画出每列对问题各 token 的注意力权重曲线,是理解模型行为的绝佳调试入口。
🎯 六个子模块逐个看
① SCP:选对列,赢在起跑线。列选择是 NL2SQL 最容易出错的一环(同义词列、相似列名)。SCP 为每一列算一个相关性分数,训练用交叉熵(Loss_sc),推理直接取 argmax(pred_sc)。
② SAP:聚合操作分类器。它取出选中列的编码做问题注意力,把上下文向量压成 6 维 logits(n_agg_ops来自 train.py 中的agg_ops = ['', 'MAX', 'MIN', 'COUNT', 'SUM', 'AVG'])。注意空串''代表"不做聚合"。
③ WNP:条件数量的门控器。输出 5 维 logits(mL_w + 1,即 0–4 个条件)。WikiSQL 的条件数很少超过 4,这一先验让后续模块只需固定长度为 4 的张量处理,大幅简化计算。
④ WCP:序列到集合的核心。与 SCP 结构几乎相同,但语义不同——每列独立判断"是否进入 WHERE"。配合 pred_wc:按分数取前 wn 高的列作为条件列集合。
⑤ WOP:操作符预测。对每条(问题 × 条件列)注意力对,拼接问题上下文c_n与列向量,分类出=、>、<、OP(其他),操作符列表同样定义在 train.py。
⑥ WVP_se:起止区间判别模型。最精巧的一环——不生成值文本,而是给问题每个 token 打两个分:start与end,选出与"列 + 操作符"上下文最匹配的连续区间,见 WVP_se 的 forward。损失函数 Loss_wv_se 对起止位置分别做交叉熵。这种 span 抽取方式避免了开放词表生成的困难,且值必然来自原问题,天然"忠实"。
⚡ 执行引导解码:SQLova-EG 的临门一脚
普通模式下各子模块贪心取最优,误差会级联放大。beam_forward(beam_forward)改为执行引导的束搜索:
- 先对"选列 × 聚合"联合概率取 top-beam,并用
check_sc_sa_pairs过滤类型不匹配的组合(如对文本列做 SUM); - 对 WHERE 条件,把
p_wc × p_wo × p_wv的联合概率排序,取出概率最高的若干候选; - 每个候选调用
engine.execute在真实表上试跑(sqlnet/dbengine.py),只保留有结果的查询; - 最后比较"带 k 个条件"的总概率,决定条件条数,输出结构化 SQL。
这套机制让 SQLova-EG 在测试集上把逻辑形式准确率从 80.7% 提到83.6%,执行准确率提到89.6%。推理入口见 predict.py,训练入口见 train.py。
✅ 小结:一张表读懂 Seq2SQL_v1
| 层次 | 关键设计 | 源码位置 |
|---|---|---|
| 词嵌入 | BERT 表感知编码,问题与列头共享上下文 | sqlova/utils/utils_wikisql.py |
| 子模块骨架 | 双双向 LSTM + 列注意力 + 填充惩罚 | wikisql_models.py |
| 条件抽取 | sigmoid 独立打分 + 起止区间判别 | wikisql_models.py |
| 解码 | 执行引导束搜索 | wikisql_models.py |
| 训练 | 六部件损失联合优化 | wikisql_models.py |
设计哲学一句话:把 NL2SQL 拆成"选列 → 聚合 → 条件骨架 → 操作符 → 值区间"的流水线,用列注意力共享语义理解,再用数据库执行结果做最后把关——结构可解释、误差可控、结果可执行。
如果你想动手实验,可从人评数据(human_eval/README.md)和标注脚本(annotate_ws.py)入手,再结合本文的模块索引阅读sqlova/model/nl2sql/wikisql_models.py,基本可以完整复现整个推理过程。
【免费下载链接】sqlova项目地址: https://gitcode.com/gh_mirrors/sq/sqlova
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考