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

基于 NNI 的 Kaggle 图像分割竞赛零改造 AutoML 实战:TGS Salt Identification 33 名方案全解析

发布时间:2026/9/23 7:32:49 来源:云帆数科 栏目:资讯中心
基于 NNI 的 Kaggle 图像分割竞赛零改造 AutoML 实战:TGS Salt Identification 33 名方案全解析
人工智能AutoML机器学习深度学习模型压缩特征工程【免费下载链接】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点击查看免费下载导读本文围绕 examples/trials/kaggle-tgs-salt/README.md 展开讲解如何在不改动一行竞赛训练代码的前提下借助 NNI 的注解式annotation-basedAPI 把 Kaggle「TGS Salt Identification Challenge」第 33 名解决方案接入自动机器学习流程用于超参调优。读完本文你将掌握NNI 与既有 PyTorch 训练脚本的集成方式get_next_parameter/report_intermediate_result/report_final_result、NNI 实验配置config.yml的完整字段含义以及该方案从数据预处理到六阶段训练、伪标签与模型集成的完整复现路径。一、项目背景竞赛代码如何「零改动」接入 AutoML该示例的核心理念是先让竞赛代码独立运行再通过 NNI 的注解与配置文件把它包装成可自动调参的 trial代码本身不需要任何结构性修改。这也是 NNI 众多 trial 示例如 examples/trials/mnist-pytorch通用的一贯思路——训练脚本保持可独立运行NNI 通过环境变量注入参数、通过注解收集中间指标与最终指标。原文明确指出这份代码仍可以单独运行standalone但完整复现竞赛结果需要至少一周的训练时间若要接入 NNI只需要配置好config.yml然后执行nnictl create --config config.yml从源码结构看train.py 中共有四处 NNI 注解式调用分别对应 NNI trial 生命周期的四个关键节点注解位置作用源码位置nni.get_next_parameter()从 NNI 获取 tuner 建议的一组超参并覆盖命令行参数train.pynni.variable(nni.choice(UNetResNetV4, UNetResNetV5, UNetResNetV6), namemodel_name)把model_name声明为可搜索的超参变量train.pynni.report_intermediate_result(iout)每个 epoch 验证后上报中间指标用于曲线绘制与 early stop 类算法train.pynni.report_final_result(best_iout)训练结束上报最终指标供 tuner 决定下一组参数train.py对应的 NNI SDK 实现在 nni/trial.pyget_next_parameter()返回 tuner 分配的参数字典report_intermediate_result()与report_final_result()负责把指标通过通信管道回传给 NNI 的调度与调优模块。正是「独立运行 注解 配置」这三者组合实现了零改造接入。二、NNI 实验配置逐字段解读 config.yml仓库提供了两份配置config.ymlLinux/macOS与 config_windows.ymlWindows两者内容除trialCommand的python3/python区别外完全一致。完整内容如下useAnnotation: true trialCommand: python3 train.py trialGpuNumber: 0 trialConcurrency: 2 maxTrialNumber: 10 tuner: name: TPE classArgs: optimize_mode: maximize trainingService: # For other platforms, check mnist-pytorch example platform: local各字段在本次实验中的实际含义useAnnotation: true开启注解模式。NNI 会解析训练脚本中以nni.*形式书写的注解上述四个 API 均依赖此开关并把可调超参注入到 trial 进程。trialCommand: python3 train.py每个 trial 实际执行的命令行。由于训练脚本通过args解析超参NNI 注入的参数与命令行默认参数自然合并这正是「不改代码」能成立的关键。trialGpuNumber: 0每个 trial 分配的 GPU 数。注意此例设为 0表示由训练脚本自行管理 GPU 资源脚本内部直接model.cuda()使用默认 GPU。trialConcurrency: 2同时并行运行的 trial 数。竞赛方案本身需要在多 fold、多模型间切换并行度为 2 能在单机多卡或交替占用 GPU 时平衡吞吐。maxTrialNumber: 10整个实验最多启动 10 个 trial用于控制总搜索预算。tuner: name: TPE使用 Tree-structured Parzen Estimator 贝叶斯优化器classArgs.optimize_mode: maximize声明目标是最大化对应分割指标 IoU 阈值均值。trainingService: platform: local在本地机器上直接运行 trial。注释提示其他平台如远程机器、Kubernetes 等可参考 mnist-pytorch 示例的配置写法。三、数据准备与预处理按 README 的 Preparation 步骤先下载 Kaggle 官方竞赛数据再运行preprocess.py生成训练所需的元数据。preprocess.py 的核心是generate_stratified_metadata()读取train.csv与depths.csv为每个样本计算 mask 覆盖率coverage与覆盖率类别coverage_classcov_to_class按 10% 区间划分 0-10 共 11 类并生成salt_exists二分类标签输出train_meta2.csv包含图像/掩膜路径、is_train、深度z、salt_exists、coverage_class等列使用StratifiedKFold(n_splits10)按覆盖率类别分层切分 10 折将每折的索引写入train_split.json——这是后续--ifolds 0等参数所依赖的折划分依据。数据路径与关键尺寸统一在 settings.py 中定义DATA_DIR默认为/mnt/chicm/data/salt训练/测试图目录、train.csv、depths.csv、meta.csv、models输出目录均由此派生模型输入为 128×128H W 128而原始图像与掩膜为 101×101ORIG_H ORIG_W 101推理时需把 128×128 的预测裁剪/缩放回 101×101。四、六阶段训练流水线详解README 给出的是完整的竞赛复现路径共六个阶段逐级提升模型质量。阶段 1基础训练100 epochs × 3 个模型对 fold 0-3每个 fold 训练三个结构不同的 UNetResNet 变体python3 train.py --ifolds 0 --epochs 100 --model_name UNetResNetV4 python3 train.py --ifolds 0 --epochs 100 --model_name UNetResNetV5 --layers 50 python3 train.py --ifolds 0 --epochs 100 --model_name UNetResNetV6--ifolds 0支持逗号分隔的多个 fold如--ifolds 0,1,2,3脚本在 train.py 中将其解析为列表并逐一训练--epochs 100基础训练轮数三个模型分别对应 models.py 中的UNetResNetV4、UNetResNetV5--layers 50指 ResNet-50 编码器、UNetResNetV6去掉了首层池化以提升分辨率并附带空掩膜分类分支logit_image。V5/V6 的解码器使用转置卷积ConvTranspose2d替代 V4 的双线性上采样。阶段 2余弦退火微调300 epochs用余弦退火学习率调度器对阶段 1 的模型做长程微调python3 train.py --ifolds 0 --epochs 300 --lrs cosine --lr 0.001 --min_lr 0.0001 --model_name UNetResNetV4--lrs cosine对应 train.py 中的CosineAnnealingLR(optimizer, args.t_max, eta_minargs.min_lr)默认--t_max 15即每 15 个 epoch 完成一个余弦周期最低学习率由--min_lr 0.0001约束训练启动时会自动加载同名模型的已有 checkpointbest_{fold}.pth继续训练从而形成「基础训练 → 微调」的衔接train.py。阶段 3加入深度通道微调在阶段 2 基础上追加--depths开关python3 train.py --ifolds 0 --epochs 300 --lrs cosine --lr 0.001 --min_lr 0.0001 --model_name UNetResNetV4 --depths--depths会触发 loader.py 中的add_depth_channel在输入张量的通道 1 写入从 0 到 1 线性渐变的标准深度张量通道 2 写入「原图 × 深度」的乘积为模型提供盐层深度先验。注意该操作发生在cuda()之前属于纯张量运算不改变网络结构。阶段 4预测与生成伪标签对阶段 3 的每个模型做推理预测然后将多个模型的预测结果集成ensemble生成测试集的伪标签pseudo labels。推理侧的实现见 predict.pydo_tta_predict实现了 4 种测试时增强TTA原图、水平翻转、垂直翻转、双向翻转预测后再翻转回来取平均单模型预测保存为{checkpoint}_out/pred.npypostprocessing.py 提供crop_image/resize_image把 128×128 还原到 101×101与binarize阈值二值化其save_pseudo_label_masks可把 RLE 编码的提交结果解码为测试集 mask 图片供下一阶段作为伪标签使用。阶段 5伪标签微调在阶段 3 模型基础上同时开启--depths --pseudopython3 train.py --ifolds 0 --epochs 300 --lrs cosine --lr 0.001 --min_lr 0.0001 --model_name UNetResNetV4 --depths --pseudo--pseudo在 loader.py 中生效get_train_loaders会把测试集元数据此时已具备伪标签 mask追加到训练集实现半监督式的伪标签学习。阶段 6模型集成将阶段 3 与阶段 5 的全部模型预测做集成对多个pred.npy求均值再二值化生成最终提交。predict.py 的ensemble_np展示了这一流程加载多个 npy、np.mean求平均、裁剪到 101×101、二值化最后通过create_submissionutils.py 中的run_length_encoding转成 Kaggle 要求的 RLE 编码提交文件。五、训练脚本核心参数速查train.py 的命令行参数构成了整个方案的「超参面」也是 NNI 可以自动搜索的对象。整理如下参数默认值说明--layers34ResNet 编码器深度可选 34/50/101/152见create_resnet--nf32模型基础通道数num_filters--lr0.001初始学习率--min_lr0.0001学习率下限cosine 的eta_min/ plateau 的min_lr--ifolds0逗号分隔的 fold 列表--batch_size32批大小--start_epoch0起始 epoch配合 checkpoint 续训--epochs200总训练轮数--optimSGD可选 SGD / AdamAdam 可叠加--adamw权重衰减--lrscosine可选 cosine / plateauReduceLROnPlateau--patience/--factor/--t_max6 / 0.5 / 15学习率调度器参数--pad_modeedge填充方式可选 reflect / edge / resize--model_nameUNetResNetV4模型结构选择--init_ckpNone从指定 checkpoint 恢复--valFalse仅验证不训练--store_loss_modelFalse额外保存混合分数最优模型--train_clsFalse开启空掩膜分类辅助损失--meta_version2元数据版本1 用普通 KFold2 用分层 10 折--pseudo/--depths/--dev_modeFalse伪标签 / 深度通道 / 开发模式各取 10 条样本快速调试值得注意的工程细节损失函数weighted_losstrain.py组合了 Lovász hinge 损失lovasz_hinge与 Focal 损失权重 0.2可选叠加salt_output的空掩膜二分类交叉熵评估指标validate每轮计算两类指标——普通 IoUintersection_over_union和竞赛官方指标 IoU 阈值均值intersection_over_union_thresholds对 0.5~0.95 每 0.05 共 10 个阈值取平均实现在 metrics.py。NNI 上报的iout即后者这也是optimize_mode: maximize的对应目标模型选择按iout保存best_{fold}.pth可按--store_loss_model追加保存混合分数最优的_loss模型。六、在 NNI 上运行与独立运行的对照两种运行方式共享同一个train.py区别仅在参数来源独立运行直接执行python3 train.py --epochs 100 --model_name UNetResNetV4 ...nni.get_next_parameter()注解在非 NNI 环境下被忽略代码退化为普通训练脚本NNI 运行nnictl create --config config.yml启动实验后NNI 为每个 trial 注入参数覆盖model_name等被nni.variable声明的超参并通过report_intermediate_result/report_final_result持续收集指标供 TPE tuner 决策下一组参数Web 界面可实时观察各 trial 的iout曲线与训练进度。七、适用前提与复现注意事项算力成本README 明确指出完整复现至少需要一周训练时间属于完整竞赛级流水线建议先用--dev_mode各折仅取 10 条样本验证链路再全量训练环境依赖需要 PyTorch、torchvision、opencv、pycocotoolsIoU 计算依赖 COCO API见 metrics.py、Keraspreprocess.py中的load_img以及 tqdm、pandas 等库路径配置所有数据路径集中在 settings.py 的DATA_DIR首次使用前必须改为本地数据目录Windows 支持仓库同时提供 config_windows.yml命令从python3调整为pythonGPU 策略trialGpuNumber: 0表示不限制单 trial 的 GPU 数量多 trial 并行trialConcurrency: 2时需自行留意显存占用。结语该示例完整展示了「竞赛代码 AutoML」的落地范式以注解为桥、以配置为闸把一套原本需要人工反复试错的六阶段分割流水线交给 NNI 的 TPE tuner 自动搜索超参同时保留代码独立运行与全量复现的能力。无论是想要复现 TGS 竞赛结果还是希望把既有 PyTorch 训练脚本快速接入 NNIexamples/trials/kaggle-tgs-salt 都是一个可以直接参照的最小完整闭环。赞分享人工智能AutoML机器学习深度学习模型压缩特征工程【免费下载链接】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 为 Kaggle 竞赛代码零改动接入 AutoML以 TGS Salt Identification 第 33 名方案为例用 NNI 为 Kaggle 竞赛代码零改动接入 AutoML以 TGS Salt Identification 第 33 名方案为例 本篇技术指南以 NNI人工智能AutoML机器学习深度学习模型压缩特征工程PaddlePaddle深度学习实战Kaggle CIFAR-10图像分类竞赛全流程解析PaddlePaddle深度学习实战Kaggle CIFAR 10图像分类竞赛全流程解析 引言 计算机视觉是深度学习最重要的应用领域之一而图像分类作为计算机示例工程教程深度学习G-Helper三步解锁华硕笔记本的完整性能控制权G Helper三步解锁华硕笔记本的完整性能控制权 你是否厌倦了官方控制软件Armoury Crate的臃肿和卡顿G Helper作为一款轻量级开源工具为桌面应用系统编程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

