news 2026/9/20 2:56:46

CardBench 零样本基数估计:Graph Transformer 数据预处理与三阶段训练实战指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CardBench 零样本基数估计:Graph Transformer 数据预处理与三阶段训练实战指南
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

CardBench 是 Google Research 开源的学习型基数估计(Learned Cardinality Estimation)基准,本指南聚焦其核心组件——graph_transformer模块,系统讲解如何把查询图(Query Graph)从稀疏的.npz格式预处理为 Transformer 可直接消费的 TFRecord 张量数据,并以"实例内训练 / 零样本训练 / 微调"三种模式运行 Graph Transformer 基数估计模型。读完本文,你将掌握 CardBench 查询图的构建缩放策略(scaling strategy)、数据预处理完整流水线,以及 train.py 全部命令行参数的含义与调优方法,能够独立复现从原始训练数据集到训练出可用的基数预测模型的完整流程。

一、模块定位与前置条件

graph_transformer是 CardBench 仓库中用于基数估计模型训练的独立子模块,位于 CardBench_zero_shot_cardinality_training/graph_transformer。它的输入是 CardBench 数据流水线产出的"带标注查询图"(每个训练样本是一条 SQL 查询及其真实基数,表示为带统计信息的图),输出是一个能够对未见过的查询预测基数的回归模型。

模块目录结构如下:

graph_transformer/ ├── constants.py # 节点/边类型、特征维度等全局常量 ├── train.py # 训练入口(三种模式共用) ├── data/ │ ├── build_scaling_strategy.py # 计算全局数值特征缩放策略 │ └── preprocess_dataset.py # 查询图 → TFRecord 密集张量 └── models/ └── graph_transformer.py # 模型架构实现(图 Transformer 编码器)

1.1 上游数据从哪来

运行本模块之前,需要先准备好 CardBench 训练数据集。两种途径:

  • 直接下载:训练查询图以database_name_<single_table|binary_join|multi_join>.npz的命名方式存放,下载说明见 DowloadArtifacts.md;
  • 自行生成:按照 README.md 描述的流水线(建表 → 计算统计 → 生成 SQL 工作负载 → 执行查询采集真实基数 → 生成查询图)产出.npz查询图文件。

查询图使用 Sparse Deferred 的GraphStruct/InMemoryDB格式存储(详见 TrainingQueryGraphs.md),每个图包含五类节点:tables(表)、attributes(列)、predicates(谓词)、ops(连接/扫描算子)、correlations(列间相关性),以及图级信息g(真实基数、执行时间、SQL 原文、query_id)。图级特征如下:

图级特征含义
cardinality查询真实返回行数(训练标签之一)
exec_time查询执行时间(毫秒,可作为另一个回归标签)
query_id查询唯一标识
querySQL 原文

1.2 运行环境

  • Python ≥ 3.10;
  • TensorFlow、TensorFlow Probability、NumPy、scikit-learn、tqdm;
  • sparse_deferred(读取.npz查询图必需);
  • 建议将training_datasets目录放置在工作目录下(下文所有命令均以DATASET_PATH="training_datasets"为前提)。

二、数据预处理:从查询图到 Transformer 输入

预处理分为两步:先构建全局缩放策略,再逐数据集预处理并落盘为 TFRecord

2.1 构建缩放策略(Build Scaling Strategy)

基数与执行时间等数值特征跨数据集差异极大(行数从数千到数亿),直接喂给模型会导致训练不稳定。因此需要先扫描全部训练数据集,统计数值特征的全局分布,生成scaling_strategy.json。执行脚本为 data/build_scaling_strategy.py:

