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

NNI TensorFlow HPO 快速入门:用 TPE 自动调优 Keras MNIST 模型的完整实战指南

发布时间:2026/9/23 10:39:58 来源:云帆数科 栏目:资讯中心
NNI TensorFlow HPO 快速入门:用 TPE 自动调优 Keras MNIST 模型的完整实战指南
NNI TensorFlow HPO 快速入门用 TPE 自动调优 Keras MNIST 模型的完整实战指南【免费下载链接】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导读本文基于 NNI 官方教程HPO Quickstart with TensorFlow仓库路径examples/tutorials/hpo_quickstart_tensorflow/完整演示如何把官方 TensorFlow 快速入门中的 Keras MNIST 分类模型改造成可自动调优的 Trial并通过 NNI 的 Python API 配置本地实验、使用 TPE 调优器搜索 4 个超参数。读完本文你将掌握模型接入 NNI 所需的 3 个 API 调用、搜索空间的 3 种核心类型choice / uniform / loguniform、nni.experiment.Experiment的完整配置流程以及实验启动、查看与停止的实战技巧。教程共分 4 步为自动调优改造模型定义超参数搜索空间配置实验Experiment运行实验。说明本文所述的“官方 TensorFlow 快速入门”对应 examples/tutorials/hpo_quickstart_tensorflow/model.py 中注释引用的 TensorFlow 官方 beginner 教程。教程源代码为两段式“文学编程”脚本主控脚本 examples/tutorials/hpo_quickstart_tensorflow/main.py 与模型脚本 examples/tutorials/hpo_quickstart_tensorflow/model.py其渲染后的文档位于 docs/source/tutorials/hpo_quickstart_tensorflow/main.rst。Step 1准备模型为自动调优改造模型改造的第一步是把要被调优的模型放到一个独立的脚本中。原因是这个脚本会在实验中被并发地评估很多次将来甚至可能被提交到分布式训练平台上运行。NNI 每次评估一组超参数称为一次trial因此这个模型脚本也被称为trial code。本教程中模型定义在 examples/tutorials/hpo_quickstart_tensorflow/model.py 中。这段代码本质上就是 TensorFlow 官方 MNIST 快速入门模型只是在原版基础上额外增加了 3 个 NNI API 调用nni.get_next_parameter()向调优算法获取本次要评估的超参数nni.report_intermediate_result()上报每个 epoch 的准确率等中间指标nni.report_final_result()上报最终准确率。模型代码逐段解读完整的模型代码如下节选自 model.pyimport nni import tensorflow as tf # Hyperparameters to be tuned params { dense_units: 128, activation_type: relu, dropout_rate: 0.2, learning_rate: 0.001, } # Get optimized hyperparameters optimized_params nni.get_next_parameter() params.update(optimized_params) print(params) # Load dataset mnist tf.keras.datasets.mnist (x_train, y_train), (x_test, y_test) mnist.load_data() x_train, x_test x_train / 255.0, x_test / 255.0 # Build model with hyperparameters model tf.keras.models.Sequential([ tf.keras.layers.Flatten(input_shape(28, 28)), tf.keras.layers.Dense(params[dense_units], activationparams[activation_type]), tf.keras.layers.Dropout(params[dropout_rate]), tf.keras.layers.Dense(10) ]) adam tf.keras.optimizers.Adam(learning_rateparams[learning_rate]) loss_fn tf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue) model.compile(optimizeradam, lossloss_fn, metrics[accuracy]) # (Optional) Report intermediate results callback tf.keras.callbacks.LambdaCallback( on_epoch_end lambda epoch, logs: nni.report_intermediate_result(logs[accuracy]) ) # Train and evaluate the model model.fit(x_train, y_train, epochs5, verbose2, callbacks[callback]) loss, accuracy model.evaluate(x_test, y_test, verbose2) # Report final result nni.report_final_result(accuracy)三个关键设计点① 默认参数 覆盖更新的双模式设计。脚本首先用一组默认参数dense_units128、activation_typerelu、dropout_rate0.2、learning_rate0.001初始化params随后用optimized_params nni.get_next_parameter()获取调优算法生成的参数并params.update(optimized_params)覆盖默认值。这种写法让脚本具备两种运行模式直接运行nni.get_next_parameter()是无操作no-op返回空字典{}此时模型与原始 TensorFlow 快速入门行为完全一致——因此教程建议先直接运行一次脚本验证环境是否就绪在 NNI 实验内运行调优算法会把超参数通过命令通道下发到 trialget_next_parameter()返回实际采样值。从源码看nni/trial.py 中的get_next_parameter()通过get_default_trial_command_channel().receive_parameter()接收参数而 trial 的环境变量NNI_PLATFORM为空即直接运行时各上报函数会自动退化为无操作这正是“双模式”实现的底层依据。源码注释还强调每个 trial 中get_next_parameter()应当且只能调用一次否则行为未定义。② 用 KerasLambdaCallback上报逐 epoch 指标。report_intermediate_result被挂在on_epoch_end回调里把每个 epoch 的logs[accuracy]上报给 NNI。这些中间指标有两个用途在 Web 门户的 Trial 详情页画出学习曲线以及配合 NNI Assessor早停评估器 实现提前终止——当某个 trial 明显没有希望时及时掐掉节省计算资源。这一部分是可选的跳过它实验也能正常运行但会失去上述能力。从 nni/trial.py 的实现可以看到report_intermediate_result会以typePERIODICAL的消息类型发送指标并自动维护递增的_intermediate_seq序号NNI 据此在 Web 门户中按顺序绘制学习曲线。③ 最终结果决定调优方向。训练结束后model.evaluate得到测试集准确率通过nni.report_final_result(accuracy)上报。从 nni/trial.py 源码可见该函数发送的指标同时携带parameter_id与trial_job_idNNI 调优器据此把“准确率”与“这组超参数”绑定起来从而在下一次采样时给出更优的候选。上报的指标既可以是浮点数也可以是字典此时必须以metric[default]存放主指标其余键用于 Web 门户可视化。Step 2定义搜索空间模型代码里准备了 4 个待调超参数dense_units、activation_type、dropout_rate、learning_rate。要在合理的范围内让调优算法采样就需要为它们定义搜索空间search space。假设我们具备如下先验知识超参数取值范围 / 分布对应 NNI 类型dense_units64、128、256 之一choiceactivation_typerelu、tanh、swish或Nonechoicedropout_rate0.5 到 0.9 之间的浮点数uniformlearning_rate0.0001 到 0.1 之间的浮点数服从指数分布loguniform在 NNI 中dense_units和activation_type的取值空间叫choice从给定列表中选一个dropout_rate的取值空间叫uniform区间内均匀采样learning_rate的取值空间叫loguniform指数分布/对数均匀采样适合跨越多个数量级的数值如学习率。细心的话你会发现这些命名来源于numpy.random的对应分布函数。搜索空间的定义如下即 main.py 中的代码search_space { dense_units: {_type: choice, _value: [64, 128, 256]}, activation_type: {_type: choice, _value: [relu, tanh, swish, None]}, dropout_rate: {_type: uniform, _value: [0.5, 0.9]}, learning_rate: {_type: loguniform, _value: [0.0001, 0.1]}, }每个键值对由_type采样类型和_value取值范围组成其语义与 NumPy 对应分布一致。NNI 支持的全部搜索空间类型与详细规范参见搜索空间参考文档。Step 3配置实验ExperimentNNI 用experiment实验来管理整个 HPO 过程experiment config实验配置决定了“如何训练模型”与“如何探索搜索空间”。本教程使用local 模式实验——模型就在本机训练不依赖任何专门的训练平台。from nni.experiment import Experiment experiment Experiment(local)接下来依次配置实验的各个组成部分。配置 Trial 代码在 NNI 中每组超参数的一次评估称为一个trial因此模型脚本被称为trial codeexperiment.config.trial_command python model.py experiment.config.trial_code_directory .trial_commandtrial 进程需要执行的命令trial_code_directorytrial 代码所在的目录。当它是相对路径时相对于当前工作目录。路径相关注意事项如果你在其他路径下运行main.py可以把 trial 代码目录设置为Path(__file__).parent注意__file__只在标准 Python 中可用Jupyter Notebook 中不可用Linux 且未使用 Conda 的环境可能要把python model.py改成python3 model.py否则可能因找不到解释器而启动失败。配置搜索空间experiment.config.search_space search_space把 Step 2 定义的search_space字典直接赋给experiment.config.search_space即可。配置调优算法这里使用TPETree-structured Parzen Estimator调优器experiment.config.tuner.name TPE experiment.config.tuner.class_args[optimize_mode] maximizetuner.name指定调优器类型NNI 内置了 TPE、Random、Anneal、Evolution 等多种算法完整清单见调优器参考tuner.class_args[optimize_mode]优化方向。本例目标是最大化测试集准确率因此设为maximize若优化的是损失越小越好则应设为minimize。配置要运行的 Trial 数量experiment.config.max_trial_number 10 experiment.config.trial_concurrency 2max_trial_number 10总共评估10 组超参数trial_concurrency 2同时并发评估2 组local 模式下它们会并行占用本机资源。运行时长控制还可以通过max_experiment_duration 1h限制总运行时长。如果既不设max_trial_number也不设max_experiment_duration实验会一直运行下去直到你按下 Ctrl-C。重要提示教程把max_trial_number设为 10 只是为了快速演示真实场景应设得更大——TPE 调优器在默认配置下需要 20 个 trial 来热身warm up即先积累一定数量的随机/历史样本用于建立概率模型样本太少时搜索效果无法体现。Step 4运行实验实验配置完成后即可启动。选择一个端口此处用 8080experiment.run(8080)启动后NNI 会在本机拉起训练管理进程并启动 Web 门户官方渲染输出见 docs/source/tutorials/hpo_quickstart_tensorflow/main.rst大致如下[2022-04-13 12:11:34] Creating experiment, Experiment ID: enw27qxj [2022-04-13 12:11:34] Starting web server... [2022-04-13 12:11:35] Setting up... [2022-04-13 12:11:35] Web portal URLs: http://127.0.0.1:8080 http://192.168.100.103:8080 True此时通过浏览器访问http://localhost:8080即可查看实验状态Trial 列表、各 trial 的学习曲线、超参数对比、最佳 Trial 等。从 nni/experiment/experiment.py 的run方法源码可见run(port, wait_completionTrue)默认会阻塞等待实验完成也可设置debugTrue输出更详细的调试日志。实验结束之后一切完成后即可安全退出。以下内容均为可选操作如果使用的是标准 Python 而非 Jupyter Notebook可以在代码末尾加上input()或signal.pause()阻止 Python 进程退出从而在实验结束后继续查看 Web 门户显式停止实验# input(Press enter to quit) experiment.stop()关于stop()的两个实用细节源码见 nni/experiment/experiment.pynni.experiment.Experiment.stop()会在Python 进程退出时自动被调用因此代码中省略它也是安全的实验停止后还可以调用nni.experiment.Experiment.view()源码重新启动 Web 门户用于事后查看已结束实验的结果。其他实验管理方式本教程全程使用 NNI Python APInni.experiment.Experiment来创建和管理实验。除此之外你也可以使用NNI 命令行工具nnictl来创建、管理、查看和停止实验例如nnictl create --config config.yml、nnictl view、nnictl stop等详见 nnictl 命令行工具教程。两种方式面向不同场景Python API 适合在脚本/Notebook 中端到端编排nnictl 适合命令行交互式管理。小结与下一步至此你已走通了 NNI 上 TensorFlow HPO 的完整链路改造模型把模型独立成脚本接入nni.get_next_parameter()、nni.report_intermediate_result()、nni.report_final_result()三个 API定义搜索空间用choice/uniform/loguniform描述每个超参数的采样方式与范围配置实验通过Experiment(local)指定 trial 命令、代码目录、搜索空间、TPE 调优器与运行规模运行与监控experiment.run(8080)启动实验Web 门户实时查看stop()/view()管理生命周期。这个示例对应的两段源码main.py 与 model.py同时被 docs/source/hpo/overview.rst 与 docs/source/hpo/quickstart.rst 列为官方入门教程入口可作为后续开发自己 HPO 任务的起点模板。想要继续深入可以阅读搜索空间规范掌握全部采样类型通过调优器参考对比不同算法的适用场景或借助 Assessor 参考为实验加入早停机制以加速搜索。【免费下载链接】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),仅供参考

