首页/新闻资讯/正文详情

NNI Trial 开发指南:如何编写运行在 NNI 上的 Trial(Tuner 参数获取与结果上报全流程)

发布时间:2026/9/23 5:14:02 来源:云帆数科 栏目:资讯中心
NNI Trial 开发指南:如何编写运行在 NNI 上的 Trial(Tuner 参数获取与结果上报全流程)
NNI Trial 开发指南如何编写运行在 NNI 上的 TrialTuner 参数获取与结果上报全流程【免费下载链接】nniAn open source AutoML toolkit for automate machine learning lifecycle, including feature engineering, neural architecture search, model compression and hyper-parameter tuning.项目地址: https://gitcode.com/gh_mirrors/nn/nniTrial 是 NNINeural Network Intelligence自动机器学习框架中执行训练任务的最小单元它从 Tuner 接收超参数/网络结构配置将训练过程中的中间结果发送给 Assessor并将最终结果回传给 Tuner。本文基于仓库中的 examples/trials/README.md 及对应的 mnist-keras 完整示例系统讲解如何把一段普通的机器学习代码改造成可运行在 NNI 上的 Trial。读完本文你将掌握 Trial 与 NNI 框架交互的四个核心步骤准备可运行的原始代码、通过nni.get_next_parameter()获取配置、通过nni.report_intermediate_result()上报中间结果、通过nni.report_final_result()上报最终结果并能独立搭建一个可被 Tuner 调参、可被 Assessor 提前终止的完整实验。Trial 在 NNI 中的角色与职责在 NNI 的自动机器学习生命周期中Trial 是承上启下的执行节点。仓库文档开篇即点明其定位Trial receive the hyper-parameter/architecture configure from Tuner, and send intermediate result to Assessor and final result to Tuner.即Trial 的输入是 Tuner 根据搜索空间search space采样出的参数输出是训练过程中产生的中间指标与训练结束后的最终指标。这一交互在源码层面有明确对应。在 nni/trial.py 中Trial 侧暴露的全部 API 为get_next_parameter()/get_next_parameters()获取 Tuner 生成的超参数get_current_parameter(tag)读取当前参数可带字段名report_intermediate_result(metric)上报中间结果report_final_result(metric)上报最终结果get_experiment_id()/get_trial_id()/get_sequence_id()获取实验 ID、Trial ID 与序号。因此编写一个运行在 NNI 上的 Trial 通常只需要四步先有一份能在本地跑通的机器学习代码再按本文所述插入 NNI 的 API 调用。第一步准备一份可运行的原始 TrialTrial 的本质是一段能够在本地直接运行的机器学习代码NNI 不要求对模型结构做任何特殊改写。文档以mnist-keras.py为例展示了一份原始代码使用 Keras 构建一个简单的卷积网络两层Conv2DMaxPooling2DFlatten 两层Dense加载 MNIST 数据集训练并评估。其关键点在于import argparse import logging import keras import numpy as np from keras import backend as K from keras.datasets import mnist from keras.layers import Conv2D, Dense, Flatten, MaxPooling2D from keras.models import Sequential K.set_image_data_format(channels_last) H, W 28, 28 NUM_CLASSES 10 def create_mnist_model(hyper_params, input_shape(H, W, 1), num_classesNUM_CLASSES): layers [ Conv2D(32, kernel_size(3, 3), activationrelu, input_shapeinput_shape), Conv2D(64, (3, 3), activationrelu), MaxPooling2D(pool_size(2, 2)), Flatten(), Dense(100, activationrelu), Dense(num_classes, activationsoftmax) ] model Sequential(layers) if hyper_params[optimizer] Adam: optimizer keras.optimizers.Adam(lrhyper_params[learning_rate]) else: optimizer keras.optimizers.SGD(lrhyper_params[learning_rate], momentum0.9) model.compile(losskeras.losses.categorical_crossentropy, optimizeroptimizer, metrics[accuracy]) return model def load_mnist_data(args): (x_train, y_train), (x_test, y_test) mnist.load_data() x_train (np.expand_dims(x_train, -1).astype(float) / 255.)[:args.num_train] x_test (np.expand_dims(x_test, -1).astype(float) / 255.)[:args.num_test] y_train keras.utils.to_categorical(y_train, NUM_CLASSES)[:args.num_train] y_test keras.utils.to_categorical(y_test, NUM_CLASSES)[:args.num_test] return x_train, y_train, x_test, y_test def train(args, params): x_train, y_train, x_test, y_test load_mnist_data(args) model create_mnist_model(params) model.fit(x_train, y_train, batch_sizeargs.batch_size, epochsargs.epochs, verbose1, validation_data(x_test, y_test), callbacks[SendMetrics()]) _, acc model.evaluate(x_test, y_test, verbose0) def generate_default_params(): return { optimizer: Adam, learning_rate: 0.001 } if __name__ __main__: PARSER argparse.ArgumentParser() PARSER.add_argument(--batch_size, typeint, default200, helpbatch size, requiredFalse) PARSER.add_argument(--epochs, typeint, default10, helpTrain epochs, requiredFalse) PARSER.add_argument(--num_train, typeint, default1000, helpNumber of train samples to be used, maximum 60000, requiredFalse) PARSER.add_argument(--num_test, typeint, default1000, helpNumber of test samples to be used, maximum 10000, requiredFalse) ARGS, UNKNOWN PARSER.parse_known_args() PARAMS generate_default_params() train(ARGS, PARAMS)这段代码没有任何 NNI 依赖可直接在本地运行用于验证代码正确性。仓库中的实际示例 examples/trials/mnist-keras/mnist-keras.py 在此基础上进一步做了两处工程化增强一是通过os.environ[NNI_OUTPUT_DIR]把 TensorBoard 日志目录指向 NNI 分配的输出目录二是将 MNIST 数据集缓存到NNI_OUTPUT_DIR下并在使用后删除避免多 Trial 并发时相互污染。注意其中的SendMetrics回调在原始代码中是空实现pass这正是后续要接入 NNI 的位置。第二步从 Tuner 获取超参数配置改造的第一处关键动作是引入nni模块并调用nni.get_next_parameter()。文档特别提醒关注示例中的第 10、24、25 行即导入语句、调用获取参数的语句以及用返回结果更新默认参数字典的语句import nni # 第 10 行导入 nni if __name__ __main__: PARSER argparse.ArgumentParser() ... ARGS, UNKNOWN PARSER.parse_known_args() PARAMS generate_default_params() RECEIVED_PARAMS nni.get_next_parameter() # 获取 Tuner 采样出的参数 PARAMS.update(RECEIVED_PARAMS) # 用 Tuner 参数覆盖默认值 train(ARGS, PARAMS)这一模式非常关键先用generate_default_params()提供一份可独立运行的默认参数再用 Tuner 返回的参数update覆盖默认值。这样既保证无 Tuner 参数时也能跑通又保证参数真正来自 Tuner。从源码看nni.get_next_parameter()在 nni/trial.py 中通过get_default_trial_command_channel().receive_parameter()从 NNI manager 接收参数记录并返回其中的parameters字段。其 docstring 给出了典型的返回形态若搜索空间为{activation: {_type: choice, _value: [relu, tanh, sigmoid]}, learning_rate: {_type: loguniform, _value: [0.0001, 0.1]}}则返回值形如{activation: relu, learning_rate: 0.02}。同时源码明确要求每个 Trial 应且只应调用一次该函数否则行为未定义见 docstring这是编写 Trial 时必须遵守的约定。值得注意的是Trial 代码在脱离 NNI 环境独立运行时receive_parameter()会走 nni/runtime/trial_command_channel/standalone.py 中的StandaloneTrialCommandChannel返回空参数集{}并发出运行时警告从而保证同一份代码既能被 NNI 调度、也能本地调试。第三步上报中间结果给 Assessor中间结果intermediate result是训练过程中周期产生的指标典型的就是每个 epoch 的 accuracy 或 loss。它的接收方是 Assessor如早停算法NNI 据此决定是否提前终止表现不佳的 Trial从而节省计算资源。在 Keras 中最自然的接入点是回调Callback。文档改造了SendMetrics回调在on_epoch_end中调用nni.report_intermediate_result(logs)class SendMetrics(keras.callbacks.Callback): def on_epoch_end(self, epoch, logs{}): nni.report_intermediate_result(logs)在model.fit(...)时把该回调传入callbacks[SendMetrics()]即可。仓库中的真实示例 examples/trials/mnist-keras/mnist-keras.py 做了一个值得借鉴的健壮性处理Keras 不同版本中验证集准确率的日志键名不一致TensorFlow 2.0 文档称其为val_acc实际为val_accuracy因此它同时兼容两种情况if val_acc in logs: nni.report_intermediate_result(logs[val_acc]) else: nni.report_intermediate_result(logs[val_accuracy])这提示了一个通用原则上报的指标值应确保是框架期望的数值形态。第四步上报最终结果给 Tuner训练结束后Trial 需要把最终指标发给 Tuner供其更新代理模型、指导下一轮采样。改造方式同样简单在model.evaluate得到准确率后调用nni.report_final_result(acc)def train(args, params): x_train, y_train, x_test, y_test load_mnist_data(args) model create_mnist_model(params) model.fit(x_train, y_train, batch_sizeargs.batch_size, epochsargs.epochs, verbose1, validation_data(x_test, y_test), callbacks[SendMetrics()]) _, acc model.evaluate(x_test, y_test, verbose0) nni.report_final_result(acc)关于两个上报 API 的取值约定nni/trial.py 的源码 docstring 给出了权威说明metric可以是float也可以是包含default键值为 float的字典若传字典Tuner 使用metric[default]其余字段可在 Web 门户中可视化report_intermediate_result内部以typePERIODICAL发送指标并为每次上报递增序列号report_final_result则以typeFINAL、sequence0发送两个 API 都断言了nni.get_next_parameter()必须在此之前被调用过否则在 NNI 平台上会直接断言失败这是 Trial 代码必须遵循的调用顺序。此外report_intermediate_result与report_final_result支持同时传更多自定义字段用于 WebUI 展示也可以调用 nni/trial.py 中的get_trial_id()、get_sequence_id()等在日志或上报中标记当前 Trial 身份。完整示例从零到可运行的 NNI Trial将以上四步合并即得到一份完整、可直接运行的 NNI Trial对应 examples/trials/mnist-keras/mnist-keras.pyimport argparse import logging import keras import numpy as np from keras import backend as K from keras.datasets import mnist from keras.layers import Conv2D, Dense, Flatten, MaxPooling2D from keras.models import Sequential import nni LOG logging.getLogger(mnist_keras) K.set_image_data_format(channels_last) H, W 28, 28 NUM_CLASSES 10 def create_mnist_model(hyper_params, input_shape(H, W, 1), num_classesNUM_CLASSES): layers [ Conv2D(32, kernel_size(3, 3), activationrelu, input_shapeinput_shape), Conv2D(64, (3, 3), activationrelu), MaxPooling2D(pool_size(2, 2)), Flatten(), Dense(100, activationrelu), Dense(num_classes, activationsoftmax) ] model Sequential(layers) if hyper_params[optimizer] Adam: optimizer keras.optimizers.Adam(lrhyper_params[learning_rate]) else: optimizer keras.optimizers.SGD(lrhyper_params[learning_rate], momentum0.9) model.compile(losskeras.losses.categorical_crossentropy, optimizeroptimizer, metrics[accuracy]) return model def load_mnist_data(args): (x_train, y_train), (x_test, y_test) mnist.load_data() x_train (np.expand_dims(x_train, -1).astype(float) / 255.)[:args.num_train] x_test (np.expand_dims(x_test, -1).astype(float) / 255.)[:args.num_test] y_train keras.utils.to_categorical(y_train, NUM_CLASSES)[:args.num_train] y_test keras.utils.to_categorical(y_test, NUM_CLASSES)[:args.num_test] return x_train, y_train, x_test, y_test class SendMetrics(keras.callbacks.Callback): def on_epoch_end(self, epoch, logs{}): LOG.debug(logs) nni.report_intermediate_result(logs) def train(args, params): x_train, y_train, x_test, y_test load_mnist_data(args) model create_mnist_model(params) model.fit(x_train, y_train, batch_sizeargs.batch_size, epochsargs.epochs, verbose1, validation_data(x_test, y_test), callbacks[SendMetrics()]) _, acc model.evaluate(x_test, y_test, verbose0) LOG.debug(Final result is: %d, acc) nni.report_final_result(acc) def generate_default_params(): return { optimizer: Adam, learning_rate: 0.001 } if __name__ __main__: PARSER argparse.ArgumentParser() PARSER.add_argument(--batch_size, typeint, default200, helpbatch size, requiredFalse) PARSER.add_argument(--epochs, typeint, default10, helpTrain epochs, requiredFalse) PARSER.add_argument(--num_train, typeint, default1000, helpNumber of train samples to be used, maximum 60000, requiredFalse) PARSER.add_argument(--num_test, typeint, default1000, helpNumber of test samples to be used, maximum 10000, requiredFalse) ARGS, UNKNOWN PARSER.parse_known_args() try: RECEIVED_PARAMS nni.get_next_parameter() LOG.debug(RECEIVED_PARAMS) PARAMS generate_default_params() PARAMS.update(RECEIVED_PARAMS) train(ARGS, PARAMS) except Exception as e: LOG.exception(e) raise配套的搜索空间与实验配置要让上述 Trial 真正在 NNI 中参与超参数搜索还需要两个配套文件搜索空间描述与实验配置。搜索空间 examples/trials/mnist-keras/search_space.json 定义了 Tuner 可在哪些参数上采样{ optimizer:{_type:choice,_value:[Adam, SGD]}, learning_rate:{_type:choice,_value:[0.0001, 0.001, 0.002, 0.005, 0.01]} }optimizer在Adam与SGD之间选择learning_rate在 5 个离散值中选择与 Trial 代码中create_mnist_model(hyper_params)读取的键一一对应——搜索空间的 key 必须与hyper_params中的字段名完全一致Tuner 采样结果才能被PARAMS.update(RECEIVED_PARAMS)正确覆盖。实验配置 examples/trials/mnist-keras/config.yml 则声明了实验的运行方式authorName: default experimentName: example_mnist-keras trialConcurrency: 1 maxExecDuration: 1h maxTrialNum: 10 #choice: local, remote, pai trainingServicePlatform: local searchSpacePath: search_space.json #choice: true, false useAnnotation: false tuner: #choice: TPE, Random, Anneal, Evolution, BatchTuner, MetisTuner #SMAC (SMAC should be installed through nnictl) builtinTunerName: TPE classArgs: #choice: maximize, minimize optimize_mode: maximize trial: command: python3 mnist-keras.py codeDir: . gpuNum: 0关键字段含义如下trainingServicePlatform: local在本地运行 Trial可选值包括 local、remote、pai以及仓库中对应的 config_pai.yml 等变体searchSpacePath指向搜索空间文件tuner.builtinTunerName: TPE使用内置 TPE 算法可选 TPE、Random、Anneal、Evolution、BatchTuner、MetisTuner 等SMAC 需通过 nnictl 另行安装tuner.classArgs.optimize_mode: maximizeTuner 按最大化方向优化因为 Trial 上报的是 accuracy若上报的是 loss 则应设为minimizetrial.command与trial.codeDir声明如何启动 Trial 及其代码目录因此 Trial 的入口脚本名、相对路径必须与之一致trial.gpuNum每个 Trial 分配的 GPU 数量本地调试可设为 0。配置就绪后即可用nnictl create --config config.yml创建实验NNI 会启动 nni manager 并按maxTrialNum/maxExecDuration调度多个 Trial 并发执行。从源码理解数据流Trial 与 NNI 框架如何通信以上 API 的背后是一条清晰的通信链路。从源码结构看Trial 侧所有上报与接收操作最终都收敛到命令通道Command Channel抽象抽象基类 nni/runtime/trial_command_channel/base.py 定义了receive_parameter()从 NNI manager 接收参数记录与send_metric()发送指标类型限定为PERIODICAL或FINAL最终指标序号必须为 0两个接口不同运行环境下有不同实现如 standalone.py脱离 NNI 运行时使用、local_legacy.py本地传统模式、v3.pyNNI v3 新式通道等nni/trial.py 中的get_next_parameter、report_intermediate_result、report_final_result均为对这些通道实现的薄封装并通过trial_env_vars定义于 nni/runtime/env_vars.py包括NNI_EXP_ID、NNI_TRIAL_JOB_ID、NNI_TRIAL_SEQ_ID、NNI_OUTPUT_DIR等获得当前实验与 Trial 的上下文。可以推断整个数据流为Tuner 依据搜索空间采样 → 参数记录经命令通道下发至 Trial → Trial 调用get_next_parameter()获取 → 训练过程中按周期调用report_intermediate_result()供 Assessor 决策 → 训练结束调用report_final_result()将最终指标传回 Tuner 更新其模型形成闭环。这一闭环正是 NNI 超参搜索能够越搜越好的底层机制。小结与编写规范速查编写一个运行在 NNI 上的 Trial 可归纳为四个步骤准备可运行的原始代码 → 用nni.get_next_parameter()获取 Tuner 参数并合并进默认参数 → 用nni.report_intermediate_result()周期上报中间结果 → 用nni.report_final_result()上报最终结果。在此之上遵循以下规范可以显著减少踩坑每个 Trial 只调用一次get_next_parameter()上报指标优先使用 float或包含 float 类型default键的字典report_intermediate_result/report_final_result必须在get_next_parameter()之后调用搜索空间的字段名与 Trial 代码读取的超参键保持一致用try/except包裹主逻辑并记录异常便于在 NNI Web 门户中定位失败原因利用NNI_OUTPUT_DIR等环境变量管理输出文件日志、模型、TensorBoard 事件保证多 Trial 并发互不干扰。更完整的示例PyTorch、TensorFlow 等不同框架可继续阅读 examples/trials/README.md 同目录下的 examples/trials/ 其他示例以及 docs/source/hpo/quickstart.rst 等官方文档。【免费下载链接】nniAn open source AutoML toolkit for automate machine learning lifecycle, including feature engineering, neural architecture search, model compression and hyper-parameter tuning.项目地址: https://gitcode.com/gh_mirrors/nn/nni创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