相关推荐

Spark与Kafka构建智能家居实时数据分析系统实战
Spark与Kafka构建智能家居实时数据分析系统实战

简介:这是一套面向智能家居设备数据分析的完整源码包,适合物联网开发者、数据工程学习者以及需要快速搭建流式处理管道的读者。项目基于Apache Spark与Kafka构建,通过MQTT协议采集传感器数据,经由HDFS存储、Spark分析后写入Postgr… · 2026/9/23 7:32:49

Thunderbird for Android 中的 OpenPGP API 协议演进全解析:从 CHANGELOG 到 AIDL 服务源码
Thunderbird for Android 中的 OpenPGP API 协议演进全解析:从 CHANGELOG 到 AIDL 服务源码

Thunderbird for Android 中的 OpenPGP API 协议演进全解析:从 CHANGELOG 到 AIDL 服务源码 【免费下载链接】thunderbird-android Thunderbird for Android – Open Source Email App for Android (fka K-9 Mail) 项目地址: https://gitcode.com/gh_mirrors/th/t… · 2026/9/23 7:32:49

Yii 2 服务定位器(Service Locator)完全指南:从 `yii\di\ServiceLocator` 到应用组件的注册、解析与模块树遍历
Yii 2 服务定位器(Service Locator)完全指南:从 `yii\di\ServiceLocator` 到应用组件的注册、解析与模块树遍历