相关推荐

YOLOv8课堂行为检测实战:从数据标注检查到训练避坑指南
YOLOv8课堂行为检测实战:从数据标注检查到训练避坑指南

简介:这份课堂行为检测数据集面向计算机视觉与教育信息化方向的学习者,聚焦真实课堂与教室场景,提供“回答问题”“板书”等4类行为的目标检测标注,可直接用于YOLO系列模型的训练与评估;数据已完成训练集、验证集划分&… · 2026/9/23 10:39:51

VFC实战:高离群点率下的点集配准与仿射模型实现
VFC实战:高离群点率下的点集配准与仿射模型实现

简介:这是一份基于VFC(变分特征对应)的点集配准MATLAB实现资源包,面向计算机视觉、医学图像分析及三维重建等领域的工程师和研究人员,用于解决不同图像间特征点集的对齐与匹配问题。包内共约2000个文件,以.… · 2026/9/23 10:39:51

3秒看懂二寸证件照尺寸,手写实现避坑指南
3秒看懂二寸证件照尺寸,手写实现避坑指南

3秒看懂二寸证件照尺寸,手写实现避坑指南 官方文档太长抓不住重点?别慌。很多应届生做图像处理或表单验证时,卡在“二寸”到底是多少像素上。PIL库的文档翻了三遍,还是不知道DPI怎么算。今天直接上 手写实现 ,用Python代码把这事说透。… · 2026/9/23 10:39:51

