- 计算机视觉
- 深度学习
- 人工智能
- 图像处理
【免费下载链接】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 不兼容。其核心变化可以归纳为四条:
SegmentationModelsBuilder.build()不再负责构建网络:旧版本需要传入model_name、encoder_name、encoder_weights、in_channels、classes、activation和**kwargs,由 Kornia 惰性导入 smp、实例化对应架构并自行查询编码器预处理参数;新版本改为接收一个已构建好的nn.Module和编码器预处理参数字典,签名变为build(model, preproc_params=None, name="segmentation_model")。preproc_params的来源明确化:它正是smp.encoders.get_preprocessing_params(encoder_name)的返回值,网络构建与参数获取全部由调用方完成。- Kornia 全仓库不再 import smp:
kornia.core.external.segmentation_models_pytorch模块被移除。 SemanticSegmentation不再是抽象类:它获得了与其他兄弟容器一致的__init__(model, pre_processor, post_processor, name=None)构造方法,builder 返回的容器因此可以被直接实例化。- 预处理数值细节修正:
input_range: [0, 255]的缩放步骤从用存储的1/255做Normalize除法,改为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:三个参数的含义如下:
| 参数 | 类型 | 说明 |
|---|---|---|
model | nn.Module | 已构建好的分割网络,将(B, 3, H, W)图像批映射为(B, C, H, W)预测;builder 会将其置为eval模式 |
preproc_params | dict或None | 编码器预处理参数字典,即smp.encoders.get_preprocessing_params(encoder_name)的返回值;None表示不做任何预处理(等价于nn.Identity) |
name | str | 包装后的模型名称,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 何时可以传入任意网络
从源码结构看,SegmentationModelsBuilder对model的唯一硬性要求是:它是一个把(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 返回的对象是SemanticSegmentation、model.model is net、网络处于非训练模式。
测试test_none_params_is_identity_preprocessing还验证了preproc_params=None时pre_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:
- 颜色空间转换:
input_space == "BGR"时插入kornia.color.BgrToRgb();"RGB"则跳过;其他值抛出ValueError("Unsupported input space: ...")。 - 范围缩放:
input_range[1] == 255时插入kornia.enhance.Rescale(255.0);== 1则跳过;其他值抛出ValueError。 - 均值方差归一化:
mean/std非None时转为张量后插入kornia.enhance.Normalize(mean=mean, std=std);为None时分别用0.0和1.0兜底,保证归一化步骤恒等。
测试test_preprocessing_bgr_255用 float64 手工计算的参考结果(x.flip(1) * 255.0 - mean_t) / std_t与管线输出对比,验证了“先翻转通道、再乘 255、最后逐通道归一化”的精确语义;test_exception_unsupported_input_space与test_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_exact用torch.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 友好性,这也是本次重构的设计目标之一:
- 管线本身是 ONNX 友好的:
get_preprocessing_pipeline返回的ImageSequential中只有BgrToRgb、Rescale、Normalize三个可导出的张量算子,没有任何 Python 控制流。测试test_preprocessing_onnx_export用torch.onnx.export(..., dynamo=True, opset_version=18)导出管线并在 onnxruntime 中推理,与 eager 模式结果在rtol=1e-5, atol=1e-5内一致。 - 导出前必须关闭便捷特性:源码 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
🐍 空间人工智能的几何计算机视觉库
相关推荐
Kornia SegmentationModelsBuilder 迁移指南:解耦 segmentation_models_pytorch 依赖与 Rescale 精度修复
Kornia SegmentationModelsBuilder 迁移指南:解耦 segmentation_models_pytorch 依赖与 Rescale
计算机视觉人工智能深度学习图像处理Kornia 可选依赖 extras 体系解析:kornia[onnx]、kornia[sd] 与懒加载依赖管理
Kornia 可选依赖 extras 体系解析:kornia onnx 、kornia sd 与懒加载依赖管理 导读 本文围绕 Kornia 的可选依赖(opt
计算机视觉人工智能深度学习图像处理Czkawka终极指南:免费快速清理重复文件,释放硬盘空间
Czkawka终极指南:免费快速清理重复文件,释放硬盘空间 还在为电脑里堆积的重复文件烦恼吗?Czkawka是一款功能强大的跨平台文件清理工具,能够快速找出并清
桌面应用
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考