后端Web框架 【免费下载链接】yii2 Yii 2: The Fast, Secure and Professional PHP Framework 项目地址: https://gitcode.com/gh_mirrors/yi/yii2 点击查看 免费下载 导读 服务定位器(Service Locator)是 Yii 2 框架中负责"按 ID 提供… · 2026/9/23 7:32:43

基于LSTM的电力负荷时间序列预测:从数据清洗到多步预测完整实践
基于LSTM的电力负荷时间序列预测:从数据清洗到多步预测完整实践

简介:这是一份基于深度学习算法实现电力负荷时间序列未来预测的 Python 源码项目,围绕负荷历史数据完成特征构造、模型训练与结果评估,覆盖 LSTM、GRU、Transformer、ARIMA、随机森林、决策树、KNN 等多种算法,适合计科、人工智能… · 2026/9/23 8:13:40

从Keil5迁移到VSCode+GCC:GD32开发环境搭建与实战指南
从Keil5迁移到VSCode+GCC:GD32开发环境搭建与实战指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/23 8:13:40

用Office文档搭建企业AI知识库:从RAG原理到Dify实操指南
用Office文档搭建企业AI知识库:从RAG原理到Dify实操指南

这两年AI大模型火得一塌糊涂,几乎每个企业都在琢磨怎么把AI真正用起来。可我接触了这么多客户和同行,发现大家碰到的第一个瓶颈往往不是模型不够聪明,而是企业自己的数据根本喂不进去。很多公司的核心经验、流程、制度、技术文档,… · 2026/9/23 8:13:40

