Generative Trees(生成树)代码实战指南:基于 Copycat 对抗训练的生成模型与采样(ICML‘22 配套代码)
2026/9/20 20:35:02 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

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

导读

本指南围绕google-research仓库中 generative_trees 目录的官方 README 与配套 Java 源码展开,系统讲解 ICML'22 论文Generative Trees: Adversarial and Copycat(作者 Richard Nock 与 Mathieu Guillame-Bert)的配套实现:如何用copycat 方法训练一棵"生成树"(Generative Tree,GT)作为数据生成模型,以及如何从预训练生成树中批量采样新样本、绘制二维密度图、完成缺失值插补。读完本文,你将掌握:编译与运行示例脚本的完整流程、WrapperGenerate两个入口的全部命令行参数及含义、copycat 训练循环的源码级运行机制,以及生成树的保存格式与统计输出结构,能够直接在自己的数据集上复现该论文的方法。


一、项目背景与核心思想

generative_trees是 ICML'22 论文Generative Trees: Adversarial and Copycat的官方配套代码(BibTeX 见 README.md)。其核心方法名为copycat:训练过程不直接优化生成树本身,而是让一个判别树(Discriminator Tree,DT)持续对"真实数据 vs 生成数据"做区分,生成树(Generator Tree,GT)则"模仿"判别树最新一次成功分裂的结构,二者交替进化,从而逐步逼近真实数据分布。README 将其概括为两大能力:

  1. 训练:使用 copycat 方法训练生成模型(入口类Wrapper);
  2. 生成:使用预训练好的模型直接生成样本或绘制密度图(入口类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

脚本的执行链路可以拆解为四步:

  1. 编译javac -d compiled -classpath src src/*.java将 src 下全部 17 个 Java 源文件编译到compiled/目录。运行示例的数据集是经典的Iris(鸢尾花)数据集(从 yggdrasil-decision-forests 项目下载,含Sepal.LengthSepal.WidthPetal.LengthPetal.Width四个数值特征和class标签)。
  2. 查看帮助java -classpath compiled Wrapper --help打印Wrapper的全部参数说明(详见下文第四节)。
  3. 训练并采样Wrapper以 Iris 数据为输入,通过 copycat 方法训练一棵生成树,并用它生成 1000 个新样本,保存到working_dir/generated.csv;运行统计信息写入working_dir/statistics.stats
  4. 展示结果:用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参数触发(见第六节)。


三、两大核心入口:WrapperGenerate

README 明确指出本项目由两个关键类构成,二者职责分离,对应源码分别位于 Wrapper.java 与 Generate.java:

入口类职责典型用法
Wrapper从数据训练生成树(copycat 方法),随后采样、保存生成器、输出统计java Wrapper --help查看全部选项
Generate加载预训练生成树,仅做生成样本或密度图java Generate --help查看全部选项

Wrapper.mainGenerate.main在无参数启动时都会提示*No parameters*. Run 'java <Class> --help' for more,并在退出前打印论文的 BibTeX 引用信息(见 Wrapper.java)。


四、Wrapper:copycat 训练生成树(参数全解)

Wrapper的完整帮助文本内嵌于源码的help()方法中(Wrapper.java),README 与脚本只展示了它的一个子集。以下按源码整理全部参数。

4.1 基础参数

参数类型必填含义
--dataset=StringCSV 数据文件路径,首行必须包含变量名(表头)
--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类型默认值含义
iterationsint无(必填GT 中的分裂次数;最终节点数 = 2 × iterations + 1
force_integer_codingbooleanfalsetrue时把可识别为整数的变量按整数编码(否则按 double 编码),生成更"干净"的 GT
force_binary_codingbooleantruetrue时把 0/1/unknown 变量识别为名义变量(nominal),否则按整数或 double 处理
faster_inductionbooleanfalsetrue时若候选分裂过多(超过Discriminator_Tree.MAX_SPLITS_BEFORE_RANDOMISATION,源码默认 1000)则对 DT 分裂做随机采样以加速训练
unknown_value_codingString"-1"数据集中"未知值"的表示符号,会写入全局常量Unknown_Feature_Value.S_UNKNOWN
number_bins_for_histogramsint19非名义变量的直方图分箱数,用于训练结束后计算 GT 边缘分布直方图(同时设置Histogram.NUMBER_CONTINUOUS_FEATURE_BINSMAX_NUMBER_INTEGER_FEATURE_BINS
copycat_local_generationbooleantruecopycat 归纳中 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
参数类型必填含义
-DString数据所在目录
-PString域(domain)前缀;数据文件必须位于Datasets/generate/open_policing_hartford.csv(即<目录>/<前缀>.csv
-LString生成树模型文件名(必须位于上述目录中)
-Nint要生成的样本数;生成文件与模型同目录,命名为<前缀>_GeneratedSample.csv不指定则只显示生成树结构
-UString数据集中"未知值"的表示,默认"-1"
-Fboolean是否强制整数编码,默认false
-X/-YString用于二维密度图的 x/y 变量名

Generate.go(Generate.java)的执行流程为:

  1. <目录>/<前缀>/<前缀>.csv定位并加载原始数据,构造数据域(Domain);
  2. 调用from_file(Generate.java)解析生成树模型文件:文件以@NODES/@ARCS两个区段分别描述节点与弧(边),逐行重建Generator_NodeGenerator_Arc及父子关系、叶子集合与树深度;
  3. -N指定了样本数,调用gt.generate_sample_with_density(number_ex)(当指定-X/-Y时)或gt.generate_sample(number_ex)生成样本;
  4. 输出<前缀>_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_SPLITSMAX_SPLITS_BEFORE_RANDOMISATION(默认 1000)、MAX_CARD_MODALITIES_BEFORE_RANDOMISATION(默认 10)等静态常量控制了名义变量候选分裂过多时的随机化加速策略,与--flags中的faster_induction直接对应。
  • 生成树Generator_Tree.java:节点上记录分裂特征、各分支概率(multi_p)与子节点,叶子节点构成可继续生长的集合。训练结束后compute_generator_histograms()会从 GT 采样一批样本,为每个特征计算边缘分布直方图(分箱数由--flagsnumber_bins_for_histograms控制)。
  • 调度入口Algorithm.java:simple_go()将参数封装为["@MatuErr", "1.0", "COPYCAT", iterations, copycat_local_generation]交给Boost;其中策略名COPYCATBoost.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时:

  1. Wrapper.simple_go在生成样本后调用impute_and_save(gt)(Wrapper.java);
  2. 逐行扫描原始训练样本,对含未知值的样本调用gt.impute_all_values_from_one_leaf(...)——让样本沿生成树落至叶子,用该叶子对应的分布对缺失特征进行最大似然补全(对应Generator_Tree.IMPUTATION_AT_MAXIMUM_LIKELIHOOD常量);
  3. 插补结果保存为<work_dir>/<spec_name>_imputed.csvspec_name来自--dataset文件名或--dataset_specname),并在运行摘要中打印该路径。

八、统计输出与结果文件

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_inductioncopycat_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

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

相关推荐

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询