news 2026/9/24 7:37:18

Kornia 语义分割模型集成重构解析:SegmentationModelsBuilder 去 smp 依赖与新预处理管线(PR 4301)

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Kornia 语义分割模型集成重构解析:SegmentationModelsBuilder 去 smp 依赖与新预处理管线(PR 4301)
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 图像处理

【免费下载链接】kornia

🐍 空间人工智能的几何计算机视觉库

项目地址:https://gitcode.com/kornia/kornia
点击查看免费下载

本指南以 changelog 片段 changelog.d/+migration-005.breaking.md 记录的一次破坏性变更(Breaking Change,对应 PR #4301)为主线,完整讲解 Kornia 中SegmentationModelsBuilder.build()的新旧两代 API、SemanticSegmentation容器的可实例化改造,以及input_range: [0, 255]预处理步骤从Normalize除法改为Rescale(255.0)乘法的底层数值原因。读者读完将掌握:如何在不再引入segmentation_models_pytorch(smp)的情况下用 Kornia 包装任意分割网络、如何获取并传入编码器预处理参数、为什么乘法替换除法能修复 bfloat16 下的精度偏差,以及 ONNX 导出时的注意事项。

一、变更总览:一次解耦外部依赖的破坏性重构

该 fragment 属于breaking类型,意味着 API 不兼容。其核心变化可以归纳为四条:

  1. SegmentationModelsBuilder.build()不再负责构建网络:旧版本需要传入model_nameencoder_nameencoder_weightsin_channelsclassesactivation**kwargs,由 Kornia 惰性导入 smp、实例化对应架构并自行查询编码器预处理参数;新版本改为接收一个已构建好的nn.Module编码器预处理参数字典,签名变为build(model, preproc_params=None, name="segmentation_model")
  2. preproc_params的来源明确化:它正是smp.encoders.get_preprocessing_params(encoder_name)的返回值,网络构建与参数获取全部由调用方完成。
  3. Kornia 全仓库不再 import smpkornia.core.external.segmentation_models_pytorch模块被移除。
  4. SemanticSegmentation不再是抽象类:它获得了与其他兄弟容器一致的__init__(model, pre_processor, post_processor, name=None)构造方法,builder 返回的容器因此可以被直接实例化。
  5. 预处理数值细节修正input_range: [0, 255]的缩放步骤从用存储的1/255Normalize除法,改为kornia.enhance.Rescale(255.0)的乘以 255 乘法,修复 bfloat16 下的精度问题。

二、新旧 API 对比与迁移路径

2.1 旧 API:builder 内部完成“导入 + 构建 + 查参”三步

在重构之前,SegmentationModelsBuilder.build()是一个典型的“便利工厂”:调用方只需描述自己想要什么模型,builder 负责其余一切,包括在运行时惰性导入 smp。这也是它存在的主要问题——Kornia 的kornia.core.external中曾长期维护segmentation_models_pytorch这样一个第三方库的惰性包装,任何 smp 的 API 变动、安装失败或版本兼容问题都可能传导到 Kornia 自身。

旧调用方式大致形如:

from kornia.models.segmentation import SegmentationModelsBuilder model = SegmentationModelsBuilder.build( model_name="Unet", encoder_name="resnet34", encoder_weights="imagenet", in_channels=3, classes=2, activation="softmax2d", )

builder 会据此拼接出 smp 构造参数,惰性导入 smp 模块,实例化架构,再自行调用编码器的预处理参数查询逻辑。

2.2 新 API:职责反转,调用方自备网络与参数

新签名(见 kornia/models/segmentation/segmentation_models.py 中的SegmentationModelsBuilder.build):

@staticmethod def build( model: nn.Module, preproc_params: Optional[dict[str, Any]] = None, name: str = "segmentation_model", ) -> SemanticSegmentation:

三个参数的含义如下:

参数类型说明
modelnn.Module已构建好的分割网络,将(B, 3, H, W)图像批映射为(B, C, H, W)预测;builder 会将其置为eval模式
preproc_paramsdictNone编码器预处理参数字典,即smp.encoders.get_preprocessing_params(encoder_name)的返回值;None表示不做任何预处理(等价于nn.Identity
namestr包装后的模型名称,SemanticSegmentation.save()会用它生成文件名;默认"segmentation_model"

标准的迁移后调用方式(同时见 segmentation_models.py 中的 doctest):

import segmentation_models_pytorch as smp from kornia.models.segmentation import SegmentationModelsBuilder net = smp.Unet( encoder_name="resnet34", encoder_weights="imagenet", classes=2, activation="softmax2d", ) params = smp.encoders.get_preprocessing_params("resnet34") model = SegmentationModelsBuilder.build(net, params, name="Unet_resnet34") model(torch.rand(1, 3, 64, 64)).shape # torch.Size([1, 2, 64, 64])

2.3 何时可以传入任意网络

从源码结构看,SegmentationModelsBuildermodel的唯一硬性要求是:它是一个把(B, 3, H, W)输入映射为(B, C, H, W)预测的nn.Module。因此它不限于 smp 网络——任何满足该形状约定的 PyTorch 模块都可以被包装。这一点在 tests/models/test_segmentation_models.py 中有直接验证:测试用nn.Sequential(nn.Conv2d(3, classes, kernel_size=1), nn.Softmax(dim=1))作为 smp 模型的替身(_stand_in_network),并断言 builder 返回的对象是SemanticSegmentationmodel.model is net、网络处于非训练模式。

测试test_none_params_is_identity_preprocessing还验证了preproc_params=Nonepre_processor就是nn.Identity,输入原样送入网络。

三、preproc_params 的字段契约与预处理管线构建

3.1 四个必需字段

get_preprocessing_pipeline(见 segmentation_models.py)要求字典必须包含以下四个键(_PREPROC_KEYS = ("input_space", "input_range", "mean", "std")),缺失任何一个都会触发KORNIA_CHECK抛出异常,错误信息为preproc_params is missing the key '<key>'

允许取值语义
input_space"RGB""BGR"网络训练时使用的颜色通道顺序;"BGR"网络需要先把 RGB 输入翻转通道
input_range[0, 1][0, 255]mean/std所基于的像素取值范围;[0, 255]表示先把[0, 1]输入乘以 255
mean每通道列表或None归一化均值;None表示不做均值归一化
std每通道列表或None归一化标准差;None表示不做标准差归一化

一个典型的 ImageNet 预训练参数(与测试文件中的IMAGENET_PARAMS完全一致):

IMAGENET_PARAMS = { "input_space": "RGB", "input_range": [0, 1], "mean": [0.485, 0.456, 0.406], "std": [0.229, 0.224, 0.225], }

3.2 管线构建顺序:颜色 → 缩放 → 归一化

get_preprocessing_pipeline按以下顺序组装一个kornia.augmentation.container.ImageSequential

  1. 颜色空间转换input_space == "BGR"时插入kornia.color.BgrToRgb()"RGB"则跳过;其他值抛出ValueError("Unsupported input space: ...")
  2. 范围缩放input_range[1] == 255时插入kornia.enhance.Rescale(255.0)== 1则跳过;其他值抛出ValueError
  3. 均值方差归一化mean/stdNone时转为张量后插入kornia.enhance.Normalize(mean=mean, std=std);为None时分别用0.01.0兜底,保证归一化步骤恒等。

测试test_preprocessing_bgr_255用 float64 手工计算的参考结果(x.flip(1) * 255.0 - mean_t) / std_t与管线输出对比,验证了“先翻转通道、再乘 255、最后逐通道归一化”的精确语义;test_exception_unsupported_input_spacetest_exception_unsupported_input_range则覆盖了非法值的报错路径。

四、核心精度修复:为什么用 Rescale(255.0) 乘法替代 Normalize(1/255) 除法

4.1 问题的根源:bfloat16 下的倒数舍入

旧实现把1/255作为缩放系数存进Normalize,等价于让每个像素乘以一个预存的1/255。问题出在这个系数本身无法在所有浮点格式中被精确表示

  • bfloat16 只能表示 8 位有效数字,1/255被舍入为约0.0039368,把它当作“乘 255”的等价操作时,实际等效乘数变成了约254.0
  • 结果就是0.5被映射到127.0而不是精确的127.5——对半精度推理,这是一个系统性的、肉眼可见的偏移。

4.2 新实现:乘以 255,一个任何 dtype 都能精确表示的数

新实现改为kornia.enhance.Rescale(255.0)直接做乘法。255是整数,在 float32、float16、bfloat16 中都可以被精确存储,因此乘 255 不会引入任何舍入误差。fragment 明确指出:bfloat16 场景下0.5现在正确映射为127.5;而 float32 输出在该步骤的移动至多 1 个 ulp(unit in the last place,最小精度单位)——即误差上界可忽略。

Rescale的实现见 kornia/enhance/rescale.py:它把因子以非持久 buffer形式注册(register_buffer("factor", factor, persistent=False)),因此.to(device)会把因子一并迁移到目标设备;forward就是一次input * self.factor。非持久设计意味着因子不会进入state_dict(),不破坏既有 checkpoint 的加载。fragment 还特别提到,Rescale把因子保存为 0 维张量,因此可以从任意设备导出 ONNX

测试test_preprocessing_255_rescale_is_exacttorch.equal做严格相等断言验证了这一点:输入[0.0, 0.5, 1.0]在 float16 与 bfloat16 下经过管线后逐位等于[0.0, 127.5, 255.0]。如果是除法实现,bfloat16 下0.5 → 127.0,该测试必然失败——这正是本次改动要修复的回归。

五、SemanticSegmentation 容器:从抽象类到可实例化

5.1 新构造签名

base.py 中的SemanticSegmentation现在拥有与 Kornia 其他模型容器一致的构造方法:

def __init__( self, model: nn.Module, pre_processor: nn.Module, post_processor: nn.Module, name: Optional[str] = None, ) -> None:
  • model:分割网络,(B, 3, H, W) → (B, C, H, W)__init__中立即执行self.model = model.eval()
  • pre_processor/post_processor:前后处理模块;builder 默认用get_preprocessing_pipeline生成的ImageSequential作为前处理、nn.Identity()作为后处理。
  • name:可选,覆盖默认的"segmentation",供save()生成文件名。

forward同时支持(B, 3, H, W)张量批和[(3, H, W), ...]图像列表两种输入,列表路径会逐图执行“预处理 → 网络 → 后处理”。测试test_list_input验证了两种路径输出一致。

5.2 抽象占位符的保留

from_config()仍然存在,但实现为raise NotImplementedError,提示用户改用SegmentationModelsBuilder.build()或直接实例化SemanticSegmentation。测试test_from_config_not_implemented专门锁定了这一行为,确保未来不会有人误以为该抽象入口可用。

5.3 可视化与保存辅助

容器还提供开箱即用的可视化能力:

  • visualize(images, semantic_masks=None, output_type="torch", colormap="random", manual_seed=2147):把概率图转成彩色分割图。注意它对 softmax 输出有硬性要求——visualize_output会先探测语义掩码的类别维求和是否接近 1(容差按 dtype 的eps缩放,以兼容半精度),若不是概率分布(如裸 logits)则抛出ValueError。因此用 smp 时需设置activation="softmax2d"让网络带 softmax 头,测试test_visualize_rejects_logits验证了裸 logits 被拒绝。
  • save(...):输出原图(_src)、彩色掩码(_mask)和加权叠加图(_overlay,通过kornia.enhance.add_weighted以 0.5/0.5 权重融合)三组结果。

六、ONNX 导出注意事项

fragment 及源码文档都强调了 ONNX 友好性,这也是本次重构的设计目标之一:

  1. 管线本身是 ONNX 友好的get_preprocessing_pipeline返回的ImageSequential中只有BgrToRgbRescaleNormalize三个可导出的张量算子,没有任何 Python 控制流。测试test_preprocessing_onnx_exporttorch.onnx.export(..., dynamo=True, opset_version=18)导出管线并在 onnxruntime 中推理,与 eager 模式结果在rtol=1e-5, atol=1e-5内一致。
  2. 导出前必须关闭便捷特性:源码 Note 明确指出,导出前应设置pipeline.disable_features = True。这是因为ImageSequential.__call__默认的输入/输出便捷转换与输出缓存会在张量上附加属性,而某些版本的torch.export会拒绝这类张量属性变更(测试注释提到 PyTorch 2.9 即如此)。
pipeline = SegmentationModelsBuilder.get_preprocessing_pipeline(params).to(device).eval() pipeline.disable_features = True torch.onnx.export(pipeline, (x,), dynamo=True, opset_version=18)

七、迁移检查清单

如果你正在升级 Kornia 并使用了分割模型集成功能,按以下步骤核对:

  • build(model_name=..., encoder_name=..., ...)改为先自行构造smp.Unet(...)等网络,再调用build(net, smp.encoders.get_preprocessing_params(encoder_name), name=...)
  • 若需要可视化输出,确认网络带了 softmax 头(smp 的activation="softmax2d");
  • 检查代码中是否引用了kornia.core.external.segmentation_models_pytorch——该模块已移除,直接删除相关引用;
  • 若对 bfloat16 推理敏感(如半精度部署),验证input_range=[0, 255]的模型输出,新管线应为逐位精确的乘 255 结果;
  • ONNX 导出时设置pipeline.disable_features = True

八、相关源码与测试索引

  • 变更记录:changelog 片段 changelog.d/+migration-005.breaking.md(PR #4301);变更收集流程见 changelog.d/README.md
  • Builder 与预处理管线实现:kornia/models/segmentation/segmentation_models.py
  • 容器类实现:kornia/models/segmentation/base.py
  • 新的缩放算子:kornia/enhance/rescale.py
  • 测试覆盖:tests/models/test_segmentation_models.py(覆盖 API 迁移、四字段校验、BGR/255 精确数学、ONNX 导出、可视化拒绝 logits 等全部行为)
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 图像处理

【免费下载链接】kornia

🐍 空间人工智能的几何计算机视觉库

项目地址:https://gitcode.com/kornia/kornia
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/24 7:31:32

ESP32-S3-BOX-3实战:智能语音与物联网联动开发指南

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

作者头像 李华
网站建设 2026/9/24 7:14:31

Modbus转MQTT数据采集全流程:从RS485到云端实战指南

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

作者头像 李华
网站建设 2026/9/24 7:06:04

在 CI/CD 流水线中使用 Regal 对 Rego 策略进行代码检查

后端认证鉴权云原生 【免费下载链接】opa Open Policy Agent (OPA) is an open source, general-purpose policy engine. 项目地址&#xff1a; https://gitcode.com/gh_mirrors/op/opa 点击查看 免费下载 Regal 是 Open Policy Agent 生态中专门用于 Rego 策略代码的 linter …

作者头像 李华