DATASET_PATH="training_datasets"; DATASET_TYPE="binary_join"; # binary_join or "single_table" python graph_transformer/data/build_scaling_strategy.py \ --dataset_names="accidents,airline,cms_synthetic_patient_data_omop,consumer,covid19_weathersource_com,crypto_bitcoin_cash,employee,ethereum_blockchain,geo_openstreetmap,github_repos,human_variant_annotation,idc_v10,movielens,open_targets_genetics,samples,stackoverflow,tpch_10G,usfs_fia,uspto_oce_claims,wikipedia" \ --input_dataset_path=$DATASET_PATH \ --dataset_type=$DATASET_TYPE \ --output_path=$DATASET_PATH

参数说明:

参数说明
--dataset_names参与统计的数据集名列表(逗号分隔),必须与training_datasets/<dataset_type>/下存在的.npz文件一一对应
--input_dataset_path查询图文件根目录,脚本会在其下按{dataset_type}/子目录查找{dataset_name}_{dataset_type}.npz
--dataset_type枚举值binary_joinsingle_table,决定读取哪套查询图
--output_path输出目录,缩放策略写入{output_path}/{dataset_type}/scaling_strategy.json

缩放策略的底层逻辑(对应 build_scaling_strategy.py#L90-L125):

  • 需要缩放的数值特征由 constants.py 中的SCALING_NUMERICAL_FEATURES定义:图级cardinalityexec_time,表级rows,列级num_unique
  • 对每个特征统计mean / std / median / min / max,同时统计log 域log_mean / log_std / log_median / log_max / log_min
  • 关键设计:所有特征一律先取log10(下限截断为1e-9)再标准化,即log_scale恒为True。这样既压缩了长尾分布,又让不同量级的特征(如行数与基数)在训练时可比。

2.2 预处理数据集(Preprocess Dataset)

缩放策略就绪后,用 data/preprocess_dataset.py 逐个数据集处理,将图结构转换为定长密集张量并写入 TFRecord:

DATASET_PATH="training_datasets"; DATASET_TYPE="binary_join"; # binary_join or "single_table" for DATASET_NAME in "accidents" "airline" "cms_synthetic_patient_data_omop" \ "consumer" "covid19_weathersource_com" "crypto_bitcoin_cash" "employee" \ "ethereum_blockchain" "geo_openstreetmap" "github_repos" "human_variant_annotation" \ "idc_v10" "movielens" "open_targets_genetics" "samples" "stackoverflow" "tpch_10G" \ "usfs_fia" "uspto_oce_claims" "wikipedia"; do python graph_transformer/data/preprocess_dataset.py \ --dataset_name=$DATASET_NAME \ --input_dataset_path=$DATASET_PATH \ --output_path=$DATASET_PATH \ --dataset_type=$DATASET_TYPE \ --scaling_strategy_filename="scaling_strategy.json" done

该脚本的核心工作流(对应 preprocess_dataset.py#L16-L28 的模块 docstring,共 9 步):

  1. 分类特征 One-Hot 编码:依据CATEGORICAL_FEATURE_UNIQUE_DICTdata_type(6 种)、operator(join/scan)、predicate_operator(10 种)、相关性validity(4 种)编码;
  2. 数值特征归一化:按 2.1 节缩放策略做log10 → 标准化
  3. 直方图归一化percentiles_100_numeric按相对值缩放到[0, 1],NaN 置为 -1;
  4. 剔除零基数图:真实基数为 0 的样本无法提供有效回归信号,直接丢弃;
  5. 剔除无谓词图:没有predicate_operator的图(对应全表扫描)也被移除;
  6. 剔除无用特征:按REMOVE_FEATURE_DICT移除namemin/max_numeric、字符串分位数、谓词常量等对模型无意义或难以数值化的字段;
  7. 相关性特征清洗:相关系数裁剪到[-1, 1],NaN 置 0(见 preprocess_dataset.py#L161-L168);
  8. 加入虚拟节点(pseudo node):在图首追加一个全零特征的pseudo_node作为读出节点(VNODE),并通过pseudo_edge连向所有真实节点,保证图连通,模型最终从该节点读取图级表示;
  9. 计算图结构张量:基于邻接矩阵计算最短距离矩阵spatial_encoding、双向空间编码、topological_order,以及两种因果掩码parent_causal_mask(距离为 1 的父节点可见)和ancestor_causal_mask(所有可达祖先可见)。

张量形状约定(定义于 constants.py):

常量含义
MAX_NUM_NODES32单图最大节点数(含虚拟节点),不足则 padding
NODE_FEATURE_DIM150节点特征维度,不足补零
NODE_TYPES6 种pseudo_node / attributes / ops / predicates / correlations / tables

预处理产出的 TFRecord 中,每个 example 包含node[32, 150])、node_padding[32])、parent_causal_mask/ancestor_causal_mask/spatial_encoding(各[32, 32])、topological_order[32, 1])以及标签cardinality/exec_time。文件命名与输入一致:{output_path}/{dataset_type}/{dataset_name}_{dataset_type}.tfrecord