obsidian-livesync 仓库的 AI 编码助手规范:读懂 AGENTS.md 中的协作、风格与发布纪律
obsidian-livesync 仓库的 AI 编码助手规范:读懂 AGENTS.md 中的协作、风格与发布纪律

数据同步 【免费下载链接】obsidian-livesync 项目地址: https://gitcode.com/gh_mirrors/ob/obsidian-livesync 点击查看 免费下载 Self-hosted LiveSync(obsidian-livesync)是一个用于跨设备同步 Obsidian 库的插件,代码库采用… · 2026/9/23 8:13:34

智能学术专著创作平台:提升效率与质量的新范式
智能学术专著创作平台:提升效率与质量的新范式

1. 项目概述:学术专著创作的新范式去年协助一位教授完成学科评估材料时,我亲眼见证了一部学术专著从构思到出版的完整历程。这位学者花费了整整八个月时间,每天工作到凌晨两点,最终交稿时体重下降了12斤。这种"学术苦修"… · 2026/9/23 8:13:34

dnf元素觉醒叫什么新手避坑
dnf元素觉醒叫什么新手避坑

5个坑让你DNF元素觉醒从入门到精通 刚接触DNF元素觉醒的玩家,是不是也遇到过这种尴尬:看着攻略把技能点加满了,结果进图一放火球,伤害低得可怜;或者明明照着视频操作,觉醒技能却放不出来,卡在原地干着急。这种“学会了操作逻辑,却不知道怎么在… · 2026/9/23 8:13:34

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

了解更多?预约专属演示

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

企业微信二维码