- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
导读
本指南围绕google-research仓库中 generative_trees 目录的官方 README 与配套 Java 源码展开,系统讲解 ICML'22 论文Generative Trees: Adversarial and Copycat(作者 Richard Nock 与 Mathieu Guillame-Bert)的配套实现:如何用copycat 方法训练一棵"生成树"(Generative Tree,GT)作为数据生成模型,以及如何从预训练生成树中批量采样新样本、绘制二维密度图、完成缺失值插补。读完本文,你将掌握:编译与运行示例脚本的完整流程、Wrapper与Generate两个入口的全部命令行参数及含义、copycat 训练循环的源码级运行机制,以及生成树的保存格式与统计输出结构,能够直接在自己的数据集上复现该论文的方法。
一、项目背景与核心思想
generative_trees是 ICML'22 论文Generative Trees: Adversarial and Copycat的官方配套代码(BibTeX 见 README.md)。其核心方法名为copycat:训练过程不直接优化生成树本身,而是让一个判别树(Discriminator Tree,DT)持续对"真实数据 vs 生成数据"做区分,生成树(Generator Tree,GT)则"模仿"判别树最新一次成功分裂的结构,二者交替进化,从而逐步逼近真实数据分布。README 将其概括为两大能力:
- 训练:使用 copycat 方法训练生成模型(入口类
Wrapper); - 生成:使用预训练好的模型直接生成样本或绘制密度图(入口类
Generate)。
代码使用纯 Java 编写,不依赖第三方机器学习框架,编译后即可运行。
二、快速上手:编译并运行官方示例
README 给出了最简运行方式:克隆仓库后进入generative_trees/目录,直接执行示例脚本:
git clone https://github.com/google-research/google-research.git cd google-research/generative_trees/ run_example.sh当前仓库中的 run_example.sh 完整实现了 README 描述的流程,整个脚本只有 20 余行,清晰展示了"编译 → 查看帮助 → 下载数据 → 训练并采样 → 展示结果"的完整链路:
set -e echo "Compile Generative Decision Trees project" javac -d compiled -classpath src src/*.java echo "Prints the help" java -classpath compiled Wrapper --help echo "Download a copy of the Iris dataset" wget https://raw.githubusercontent.com/google/yggdrasil-decision-forests/main/yggdrasil_decision_forests/test_data/dataset/iris.csv -O iris.csv echo "Train and sample a new generator" mkdir -p working_dir java -classpath compiled Wrapper \ --dataset=iris.csv \ --work_dir=working_dir \ --num_samples=1000 \ --output_samples=working_dir/generated.csv \ --output_stats=working_dir/statistics.stats echo "Display some of the generated samples" head working_dir/generated.csv脚本的执行链路可以拆解为四步:
- 编译:
javac -d compiled -classpath src src/*.java将 src 下全部 17 个 Java 源文件编译到compiled/目录。运行示例的数据集是经典的Iris(鸢尾花)数据集(从 yggdrasil-decision-forests 项目下载,含Sepal.Length、Sepal.Width、Petal.Length、Petal.Width四个数值特征和class标签)。 - 查看帮助:
java -classpath compiled Wrapper --help打印Wrapper的全部参数说明(详见下文第四节)。 - 训练并采样:
Wrapper以 Iris 数据为输入,通过 copycat 方法训练一棵生成树,并用它生成 1000 个新样本,保存到working_dir/generated.csv;运行统计信息写入working_dir/statistics.stats。 - 展示结果:用
head打印生成样本的前几行。
README 给出了脚本执行完毕后的预期输出——一组带class标签的合成样本,其数值与真实 Iris 分布接近但又不完全相同,例如:
Display some of the generated samples Sepal.Length,Sepal.Width,Petal.Length,Petal.Width,class 5.117246154727025,3.294665099621395,1.5415873061790373,0.34693251377205403,setosa 4.938340282187983,3.306772168630169,2.1019090151019775,0.4123936617890174,setosa 5.577975609495907,3.453786064420899,3.671345561310016,0.7218885473617979,versicolor 6.146461600520874,3.7586348414987745,5.222165962947139,2.3442290234292913,versicolor运行前提:机器上需要安装 JDK(含javac/java)以及wget用于下载数据集。
说明:README 还提到用于缺失数据插补的脚本
script-missing-data-imputation.sh("automates the process, can be edited easily"),不过在本文所基于的仓库快照中该脚本并未包含,缺失值插补能力可以直接通过Wrapper的--impute_missing=true参数触发(见第六节)。
三、两大核心入口:Wrapper与Generate
README 明确指出本项目由两个关键类构成,二者职责分离,对应源码分别位于 Wrapper.java 与 Generate.java:
| 入口类 | 职责 | 典型用法 |
|---|---|---|
Wrapper | 从数据训练生成树(copycat 方法),随后采样、保存生成器、输出统计 | java Wrapper --help查看全部选项 |
Generate | 加载预训练生成树,仅做生成样本或密度图 | java Generate --help查看全部选项 |
Wrapper.main与Generate.main在无参数启动时都会提示*No parameters*. Run 'java <Class> --help' for more,并在退出前打印论文的 BibTeX 引用信息(见 Wrapper.java)。
四、Wrapper:copycat 训练生成树(参数全解)
Wrapper的完整帮助文本内嵌于源码的help()方法中(Wrapper.java),README 与脚本只展示了它的一个子集。以下按源码整理全部参数。
4.1 基础参数
| 参数 | 类型 | 必填 | 含义 |
|---|---|---|---|
--dataset= | String | 是 | CSV 数据文件路径,首行必须包含变量名(表头) |
--dataset_spec= | JSON 字符串 | 否 | 数据集规格说明,含name/path/label/task四个字段(见 4.2) |
--num_samples= | int | 是 | 要生成的样本数量 |
--work_dir= | String | 是 | 生成树模型与密度图文件的保存目录 |
--output_samples= | String | 是 | 生成样本的输出文件名 |
--output_stats= | String | 是 | 运行统计文件(运行时长、GT 边缘直方图、GT 树节点统计等) |
--x=--y= | String | 否 | (可选)用于保存二维密度图的两个变量名,输出格式为(x, y, density_value_at_(x,y)) |
--flags= | JSON 字符串 | 否 | 训练超参数(见 4.3) |
--impute_missing= | boolean | 否 | 若为true,用生成树对训练数据中的缺失值进行插补 |
源码中的参数解析位于Wrapper.fit_vars(Wrapper.java),例如--output_samples=会被拆分为输出目录path_to_generated_samples与文件名blueprint_save_name,而生成树模型文件则自动命名为generator_<文件名>(前缀由常量PREFIX_GENERATOR = "generator_"定义),保存在--work_dir下。
4.2--dataset_spec数据集规格说明
--dataset_spec接收一段 JSON 形式的字符串,源码按顺序解析四个 token(见 Wrapper.java 中定义的DATASET_TOKENS):
'--dataset_spec={"name": "iris", "path": "${ANYDIR}/Datasets/iris/iris.csv", "label": "class", "task": "BINARY_CLASSIFICATION"}'name:数据集前缀名(同时也用于缺失值插补输出文件的命名,见第六节);path:数据文件路径;label:标签列名;task:任务类型(如BINARY_CLASSIFICATION)。
源码要求这四个 token 在字符串中按上述顺序出现且各出现一次,否则会报错("more than one occurrence of ..." 或 "zero occurrence of ...")。若同时给出了--dataset=与--dataset_spec中的path,两者不一致时会打印Non identical information in --dataset_spec path vs --dataset警告。
4.3--flags训练超参数
--flags接收一段 JSON 格式的{"name" : value, ...}字符串,源码支持的全部 flag 及其默认值(定义在 Wrapper.java 的ALL_FLAGS数组,默认值见 L76-L81):
'--flags={"iterations" : "10", "force_integer_coding" : "true", "force_binary_coding" : "true", "faster_induction" : "true", "unknown_value_coding" : "?", "number_bins_for_histograms" : "11"}'| Flag | 类型 | 默认值 | 含义 |
|---|---|---|---|
iterations | int | 无(必填) | GT 中的分裂次数;最终节点数 = 2 × iterations + 1 |
force_integer_coding | boolean | false | 为true时把可识别为整数的变量按整数编码(否则按 double 编码),生成更"干净"的 GT |
force_binary_coding | boolean | true | 为true时把 0/1/unknown 变量识别为名义变量(nominal),否则按整数或 double 处理 |
faster_induction | boolean | false | 为true时若候选分裂过多(超过Discriminator_Tree.MAX_SPLITS_BEFORE_RANDOMISATION,源码默认 1000)则对 DT 分裂做随机采样以加速训练 |
unknown_value_coding | String | "-1" | 数据集中"未知值"的表示符号,会写入全局常量Unknown_Feature_Value.S_UNKNOWN |
number_bins_for_histograms | int | 19 | 非名义变量的直方图分箱数,用于训练结束后计算 GT 边缘分布直方图(同时设置Histogram.NUMBER_CONTINUOUS_FEATURE_BINS与MAX_NUMBER_INTEGER_FEATURE_BINS) |
copycat_local_generation | boolean | true | copycat 归纳中 GT 每新增一次分裂后,只对受影响叶子的本地生成样本替换对应特征;为false时用整棵 GT 重新生成全部样本(对应Boost.COPYCAT_GENERATE_WITH_WHOLE_GT,见第五节) |
从源码看,--flags解析要求字符串必须以{开头、以}结尾,每个条目为tag:value形式,且 tag 必须在ALL_FLAGS白名单内,否则直接报错终止(Wrapper.java)。
4.4 完整示例命令行
源码help()中给出的完整示例(--x/--y指定密度图坐标轴、--impute_missing=true开启插补):
java Wrapper --dataset=${ANYDIR}/Datasets/iris/iris.csv \ '--dataset_spec={"name": "iris", "path": "${ANYDIR}/Datasets/iris/iris.csv", "label": "class", "task": "BINARY_CLASSIFICATION"}' \ --num_samples=10000 \ --work_dir=${ANYDIR}/Datasets/iris/working_dir \ --output_samples=${ANYDIR}/Datasets/iris/output_samples/iris_gt_generated.csv \ --output_stats=${ANYDIR}/Datasets/iris/results/generated_examples.stats \ --x=Sepal.Length --y=Sepal.Width \ '--flags={"iterations" : "10", "force_integer_coding" : "true", "force_binary_coding" : "true", "faster_induction" : "true", "unknown_value_coding" : "?", "number_bins_for_histograms" : "11"}' \ --impute_missing=true注意:源码不允许--x与--y指向同一变量(会报density plot requested on the same X and Y variable)。
五、Generate:从预训练生成树采样
Generate用于加载由Wrapper保存的生成树文件,仅做推断与生成。其帮助文本位于 Generate.java,参数以短选项形式提供:
java -Xmx10000m Generate -D Datasets/generate/ -P open_policing_hartford -U NA -F true -N 1000 -L example-generator_open_policing_hartford.csv| 参数 | 类型 | 必填 | 含义 |
|---|---|---|---|
-D | String | 是 | 数据所在目录 |
-P | String | 是 | 域(domain)前缀;数据文件必须位于Datasets/generate/open_policing_hartford.csv(即<目录>/<前缀>.csv) |
-L | String | 是 | 生成树模型文件名(必须位于上述目录中) |
-N | int | 否 | 要生成的样本数;生成文件与模型同目录,命名为<前缀>_GeneratedSample.csv;不指定则只显示生成树结构 |
-U | String | 否 | 数据集中"未知值"的表示,默认"-1" |
-F | boolean | 否 | 是否强制整数编码,默认false |
-X/-Y | String | 否 | 用于二维密度图的 x/y 变量名 |
Generate.go(Generate.java)的执行流程为:
- 按
<目录>/<前缀>/<前缀>.csv定位并加载原始数据,构造数据域(Domain); - 调用
from_file(Generate.java)解析生成树模型文件:文件以@NODES/@ARCS两个区段分别描述节点与弧(边),逐行重建Generator_Node、Generator_Arc及父子关系、叶子集合与树深度; - 若
-N指定了样本数,调用gt.generate_sample_with_density(number_ex)(当指定-X/-Y时)或gt.generate_sample(number_ex)生成样本; - 输出
<前缀>_GeneratedSample.csv;若指定了-X/-Y,另输出<前缀>_GeneratedSample_DENSITY_X_<x>_Y_<y>.csv,列为x,y,generated_density。
六、源码级原理:copycat 训练循环
Wrapper的训练核心由 Boost.java 的simple_boost_copycat实现,它把"判别树 vs 生成树"的交替博弈固化为一个循环:
初始化: 1) 创建生成树 GT(Generator_Tree),并初始化根节点 2) 创建判别树 DT(Discriminator_Tree),把所有真实训练样本挂到根叶子 3) 用 GT 生成一批"假"样本(myDomain.myDS.generate_examples(gt)) 循环(直到达到 iterations 或无法再分裂): 4) 计算假样本在 DT 各节点中的训练折叠索引 5) DT 执行一步生长 one_step_grow() —— 找到最佳分裂 6) 若分裂成功(DT_SPLIT_OK): - 找到 GT 中与 DT 被分裂叶子"对应"的叶子(gt.get_leaf_to_be_split) - GT 以 copycat 方式生长(gt.one_step_grow_copycat),模仿 DT 的新分裂 - 根据 Boost.COPYCAT_GENERATE_WITH_WHOLE_GT 决定: 为 true(默认):用整棵 GT 重新生成全部假样本 为 false:仅对刚分裂的 GT 叶子局部重生成(generate_and_replace_examples)相关核心结构:
- 判别树Discriminator_Tree.java:负责在真实/假样本上找最佳分裂。源码中
RANDOMISE_SPLIT_FINDING_WHEN_TOO_MANY_SPLITS、MAX_SPLITS_BEFORE_RANDOMISATION(默认 1000)、MAX_CARD_MODALITIES_BEFORE_RANDOMISATION(默认 10)等静态常量控制了名义变量候选分裂过多时的随机化加速策略,与--flags中的faster_induction直接对应。 - 生成树Generator_Tree.java:节点上记录分裂特征、各分支概率(
multi_p)与子节点,叶子节点构成可继续生长的集合。训练结束后compute_generator_histograms()会从 GT 采样一批样本,为每个特征计算边缘分布直方图(分箱数由--flags的number_bins_for_histograms控制)。 - 调度入口Algorithm.java:
simple_go()将参数封装为["@MatuErr", "1.0", "COPYCAT", iterations, copycat_local_generation]交给Boost;其中策略名COPYCAT由Boost.KEY_NAME白名单校验(Boost.java)。 - 数据域Domain.java:负责加载特征与样本(
Dataset)、计算域直方图,并挂载内存监控器(MemoryMonitor)。
Wrapper.simple_go(Wrapper.java)按固定流水线串联:加载数据 → 学习 GT → 计算 GT 边缘直方图 → 保存 GT 到work_dir/generator_<name>→ 生成样本 →(可选)缺失值插补 → 保存样本 →(可选)保存二维密度图 → 保存统计文件,每一步都打印耗时(毫秒)。
七、缺失值插补(--impute_missing)
README 明确将缺失值插补列为该代码的关键能力之一。当训练数据含有未知值(默认编码为"-1",可用unknown_value_coding修改)且传入--impute_missing=true时:
Wrapper.simple_go在生成样本后调用impute_and_save(gt)(Wrapper.java);- 逐行扫描原始训练样本,对含未知值的样本调用
gt.impute_all_values_from_one_leaf(...)——让样本沿生成树落至叶子,用该叶子对应的分布对缺失特征进行最大似然补全(对应Generator_Tree.IMPUTATION_AT_MAXIMUM_LIKELIHOOD常量); - 插补结果保存为
<work_dir>/<spec_name>_imputed.csv(spec_name来自--dataset文件名或--dataset_spec的name),并在运行摘要中打印该路径。
八、统计输出与结果文件
Wrapper每次运行会产出多类文件,--output_stats指定主统计文件路径:
- 主统计文件(JSON 格式,Wrapper.java):包含
running_time_seconds(训练+生成总时长)、gt_number_nodes(GT 节点数)、gt_depth(GT 深度)、running_time_gt_training_plus_exemple_generation,以及开启插补时的running_time_gt_training_plus_imputation。 - 附加统计文件(
<output_stats>_more.txt,L429-L456):记录本次运行使用的全部 flag 值、GT 训练与采样各自耗时、每个特征的 GT 边缘分布直方图(便于与真实数据分布对比),以及GT node counts per feature name(每个特征在 GT 中的节点计数)。 - 生成树模型:
<work_dir>/generator_<输出样本名>,以@NODES/@ARCS文本格式序列化(见 Generator_Tree.java),可被Generate -L重新加载。 - 二维密度图:指定
--x/--y后输出<工作目录>/<样本名>_2DDensity_plot_X_<x>_Y_<y>.csv,格式为x,y,density_value。
九、小结与扩展阅读
generative_trees提供了一个无第三方依赖的完整"训练—生成"闭环:Wrapper以 copycat 方式在判别树与生成树的交替博弈中学习数据分布,Generate负责从已保存的生成树批量采样与绘制密度图,--impute_missing则让同一棵生成树兼任缺失值插补器。结合源码可知,--flags中的iterations直接决定生成树规模(节点数 = 2×iterations+1),faster_induction、copycat_local_generation等开关则分别控制训练加速与局部/全局重生成策略,为复现论文实验或调整自己的数据管线提供了清晰的旋钮。
进一步探索可阅读:
- generative_trees/README.md:官方说明与引用信息;
- generative_trees/run_example.sh:开箱即用的端到端示例;
- generative_trees/src/Wrapper.java:全部训练参数的内置帮助文档;
- generative_trees/src/Generate.java:生成入口参数说明;
- generative_trees/src/Boost.java:copycat 训练循环核心实现;
- generative_trees/src/Generator_Tree.java 与 generative_trees/src/Discriminator_Tree.java:生成树/判别树的数据结构与生长逻辑。
使用该代码复现论文结果时,请引用:
@inproceedings{ngbGT, title={Generative Trees: Adversarial and Copycat}, author={R. Nock and M. Guillame-Bert}, booktitle={39$^{~th}$ International Conference on Machine Learning}, year={2022} }- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
Generative Forests 实战指南:基于 google-research 生成式树集成模型的编译、训练与数据生成
Generative Forests 实战指南:基于 google research 生成式树集成模型的编译、训练与数据生成 本文是 Google Resear
人工智能深度学习NLP计算机视觉强化学习RP2040 CAN总线通信终极指南:5个步骤掌握can2040实战应用
RP2040 CAN总线通信终极指南:5个步骤掌握can2040实战应用 在嵌入式系统和物联网设备开发中,CAN总线通信一直是工业控制、汽车电子和机器人领域的核
基于 fairseq 的分层神经故事生成实战指南:WritingPrompts 数据预处理、卷积模型训练与采样生成
基于 fairseq 的分层神经故事生成实战指南:WritingPrompts 数据预处理、卷积模型训练与采样生成 导读 本文基于 kosmos 2/fairs
人工智能大模型预训练深度学习NLP计算机视觉多模态语音音频微调
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考