三、训练 Graph Transformer 模型

graph_transformer/train.py 是统一训练入口,通过参数组合实现三种训练模式,并内置了 QError(q-error)评估体系。

3.1 三种训练模式

模式一:实例内训练(Instance Based Model)——训练集、测试集为同一数据集,验证模型在该数据分布内的拟合能力:

DATASET_PATH="training_datasets"; MODEL_PATH="models" DATASET_TYPE="binary_join"; # binary_join or "single_table" TRAINING_DATASET="accidents"; TEST_DATASET="accidents"; python graph_transformer/train.py \ --training_dataset_names=$TRAINING_DATASET \ --test_dataset_name=$TEST_DATASET \ --input_dataset_path=$DATASET_PATH \ --model_path=$MODEL_PATH \ --dataset_type=$DATASET_TYPE \ --scaling_strategy_filename="scaling_strategy.json" \ --label="cardinality" \ --batch_size=128 \ --train_val_sample_size=5000 \ --test_sample_size=500

模式二:零样本训练(Zero-Shot Model)——用其余 19 个数据集训练,accidents完全留作测试,检验模型的跨库泛化能力,这也是 CardBench 的核心卖点:

DATASET_PATH="training_datasets"; MODEL_PATH="models"; DATASET_TYPE="binary_join"; # binary_join or "single_table" TRAINING_DATASETS="airline,cms_synthetic_patient_data_omop,consumer,covid19_weathersource_com,crypto_bitcoin_cash,employee,ethereum_blockchain,geo_openstreetmap,github_repos,human_variant_annotation,idc_v10,movielens,open_targets_genetics,samples,stackoverflow,tpch_10G,usfs_fia,uspto_oce_claims,wikipedia"; TEST_DATASET="accidents"; python graph_transformer/train.py \ --training_dataset_names=$TRAINING_DATASETS \ --test_dataset_name=$TEST_DATASET \ --input_dataset_path=$DATASET_PATH \ --model_path=$MODEL_PATH \ --dataset_type=$DATASET_TYPE \ --scaling_strategy_filename="scaling_strategy.json" \ --label="cardinality" \ --batch_size=128 \ --train_val_sample_size=5000 \ --test_sample_size=500

模式三:微调模型(Finetuned Model)——在零样本模型权重基础上,用目标数据集的小样本继续训练,兼顾泛化与适配:

DATASET_PATH="training_datasets"; MODEL_PATH="models"; DATASET_TYPE="binary_join"; # binary_join or "single_table" TRAINING_DATASET="accidents"; TEST_DATASET="accidents"; BASE_MODEL_CKPT_PATH="models/graph_transformer.ckpt" python graph_transformer/train.py \ --training_dataset_names=$TRAINING_DATASET \ --test_dataset_name=$TEST_DATASET \ --input_dataset_path=$DATASET_PATH \ --model_path=$MODEL_PATH \ --dataset_type=$DATASET_TYPE \ --scaling_strategy_filename="scaling_strategy.json" \ --label="cardinality" \ --batch_size=128 \ --train_val_sample_size=500 \ --test_sample_size=500 \ --base_model_checkpoint_path=$BASE_MODEL_CKPT_PATH

