Kornia GuidedBlur 半精度修复:多通道引导下 float16/bfloat16 引导滤波的求解器适配
【免费下载链接】kornia🐍 Geometric Computer Vision Library for Spatial AI项目地址: https://gitcode.com/gh_mirrors/ko/kornia
本篇文章聚焦 Kornia 图像滤波模块中guided_blur/GuidedBlur(引导滤波)的一次关键修复:当引导图(guidance)包含多个通道时,算子此前在float16/bfloat16半精度下会因torch.linalg.solve缺乏半精度 LU 分解 kernel 而直接报错甚至中止进程。文章将完整还原问题根因、修复策略、源码级实现细节与回归测试验证,帮助读者理解如何在自己的项目中安全地以半精度运行多通道引导滤波,并给出可复制的函数式与模块式调用示例。
背景:引导滤波在 Kornia 中的实现
引导滤波(Guided Image Filtering)是一种经典的保边平滑算子,由 He 等人于 2010 年提出(对应仓库 docs/source/references.bib 中的he2010guided),其后 He 与 Sun 又提出加速版本 Fast Guided Filter(he2015fast)。其核心假设是:在每个局部窗口内,输出可以被表示为引导图的一个线性变换,因此平坦区域被平滑平均,而引导图中较强的结构被保留下来。
Kornia 在 kornia/filters/guided.py 中实现了完整的引导滤波逻辑,并通过 kornia/filters/init.py 对外导出两个 API:
- 函数式:
guided_blur(guidance, input, kernel_size, eps, ...) - 模块式:
GuidedBlur(kernel_size, eps, ...)(继承nn.Module)
两者支持引导图与输入图通道数不一致(如 3 通道 RGB 引导 4 通道输入),参数包括border_type、subsample(Fast Guided Filter 下采样)与separable(可分离盒式滤波)。根据输入guidance.shape[1]是否为 1,实现会分流到两条路径(见 kornia/filters/guided.py):
- 单通道引导:
_guided_blur_grayscale_guidance,纯盒式滤波与逐元素运算,不涉及矩阵求解; - 多通道引导:
_guided_blur_multichannel_guidance,每个像素需要求解一个C x C的线性系统,这是本次修复的核心对象。
问题现场:半精度下的两类崩溃
在修复之前,当引导图通道数C > 1(即进入_guided_blur_multichannel_guidance)且输入为float16或bfloat16时,每次调用torch.linalg.solve都会失败,且不同后端表现不同(见 changelog.d/+migration-116.fixed.md):
- CPU 后端:抛出
NotImplementedError: "lu_cpu" not implemented for 'Half',即 PyTorch 的 LU 分解在 CPU 上没有半精度实现; - MPS 后端:直接触发硬断言
Only MPSDataTypeFloat32 is supported,这不是普通异常而是会直接中止进程,危害更大。
值得注意的是,单通道引导路径(灰度引导)永远不会进入求解器,因此完全不受影响——崩溃仅出现在guide_dim > 1的场景(这一结论在 tests/filters/test_guided.py 的回归测试文档字符串中有明确说明)。
隐藏的第二个 Bug:张量 eps 参与 dtype 提升
除了求解器本身缺乏半精度 kernel,还存在一个容易被忽视的 dtype 提升陷阱。eps(正则化参数)既可以是普通 Python 浮点数,也可以是torch.Tensor。当用户传入常见的eps=torch.tensor(0.1)时,该张量默认是float32:
- 构造系数矩阵
A = var_I + _eps时,float32的_eps会把float16的var_I提升(promote)为 float32; - 而右侧项
cov_Ip仍保持float16; - 结果就是交给
solve的一对操作数 dtype 不一致(矩阵比右端项更宽),torch.linalg.solve会拒绝这对不匹配的输入。
也就是说,即便先绕过半精度 LU kernel 的问题,eps的 dtype 提升也会让求解器直接报错。这是"多通道 + 半精度"场景下同一崩溃的两个叠加根因。
修复方案:统一求解 dtype,float32 求解后转回
修复的核心思路非常克制:在进入torch.linalg.solve之前,把两个操作数统一到一个求解 dtype,半精度一律提升到 float32,求解完成后把结果 cast 回原 dtype。对应源码位于 kornia/filters/guided.py:
if isinstance(eps, torch.Tensor): _eps = torch.eye(C, device=guidance.device, dtype=guidance.dtype).view(1, 1, 1, C, C) * eps.view(-1, 1, 1, 1, 1) else: _eps = guidance.new_full((C,), eps).diag().view(1, 1, 1, C, C) A = var_I + _eps solve_dtype = torch.promote_types(A.dtype, cov_Ip.dtype) if solve_dtype in (torch.float16, torch.bfloat16): solve_dtype = torch.float32 a = torch.linalg.solve(A.to(solve_dtype), cov_Ip.to(solve_dtype)).to(cov_Ip.dtype)关键设计决策(源码注释中均有明确交代):
- 先提升、再判半精度:先通过
torch.promote_types求出两个操作数的共同 dtype,若恰好是float16或bfloat16再抬升到float32。这样既覆盖了半精度输入,也天然解决了eps为 float32 张量导致的 dtype 分裂——两个操作数最终被统一到同一个 dtype。 - 刻意不使用
_torch_solve_cast这类工具:因为那类封装会把float32也提升到float64,不仅改变默认路径的数值结果,还会让每个像素都付出一次 float64 求解的代价。本次修复只在半精度时抬升,float32/float64输入完全走原来的路径。 Tensor.to在 dtype 已匹配时是 no-op:因此全部输入均为 float32 或 float64 的调用,求解前后不做任何转换,行为与修复前逐位一致(bit-identical)。
兼容性保证:默认路径数值不变、编译路径不受影响
修复对既有用户是透明的,changelog 明确给出了三项兼容性保证(见 changelog.d/+migration-116.fixed.md):
float32与float64结果与修复前逐位一致:因为统一 dtype 的逻辑对非半精度是 no-op;- 单通道引导不受影响:该路径从不进入求解器;
- 算子仍可在
fullgraph=True下以单图(single graph)编译:Tensor.to与条件提升均不破坏 TorchScript /torch.compile的图结构。
这一点在 tests/filters/test_guided.py 的test_dynamo中有直接验证:该测试对GuidedBlur(含float与torch.Tensor两种eps形态)使用torch_optimizer优化后与 eager 模式结果对比,覆盖kernel_size为5/(5, 7)、subsample为1/2、separable为False/True的组合。
回归测试:失败率从 98 到 0
本次修复配套了专门的回归测试test_multichannel_guidance_in_half_precision(见 tests/filters/test_guided.py),其测试策略值得借鉴:
- 仅在
float16/bfloat16下运行(其他 dtype 直接pytest.skip),是纯粹针对求解器路径的回归测试; - 使用
border_type="constant"而非默认reflect:因为默认反射边界需要半精度的reflection_pad2d,而 CPU PyTorch 2.5.1 对 float16 没有该实现,用constant可以确保测试真正落在求解器上(排除边界填充的干扰); - 同时覆盖
eps=0.1(float)与eps=torch.tensor(0.1)(tensor)两种形态,恰好对应上述两个根因; - 断言输出 dtype 保持为原半精度、所有值有限(
torch.isfinite),并以 float32 版本结果为基准,用8 * torch.finfo(dtype).eps的容差做数值对比——注释说明:由于求解在 float32 中进行,结果误差应受"周围半精度算术累积误差"支配,而非求解器自身精度。
整体效果在 changelog 中有量化数据:tests/filters/test_guided.py在--dtype=float16,bfloat16下从98 失败 / 67 通过提升到169 全部通过。
API 参考与实战用法
函数式接口guided_blur
完整签名见 kornia/filters/guided.py:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
guidance | Tensor | — | 引导图,形状(B, C, H, W),通道数可为 1 或多通道 |
input | Tensor | — | 待滤波输入,形状(B, C, H, W),通道数可与 guidance 不同 |
kernel_size | int/tuple[int, int] | — | 滤波核大小 |
eps | float/Tensor | — | 正则化参数,值越小保留的边缘越多;可以是逐 batch 的张量 |
border_type | str | 'reflect' | 边界填充模式:'constant'、'reflect'、'replicate'、'circular' |
subsample | int | 1 | Fast Guided Filter 下采样因子,1表示不采样 |
separable | bool | False | 用两次一维盒式滤波代替二维,大窗口下可减少计算量 |
函数内部会校验 guidance/input 的 batch 与空间维度一致(KORNIA_CHECK系列断言,见 kornia/filters/guided.py),并根据 guidance 通道数自动分派到灰度或矩阵求解路径。
最小可运行示例
import torch import kornia # 函数式:3 通道引导、4 通道输入,直接在半精度下运行 guidance = torch.rand(2, 3, 5, 5, dtype=torch.float16) input = torch.rand(2, 4, 5, 5, dtype=torch.float16) output = kornia.filters.guided_blur(guidance, input, kernel_size=3, eps=0.1) print(output.shape, output.dtype) # torch.Size([2, 4, 5, 5]) torch.float16 # bfloat16 同样支持 output_bf = kornia.filters.guided_blur( guidance.to(torch.bfloat16), input.to(torch.bfloat16), 5, torch.tensor(0.1) ) # 模块式:与 nn.Module 生态无缝集成 from kornia.filters import GuidedBlur blur = GuidedBlur(kernel_size=5, eps=0.1, border_type="reflect", subsample=1, separable=False) out_module = blur(guidance, input)值得注意的细节
- 张量 eps 现在安全:修复后
eps=torch.tensor(0.1)(float32)与float16/bfloat16输入混用不会再触发 dtype 不匹配,因为两个操作数会被统一到 float32 再求解(参见 kornia/filters/guided.py 的注释)。 - 默认
reflect边界在半精度下的边界情况:回归测试特意改用constant以隔离求解器行为;如果你的运行时(如 CPU PyTorch 2.5.x)缺少半精度reflection_pad2d,遇到 float16 与reflect组合报错时,可考虑更换border_type或升级 PyTorch 版本。 - 数值精度参考:多通道路径链式执行多次盒式滤波、一次
C x C求解和一次einsum,误差通常在几个 eps 量级(bfloat16 因输入表示精度本身有限,误差预算会略宽),测试中为 bfloat16 放宽到4 * eps容差(见 tests/filters/test_guided.py)。
小结
changelog.d/+migration-116.fixed.md记录的这次修复,本质上是"算子级半精度适配"的一个范本:先精确定位崩溃发生的 kernel(torch.linalg.solve无半精度 LU),再排查 dtype 提升这类隐性问题,最终用最小侵入的方式(求解阶段统一到 float32、完成后转回)同时解决两者,并保证非半精度路径逐位不变、编译兼容、单通道路径零改动。修复后,guided_blur/GuidedBlur在float16/bfloat16下配合多通道引导图即可稳定运行,为移动端、MPS 等偏好半精度算力的部署场景扫清了障碍。
【免费下载链接】kornia🐍 Geometric Computer Vision Library for Spatial AI项目地址: https://gitcode.com/gh_mirrors/ko/kornia
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考