cisco1841常见报错与解决
cisco1841常见报错与解决

Cisco 1841选型指南:3个维度看清它还能打吗,附完整示例 很多刚入行的运维或网络工程师,手里攥着几本 Cisco 官方文档,背下了 ACL、OSPF、BGP 的语法,但真到了要搭一个能跑业务的拓扑,或者面对一台老旧的 Cisco… · 2026/9/23 11:25:51

ESD防护设计全解析:从HBM/CDM模型到GGNMOS与版图验证
ESD防护设计全解析:从HBM/CDM模型到GGNMOS与版图验证

简介:这份文档面向集成电路设计、模拟电路与芯片可靠性方向的工程师及学生,系统梳理ESD防护从原理到器件、电路与工艺的完整知识链条,帮助读者建立应对静电放电失效的分析框架。压缩包内为1个docx文件,约10.51MB,内容以… · 2026/9/23 11:25:51

Claude Code Token 消耗暴降65倍:代码地图实战指南
Claude Code Token 消耗暴降65倍:代码地图实战指南

Claude Code 最近在 GitHub 上已经有 30K Star,用过的朋友应该都有同感:这玩意写代码是真猛,Token 消耗也是真夸张。你要是不管它,改一个小功能它能翻遍整个仓库,每次工具调用的中间结果全要送进模型重新算一遍&#x… · 2026/9/23 11:25:51