注意微调场景下--train_val_sample_size收窄到 500,体现"用小量标注数据适配新库"的意图;模型参数会从graph_transformer.ckpt恢复,学习率重置为--init_lr(见 train.py#L455-L457)。

3.2 全部训练参数速查表

以下参数均在 train.py 中定义,可直接覆盖默认值:

参数默认值说明
--training_dataset_names必填训练数据集列表,逗号分隔;同路径时按比例切分 train/val/test
--test_dataset_name必填测试数据集名
--input_dataset_path必填TFRecord 数据根目录
--model_path必填模型检查点输出目录(自动创建),权重写入{model_path}/graph_transformer.ckpt
--dataset_type必填binary_joinsingle_table
--scaling_strategy_filename缩放策略文件名,从{input_dataset_path}/{dataset_type}/下读取
--labelcardinality回归标签,可选cardinalityexec_time(枚举约束)
--batch_size64批大小
--train_val_sample_size5000train+val 总样本数,按--train_ratio切分
--test_sample_size500测试集样本数
--train_ratio0.85train/(train+val) 比例
--num_epochs200最大训练轮数
--init_lr1e-3初始学习率
--min_lr1e-5学习率衰减下限
--num_encoding_layers16Transformer 编码器层数
--num_embedding_layers3节点特征嵌入 MLP 层数
--num_output_layers3输出头 MLP 层数
--model_dim128模型隐藏维度(编码器与输出头共用)
--num_heads8多头注意力头数(model_dim需可整除)
--dropout0.0Dropout 比率
--mask_typeancestor_causal_mask因果掩码类型,可选parent_causal_maskancestor_causal_mask
--reduce_lr_patience5验证损失不降 5 轮后衰减学习率
--reduce_lr_factor0.7学习率衰减系数
--early_stopping_patience10验证损失不降 10 轮即早停(恢复最优权重)
--base_model_checkpoint_path微调起始权重路径

3.3 数据读取与切分逻辑

read_data(train.py#L172-L266)按以下规则组织数据:

  • {input_dataset_path}/{dataset_type}/{name}_{dataset_type}.tfrecord读取,parse_example解析出特征字典与标签(train.py#L141-L169);
  • 测试数据集与训练集路径相同(实例内训练/微调)时,先从数据集头部切出test_size个样本作测试,其余再按train_ratio分成训练与验证集;
  • 当测试集不在训练列表中(零样本训练)时,训练/验证/测试分别从各自文件读取;
  • 所有序列均按[32, 150]等形状padded_batch补齐,训练集每轮reshuffle_each_iteration=True重打乱。

3.4 模型架构要点

模型实现于 models/graph_transformer.py,核心组件:

  • MultiplexNodeFeatureEncoder:按节点类型分路复用嵌入。将节点特征一次性嵌入为[B, S, num_node_types × model_dim],再用 one-hot 节点类型向量做多路选择(mux),得到各类型节点共享参数、又彼此独立的表示(graph_transformer.py#L319-L369);
  • GraphTransformerEncodernum_encoding_layers层 Transformer 编码器堆叠,每层为 LayerNorm → 多头自注意力 → 残差 → FFN(GELU)。注意力偏置来自空间位置编码spatial_pos_encoder对最短距离矩阵做 Embedding,转置为[B, H, S, S]加入 logits),注意力掩码采用causal_mask(graph_transformer.py#L449-L500);
  • Predictor读出层:取编码器输出中index=0 的虚拟节点(VNODE)表示作为图嵌入,经 LayerNorm 与num_output_layers层 MLP 回归出标量基数(graph_transformer.py#L609-L614)。

训练配置:Adam 优化器 + MeanAbsoluteError 损失;训练回调包括 EarlyStopping、ReduceLROnPlateau、按 epoch 保存权重的 ModelCheckpoint,以及自定义TestModelCallback(每轮在测试集上滚动记录 QError)。

3.5 QError 评估指标

基数估计领域通用评估指标是q-error:真实基数与预测基数之比中较大的一个(max(real/pred, pred/real),越接近 1 越好)。QErrorMetric(train.py#L269-L315)在计算前先把标准化标签按缩放策略还原为真实基数:

unscaled = 10 ** (y * log_std + log_mean) q_error = max(unscaled_true / unscaled_pred, unscaled_pred / unscaled_true)

训练与测试过程分别报告mean_q_errorp50 / p75 / p90 / p95 / p99分位 q-error,兼顾平均表现与长尾表现;测试集上的分位误差由TestModelCallback每轮写入日志(train.py#L318-L367)。

四、端到端复现流程总结

将上述步骤串联,一次完整的零样本训练实验路径为:

  1. 准备数据:按 DowloadArtifacts.md 下载查询图,目录结构为training_datasets/{single_table|binary_join}/*.npz
  2. 构建缩放策略:运行build_scaling_strategy.py,生成training_datasets/{dataset_type}/scaling_strategy.json
  3. 预处理:循环运行preprocess_dataset.py,产出各数据集.tfrecord
  4. 训练:按需选择实例内 / 零样本 / 微调模式运行train.py,监控test_p50_q_errortest_p95_q_error等指标,模型权重保存于models/graph_transformer.ckpt

关于查询图的结构细节(节点/边类型、单表/二表连接/多连接数据集的查询规模表)可继续参阅 TrainingQueryGraphs.md,其中包含 20 个数据集的图规模统计与 Sparse Deferred 读取示例;本模块涉及的关键常量(节点类型、边类型、特征裁剪清单)则集中在 constants.py,是理解数据形态与模型输入之间映射关系的最佳入口。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 2:55:34

方案 A vs 方案 B

方案 A vs 方案 B 【免费下载链接】baoyu-skills 项目地址: https://gitcode.com/gh_mirrors/ba/baoyu-skills Overview 一张对比两种技术选型优劣势的双栏信息图。 Learning Objectives 观众将理解&#xff1a;两方案的核心差异、各自适用场景、最终推荐。 Section…

作者头像 李华
网站建设 2026/9/20 2:54:21

多代理编排实战:5分钟让一句提问被路由到最合适的AI代理

多代理编排实战&#xff1a;5分钟让一句提问被路由到最合适的AI代理 【免费下载链接】agent-squad Flexible and powerful framework for managing multiple AI agents and handling complex conversations 项目地址: https://gitcode.com/GitHub_Trending/mu/agent-squad …

作者头像 李华
网站建设 2026/9/20 2:50:46

C盘爆红不用怕:4个文件夹清理法,10分钟释放100G

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/20 2:50:15

KingbaseES数据库对象权限管理实战:从授权到角色设计

1. 项目概览&#xff1a;为什么数据库对象权限管理是刚需先说一个扎心的事实&#xff1a;大部分数据库安全问题&#xff0c;不是被外部攻击攻破的&#xff0c;而是内部权限失控导致的。某个开发同事离职后账号没回收、某个应用账号用了超级用户权限跑业务、某张薪酬表人人都能S…

作者头像 李华
网站建设 2026/9/20 2:46:58

DeepSeek Harness插件实战:生产级AI应用编排必备10款插件解析

用DeepSeek Harness做AI应用编排的人&#xff0c;最近应该都感受到了插件生态带来的变化。以前我们搭一个Agent工作流&#xff0c;路由、诊断、知识库、日志全要自己写代码&#xff0c;现在dsh插件市场里几十个现成插件可以装&#xff0c;其中一个叫Code Sentinel的代码诊断插件…

作者头像 李华