十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Generative Trees(生成树)代码实战指南:基于 Copycat 对抗训练的生成模型与采样(ICML‘22 配套代码)

Generative Trees(生成树)代码实战指南:基于 Copycat 对抗训练的生成模型与采样(ICML‘22 配套代码) 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本指南围绕google-research仓库中 generative_trees 目录的官方 README 与配套 Java 源码展开系统讲解 ICML22 论文Generative Trees: Adversarial and Copycat作者 Richard Nock 与 Mathieu Guillame-Bert的配套实现如何用copycat 方法训练一棵生成树Generative TreeGT作为数据生成模型以及如何从预训练生成树中批量采样新样本、绘制二维密度图、完成缺失值插补。读完本文你将掌握编译与运行示例脚本的完整流程、Wrapper与Generate两个入口的全部命令行参数及含义、copycat 训练循环的源码级运行机制以及生成树的保存格式与统计输出结构能够直接在自己的数据集上复现该论文的方法。一、项目背景与核心思想generative_trees是 ICML22 论文Generative Trees: Adversarial and Copycat的官方配套代码BibTeX 见 README.md。其核心方法名为copycat训练过程不直接优化生成树本身而是让一个判别树Discriminator TreeDT持续对真实数据 vs 生成数据做区分生成树Generator TreeGT则模仿判别树最新一次成功分裂的结构二者交替进化从而逐步逼近真实数据分布。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 \ --datasetiris.csv \ --work_dirworking_dir \ --num_samples1000 \ --output_samplesworking_dir/generated.csv \ --output_statsworking_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.shautomates the process, can be edited easily不过在本文所基于的仓库快照中该脚本并未包含缺失值插补能力可以直接通过Wrapper的--impute_missingtrue参数触发见第六节。三、两大核心入口Wrapper与GenerateREADME 明确指出本项目由两个关键类构成二者职责分离对应源码分别位于 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。四、Wrappercopycat 训练生成树参数全解Wrapper的完整帮助文本内嵌于源码的help()方法中Wrapper.javaREADME 与脚本只展示了它的一个子集。以下按源码整理全部参数。4.1 基础参数参数类型必填含义--datasetString是CSV 数据文件路径首行必须包含变量名表头--dataset_specJSON 字符串否数据集规格说明含name/path/label/task四个字段见 4.2--num_samplesint是要生成的样本数量--work_dirString是生成树模型与密度图文件的保存目录--output_samplesString是生成样本的输出文件名--output_statsString是运行统计文件运行时长、GT 边缘直方图、GT 树节点统计等--x--yString否可选用于保存二维密度图的两个变量名输出格式为(x, y, density_value_at_(x,y))--flagsJSON 字符串否训练超参数见 4.3--impute_missingboolean否若为true用生成树对训练数据中的缺失值进行插补源码中的参数解析位于Wrapper.fit_varsWrapper.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 1force_integer_codingbooleanfalse为true时把可识别为整数的变量按整数编码否则按 double 编码生成更干净的 GTforce_binary_codingbooleantrue为true时把 0/1/unknown 变量识别为名义变量nominal否则按整数或 double 处理faster_inductionbooleanfalse为true时若候选分裂过多超过Discriminator_Tree.MAX_SPLITS_BEFORE_RANDOMISATION源码默认 1000则对 DT 分裂做随机采样以加速训练unknown_value_codingString-1数据集中未知值的表示符号会写入全局常量Unknown_Feature_Value.S_UNKNOWNnumber_bins_for_histogramsint19非名义变量的直方图分箱数用于训练结束后计算 GT 边缘分布直方图同时设置Histogram.NUMBER_CONTINUOUS_FEATURE_BINS与MAX_NUMBER_INTEGER_FEATURE_BINScopycat_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_missingtrue开启插补java Wrapper --dataset${ANYDIR}/Datasets/iris/iris.csv \ --dataset_spec{name: iris, path: ${ANYDIR}/Datasets/iris/iris.csv, label: class, task: BINARY_CLASSIFICATION} \ --num_samples10000 \ --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 \ --xSepal.Length --ySepal.Width \ --flags{iterations : 10, force_integer_coding : true, force_binary_coding : true, faster_induction : true, unknown_value_coding : ?, number_bins_for_histograms : 11} \ --impute_missingtrue注意源码不允许--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.goGenerate.java的执行流程为按目录/前缀/前缀.csv定位并加载原始数据构造数据域Domain调用from_fileGenerate.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) 创建生成树 GTGenerator_Tree并初始化根节点 2) 创建判别树 DTDiscriminator_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.javasimple_go()将参数封装为[MatuErr, 1.0, COPYCAT, iterations, copycat_local_generation]交给Boost其中策略名COPYCAT由Boost.KEY_NAME白名单校验Boost.java。数据域Domain.java负责加载特征与样本Dataset、计算域直方图并挂载内存监控器MemoryMonitor。Wrapper.simple_goWrapper.java按固定流水线串联加载数据 → 学习 GT → 计算 GT 边缘直方图 → 保存 GT 到work_dir/generator_name→ 生成样本 →可选缺失值插补 → 保存样本 →可选保存二维密度图 → 保存统计文件每一步都打印耗时毫秒。七、缺失值插补--impute_missingREADME 明确将缺失值插补列为该代码的关键能力之一。当训练数据含有未知值默认编码为-1可用unknown_value_coding修改且传入--impute_missingtrue时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.csvspec_name来自--dataset文件名或--dataset_spec的name并在运行摘要中打印该路径。八、统计输出与结果文件Wrapper每次运行会产出多类文件--output_stats指定主统计文件路径主统计文件JSON 格式Wrapper.java包含running_time_seconds训练生成总时长、gt_number_nodesGT 节点数、gt_depthGT 深度、running_time_gt_training_plus_exemple_generation以及开启插补时的running_time_gt_training_plus_imputation。附加统计文件output_stats_more.txtL429-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×iterations1faster_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.javacopycat 训练循环核心实现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-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/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),仅供参考
返回列表