相关推荐

Java异步编程完全指南:从线程池、CompletableFuture到虚拟线程
Java异步编程完全指南:从线程池、CompletableFuture到虚拟线程

1. 为什么Java的异步编程绕不开线程模型1.1 一个真实接口延迟案例:从同步瓶颈说起我在维护一个电商后台项目时,遇到过这么一个问题:订单详情接口的响应时间从平均200毫秒慢慢涨到了1.2秒,而且是线性恶化的趋势。业务量没涨那么多&… · 2026/9/23 5:13:56

swagger-codegen 生成的 C Pet 模型深度解析:以 SwaggerClientWithPropertyChanged 为例
swagger-codegen 生成的 C Pet 模型深度解析:以 SwaggerClientWithPropertyChanged 为例

swagger-codegen 生成的 C# Pet 模型深度解析:以 SwaggerClientWithPropertyChanged 为例 【免费下载链接】swagger-codegen swagger-codegen contains a template-driven engine to generate documentation, API clients and server stubs in different languages … · 2026/9/23 5:13:56

redux-form v5 → v6 迁移完全指南:控制反转、Field 组件与状态结构重塑
redux-form v5 → v6 迁移完全指南:控制反转、Field 组件与状态结构重塑

前端UI组件 【免费下载链接】redux-form A Higher Order Component using react-redux to keep form state in a Redux store 项目地址: https://gitcode.com/gh_mirrors/re/redux-form 点击查看 免费下载 导读 v6 是 redux-form 历史上一次彻底重写(c… · 2026/9/23 5:13:56