邻信AI伴侣游戏:LLM驱动分层记忆与主动消息调度实战
邻信AI伴侣游戏:LLM驱动分层记忆与主动消息调度实战

1. 项目缘起与整体设计思路1.1 为什么想到做“邻信”这个AI伴侣游戏做“邻信”这个项目的起点其实很朴素:市面上大多数所谓AI伴侣产品,本质就是一个聊天框加一个系统提示词,用户发一句、模型回一句,聊上十几轮就开始重复、失忆、人… · 2026/9/23 11:25:51

基于Matlab的齿轮箱传递路径分析(TPA)故障诊断实战
基于Matlab的齿轮箱传递路径分析(TPA)故障诊断实战

写这篇东西的起因,是我前段时间帮朋友处理一套减速机试验台的异常振动。传感器装在箱体表面,频谱一看就是典型的齿轮啮合频率边带,但问题在于——传感器测点离故障齿轮隔了好几根轴,中间经过轴承、箱体、螺栓连接面,振… · 2026/9/23 11:25:45

Blender 键盘快捷键速查手册:187 个高频操作一览(Quick Reference 备忘清单版)
Blender 键盘快捷键速查手册:187 个高频操作一览(Quick Reference 备忘清单版)

文档知识库教程开发工具 【免费下载链接】reference 为开发人员分享快速参考备忘清单(速查表) 项目地址: https://gitcode.com/jaywcjlove/reference 点击查看 免费下载 本文是基于 Quick Reference 开源项目 docs/blender.md 备忘清单整理而成的 Blender 快捷键速… · 2026/9/23 11:25:45

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

了解更多?预约专属演示

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

企业微信二维码