拓冰建站拓冰建站
首页 / 资讯中心 / 正文

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

人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle 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.10TensorFlow、TensorFlow Probability、NumPy、scikit-learn、tqdmsparse_deferred读取.npz查询图必需建议将training_datasets目录放置在工作目录下下文所有命令均以DATASET_PATHtraining_datasets为前提。二、数据预处理从查询图到 Transformer 输入预处理分为两步先构建全局缩放策略再逐数据集预处理并落盘为 TFRecord。2.1 构建缩放策略Build Scaling Strategy基数与执行时间等数值特征跨数据集差异极大行数从数千到数亿直接喂给模型会导致训练不稳定。因此需要先扫描全部训练数据集统计数值特征的全局分布生成scaling_strategy.json。执行脚本为 data/build_scaling_strategy.pyDATASET_PATHtraining_datasets; DATASET_TYPEbinary_join; # binary_join or single_table python graph_transformer/data/build_scaling_strategy.py \ --dataset_namesaccidents,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_join或single_table决定读取哪套查询图--output_path输出目录缩放策略写入{output_path}/{dataset_type}/scaling_strategy.json缩放策略的底层逻辑对应 build_scaling_strategy.py#L90-L125需要缩放的数值特征由 constants.py 中的SCALING_NUMERICAL_FEATURES定义图级cardinality、exec_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 逐个数据集处理将图结构转换为定长密集张量并写入 TFRecordDATASET_PATHtraining_datasets; DATASET_TYPEbinary_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_filenamescaling_strategy.json done该脚本的核心工作流对应 preprocess_dataset.py#L16-L28 的模块 docstring共 9 步分类特征 One-Hot 编码依据CATEGORICAL_FEATURE_UNIQUE_DICT对data_type6 种、operatorjoin/scan、predicate_operator10 种、相关性validity4 种编码数值特征归一化按 2.1 节缩放策略做log10 → 标准化直方图归一化percentiles_100_numeric按相对值缩放到[0, 1]NaN 置为 -1剔除零基数图真实基数为 0 的样本无法提供有效回归信号直接丢弃剔除无谓词图没有predicate_operator的图对应全表扫描也被移除剔除无用特征按REMOVE_FEATURE_DICT移除name、min/max_numeric、字符串分位数、谓词常量等对模型无意义或难以数值化的字段相关性特征清洗相关系数裁剪到[-1, 1]NaN 置 0见 preprocess_dataset.py#L161-L168加入虚拟节点pseudo node在图首追加一个全零特征的pseudo_node作为读出节点VNODE并通过pseudo_edge连向所有真实节点保证图连通模型最终从该节点读取图级表示计算图结构张量基于邻接矩阵计算最短距离矩阵spatial_encoding、双向空间编码、topological_order以及两种因果掩码parent_causal_mask距离为 1 的父节点可见和ancestor_causal_mask所有可达祖先可见。张量形状约定定义于 constants.py常量值含义MAX_NUM_NODES32单图最大节点数含虚拟节点不足则 paddingNODE_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 是统一训练入口通过参数组合实现三种训练模式并内置了 QErrorq-error评估体系。3.1 三种训练模式模式一实例内训练Instance Based Model——训练集、测试集为同一数据集验证模型在该数据分布内的拟合能力DATASET_PATHtraining_datasets; MODEL_PATHmodels DATASET_TYPEbinary_join; # binary_join or single_table TRAINING_DATASETaccidents; TEST_DATASETaccidents; 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_filenamescaling_strategy.json \ --labelcardinality \ --batch_size128 \ --train_val_sample_size5000 \ --test_sample_size500模式二零样本训练Zero-Shot Model——用其余 19 个数据集训练accidents完全留作测试检验模型的跨库泛化能力这也是 CardBench 的核心卖点DATASET_PATHtraining_datasets; MODEL_PATHmodels; DATASET_TYPEbinary_join; # binary_join or single_table TRAINING_DATASETSairline,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_DATASETaccidents; 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_filenamescaling_strategy.json \ --labelcardinality \ --batch_size128 \ --train_val_sample_size5000 \ --test_sample_size500模式三微调模型Finetuned Model——在零样本模型权重基础上用目标数据集的小样本继续训练兼顾泛化与适配DATASET_PATHtraining_datasets; MODEL_PATHmodels; DATASET_TYPEbinary_join; # binary_join or single_table TRAINING_DATASETaccidents; TEST_DATASETaccidents; BASE_MODEL_CKPT_PATHmodels/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_filenamescaling_strategy.json \ --labelcardinality \ --batch_size128 \ --train_val_sample_size500 \ --test_sample_size500 \ --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_join或single_table--scaling_strategy_filename无缩放策略文件名从{input_dataset_path}/{dataset_type}/下读取--labelcardinality回归标签可选cardinality或exec_time枚举约束--batch_size64批大小--train_val_sample_size5000trainval 总样本数按--train_ratio切分--test_sample_size500测试集样本数--train_ratio0.85train/(trainval) 比例--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_mask或ancestor_causal_mask--reduce_lr_patience5验证损失不降 5 轮后衰减学习率--reduce_lr_factor0.7学习率衰减系数--early_stopping_patience10验证损失不降 10 轮即早停恢复最优权重--base_model_checkpoint_path无微调起始权重路径3.3 数据读取与切分逻辑read_datatrain.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_iterationTrue重打乱。3.4 模型架构要点模型实现于 models/graph_transformer.py核心组件MultiplexNodeFeatureEncoder按节点类型分路复用嵌入。将节点特征一次性嵌入为[B, S, num_node_types × model_dim]再用 one-hot 节点类型向量做多路选择mux得到各类型节点共享参数、又彼此独立的表示graph_transformer.py#L319-L369GraphTransformerEncodernum_encoding_layers层 Transformer 编码器堆叠每层为 LayerNorm → 多头自注意力 → 残差 → FFNGELU。注意力偏置来自空间位置编码spatial_pos_encoder对最短距离矩阵做 Embedding转置为[B, H, S, S]加入 logits注意力掩码采用causal_maskgraph_transformer.py#L449-L500Predictor读出层取编码器输出中index0 的虚拟节点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 越好。QErrorMetrictrain.py#L269-L315在计算前先把标准化标签按缩放策略还原为真实基数unscaled 10 ** (y * log_std log_mean) q_error max(unscaled_true / unscaled_pred, unscaled_pred / unscaled_true)训练与测试过程分别报告mean_q_error及p50 / p75 / p90 / p95 / p99分位 q-error兼顾平均表现与长尾表现测试集上的分位误差由TestModelCallback每轮写入日志train.py#L318-L367。四、端到端复现流程总结将上述步骤串联一次完整的零样本训练实验路径为准备数据按 DowloadArtifacts.md 下载查询图目录结构为training_datasets/{single_table|binary_join}/*.npz构建缩放策略运行build_scaling_strategy.py生成training_datasets/{dataset_type}/scaling_strategy.json预处理循环运行preprocess_dataset.py产出各数据集.tfrecord训练按需选择实例内 / 零样本 / 微调模式运行train.py监控test_p50_q_error、test_p95_q_error等指标模型权重保存于models/graph_transformer.ckpt。关于查询图的结构细节节点/边类型、单表/二表连接/多连接数据集的查询规模表可继续参阅 TrainingQueryGraphs.md其中包含 20 个数据集的图规模统计与 Sparse Deferred 读取示例本模块涉及的关键常量节点类型、边类型、特征裁剪清单则集中在 constants.py是理解数据形态与模型输入之间映射关系的最佳入口。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐TensorFlow Models Transformer 机器翻译模型实战从 WMT 数据预处理到 Keras 训练与 BLEU 评估TensorFlow Models Transformer 机器翻译模型实战从 WMT 数据预处理到 Keras 训练与 BLEU 评估 本文围绕 Tenso人工智能深度学习计算机视觉NLP语音用自有数据预训练 RoBERTa基于 fairseq 的完整实战指南数据处理 · 训练 · 加载用自有数据预训练 RoBERTa基于 fairseq 的完整实战指南数据处理 · 训练 · 加载 导读 本文以 decoding/IAD/fairseq/人工智能大模型预训练深度学习NLP计算机视觉多模态语音音频微调Neurite高级定制指南如何创建自定义节点类型和扩展功能Neurite高级定制指南如何创建自定义节点类型和扩展功能 Neurite是一款功能强大的分形思维导图工具专为AI代理、网络链接、笔记和代码设计。本指南将详创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门