Python打包成exe避坑指南
Python打包成exe避坑指南

Python 打包成 exe 避坑指南(PyInstaller 实战)写 Python 的第十五年,我被问得最多的问题之一就是:"我写了个脚本,怎么发给不会装 Python 的同事用?"答案就是 PyInstaller——把 .py 打包成 .exe… · 2026/9/23 7:45:16

告别复制即崩:先锋网站开发避坑与速查手册实战指南
告别复制即崩:先锋网站开发避坑与速查手册实战指南

告别复制即崩:先锋网站开发避坑与速查手册实战指南 刚接手一个嵌入式项目的前端展示页,也就是俗称的“先锋网站”,直接从网上扒了一套开源模板。代码贴进去,本地 npm run dev… · 2026/9/23 7:45:04

Octopress静态博客实战:从WordPress迁移到高效写作的完整指南
Octopress静态博客实战:从WordPress迁移到高效写作的完整指南

凌晨一点半,我刚把一篇三千字的长文从剪贴板粘进 WordPress 后台,正要点击发布,页面忽然变成一片空白——数据库连接错误。那一瞬间我真想把电脑从窗口扔出去。三年积累的两百多篇帖子、几百条评论,全躺在 MySQL 里,而… · 2026/9/23 7:45:04

13邀避坑指南:跨省转介与报考门槛深度解析
13邀避坑指南:跨省转介与报考门槛深度解析

