如何用 graph.v1 的 train_schema 与 predict_schema 自定义 Rasa 执行图
【免费下载链接】rasa💬 Open source machine learning framework to automate text- and voice-based conversations: NLU, dialogue management, connect to Slack, Facebook, and more - Create chatbots and voice assistants项目地址: https://gitcode.com/GitHub_Trending/ra/rasa
Rasa 默认用default.v1配方从config.yml的pipeline/policies两个键自动搭建训练与预测图。当你需要更细粒度地控制每个执行步骤、做消融实验或接入自定义图组件时,可以改用实验性的graph.v1配方:把配置改写成train_schema与predict_schema两张显式节点图。这篇文档基于 Graph Recipe 文档 和仓库中的真实示例配置,说明如何写出这样的配置并验证它是否生效。
需要说明两点前提(均来自 Graph Recipe 文档):
graph.v1是 Rasa 3.1 引入的实验特性,功能后续可能变更或移除;- 与默认配方不同,
graph.v1不会自动补全任何缺省配置,也不支持 CLI 参数覆盖和 finetuning——配置里没写的就是没有,"what you see is what you get"。
配置文件的整体结构
graph.v1与default.v1共享recipe和language两个键,含义不变;区别在于不再有pipeline和policies,取而代之的是:
| 键 | 作用 |
|---|---|
train_schema | 训练时执行图的节点定义 |
predict_schema | 预测时执行图的节点定义 |
nlu_target | 显式指定 NLU 侧的终点节点名 |
core_target | 显式指定核心(对话管理)侧的终点节点名 |
关于 target 的两个要点:
- 如果省略
nlu_target/core_target,将分别采用默认配方的节点名run_RegexMessageHandler和select_prediction,因此你的 schema 中必须存在同名的节点; - 两张图的第一个节点依赖固定资源:训练任务用
__importer__(代表配置、训练数据等),预测任务用__message__(代表收到的消息)。core 预测节点还可以从__tracker__取对话追踪数据(见下文示例)。
target 缺失时的行为可以直接在源码 GraphV1Recipe 中确认:nlu_target对所有训练类型都是必需的;只要不是纯 NLU 训练,core_target也是必需的。两者缺失会抛出InvalidConfigException,报错信息为:
Can't find target names for NLU and/or core. Please make sure to provide 'nlu_target' (required for all training types) and 'core_target' (required if training is not just NLU) values in your config.yml file.所以只训练 NLU 时可以只写nlu_target;做端到端训练则两个都要写。
节点里各个键的含义
train_schema.nodes和predict_schema.nodes下的每个节点都用同一组键描述,Graph Recipe 文档 对它们的定义如下:
needs:声明节点需要哪些数据、数据来自哪个父节点。键是数据名,值是节点名:
needs: messages: nlu_message_converteruses:实例化该节点所用的类,必须写完整的 Python 路径:
uses: rasa.graph_components.converters.nlu_message_converter.NLUMessageConverter不要求必须用 Rasa 内置类,也可以放你自己实现的图组件(写法参见 自定义图组件文档)。
constructor_name:用于实例化组件的构造方法名,例如create(训练)或load(推理):
constructor_name: loadfn:执行该节点时调用的函数名:
fn: combine_predictions_from_kwargsconfig:传给组件的任意配置参数:
config: language: en persist: falseeager:是否在构图阶段立即实例化。通常训练时懒加载(false)、推理时立即加载(true),以避免首次预测变慢。resource:若给定,节点从该资源加载而不是从头实例化,例如用于预测阶段加载训练好的组件:
resource: name: train_RulePolicy1is_target:true时该节点不会被指纹计算剪掉(结果可能仍用缓存)。所有会训练的组件通常设为true,因为训练结果必须进模型归档,供推理使用。is_input:true的节点在指纹运行中也会始终执行,用于保证能检测到文件内容变化。
用真实示例看两张图如何对应
仓库里有两个可直接参考的graph.v1配置:
- 最小示例配置:训练与预测各只有极少的节点(MemoizationPolicy 等),适合理解结构;
- 完整示例配置:测试代码注释说明它与默认配方的默认配置在图 schema 形式下等价,适合对照你现有的
default.v1项目迁移。
以最小示例配置为例,它声明了自定义的 target,然后给出训练节点和对应的预测节点。训练侧(摘自该文件):
recipe: graph.v1 language: en core_target: custom_core_target nlu_target: custom_nlu_target train_schema: nodes: train_MemoizationPolicy0: needs: training_trackers: training_tracker_provider domain: domain_for_core_training_provider uses: rasa.core.policies.memoization.MemoizationPolicy constructor_name: create fn: train config: { } eager: false is_target: true is_input: false resource: null预测侧(同文件摘录):
predict_schema: nodes: run_MemoizationPolicy0: needs: domain: domain_provider tracker: __tracker__ rule_only_data: rule_only_data_provider uses: rasa.core.policies.memoization.MemoizationPolicy constructor_name: load fn: predict_action_probabilities config: {} eager: true is_target: false is_input: false resource: name: train_MemoizationPolicy0两段之间的对应关系是核心规则:预测节点通过resource.name指向训练节点的节点名(这里是train_MemoizationPolicy0),推理时从模型存储加载训练产物,而不是重新实例化一个未训练的组件。自定义nlu_target/core_target(custom_nlu_target、custom_core_target)则必须能在predict_schema.nodes里找到同名节点。完整的节点依赖链(nlu_training_data_provider→ 各 featurizer → 分类器,story_graph_provider→training_tracker_provider→ 各 policy)可以参考 完整示例配置。
训练与验证
配置写好后,在包含config.yml的项目目录运行rasa train。针对graph.v1,可以用以下三类可核对的现象判断配置是否符合预期:
- target 缺失:如上所述,缺少
nlu_target或非 NLU 训练时缺少core_target,训练在构建图之前就会以InvalidConfigException报错,并按报错信息提示补上对应键。 - 节点引用了不存在的模块:
uses中写了无法解析的完整路径时,GraphSchema.from_dict会抛出GraphSchemaException。仓库测试 test_graph_recipe.py 就是用不存在的rasa.core.policies.ted_policy.TEDPolicy1000来验证这一报错的。 - 误传 CLI 参数或开启 finetuning:GraphV1Recipe 源码 显示,只要检测到 CLI 参数或 finetuning,就会发出
UserWarning:
Unlike the Default Recipe, Graph Recipe does not utilize CLI parameters or finetuning and these configurations will be ignored. Add configuration to the recipe itself if you want them to be used.也就是说,想改epochs、max_history这类参数,不要加在命令行上,而是写进对应节点的config(如 完整示例配置 中train_TEDPolicy3的config里就写了max_history: 5和epochs: 100)。
限制与注意事项
- 该配方标注为实验特性(New in 3.1),接口可能变化;生产项目是否切换需要自行评估。
- 它不利用 CLI 参数和 finetuning,所有参数必须显式落在 schema 节点的
config里。 default.v1在配置节缺失时会自动写回推荐配置;graph.v1完全没有这个行为,缺了就是缺了。- 自定义图组件接入
graph.v1时,要求组件实现GraphComponent接口、使用类型注解(不允许 forward references),并在uses中写完整模块路径,详见 自定义图组件文档。
【免费下载链接】rasa💬 Open source machine learning framework to automate text- and voice-based conversations: NLU, dialogue management, connect to Slack, Facebook, and more - Create chatbots and voice assistants项目地址: https://gitcode.com/GitHub_Trending/ra/rasa
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考