13邀避坑指南:跨省转介与报考门槛深度解析 配置环境就卡半天?别急,这次咱们聊点更“硬核”的。很多刚入行或者准备转型的朋友,一提到 13邀… · 2026/9/23 7:44:57

理想汽车:从增程式到纯电转型的战略挑战
理想汽车:从增程式到纯电转型的战略挑战

1. 理想汽车的崛起与困境:从增程式优等生到纯电转型的阵痛理想汽车最初凭借精准的市场定位和增程式技术路线,在中国新能源汽车市场迅速崛起。2019年推出的理想ONE以"为家庭用户打造"为核心理念,成功抓住了中国家庭对空间、舒适性和… · 2026/9/23 7:44:57

投资战略:识别不变性与三维时间框架
投资战略:识别不变性与三维时间框架

1. 投资战略的本质解析"在变化中寻找不变"这句话听起来像哲学命题,但恰恰是顶级投资者每天都在践行的生存法则。我从业十五年,见过太多人沉迷于追逐市场热点、政策风向和技术形态,最终却被市场反复收割。真正的战略思维&#xff0c… · 2026/9/23 7:44:57

3招搞定手机怎么下载微信面试难题实战项目解析
3招搞定手机怎么下载微信面试难题实战项目解析

3招搞定手机怎么下载微信面试难题实战项目解析 面试被问“手机怎么下载微信”背后的原理,90%的人答不上来。别笑,这看似弱智的问题,实则是考察你对移动应用分发机制、安全校验及网络协议理解的试金石。我带过不少校招新人,他们背了八股文,却连一个A… · 2026/9/23 0:00:03

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型
你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型 面试被问“高并发下如何保证消息不丢失”,你张口就是“用Redis”,结果面试官追问“如果Redis宕机了怎么办”,你瞬间卡壳。这种场景太常见了,很多新手在背八股文时,只记住了技术名词… · 2026/9/23 0:00:29

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧
Win7无线热点配置工具源码解析:解决API失效的3个实战技巧

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧 Win7无线热点配置工具在Win10/11上跑不动?不是你的问题,是版本升级后 API 全变了。很多老项目里的 netsh wlan… · 2026/9/23 0:00:36

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码