ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Kornia 增强基类完全指南:从零实现自定义 2D/3D 数据增强

Kornia 增强基类完全指南:从零实现自定义 2D/3D 数据增强 计算机视觉人工智能深度学习图像处理【免费下载链接】kornia Geometric Computer Vision Library for Spatial AI项目地址https://gitcode.com/gh_mirrors/ko/kornia点击查看免费下载本文是一份面向开发者的技术指南围绕 Kornia 空间 AI 几何视觉库Geometric Computer Vision Library for Spatial AI中 docs/source/augmentation.base.rst 所定义的增强Augmentation基类体系展开。你将掌握 Kornia 增强系统的sample-apply核心流程、基类继承链中各层职责、如何通过IntensityAugmentationBase2D与GeometricAugmentationBase2D快速编写生产级自定义增强算子以及概率门控、随机参数生成、可复现性与序列化等进阶要点。阅读源码实现可同时参考 kornia/augmentation/base.py 与 kornia/augmentation/_2d/base.py。导读为什么需要理解增强基类Kornia 的kornia.augmentation子包内置了大量开箱即用的增强算子但它们全部构建在一套统一的基础设施之上。理解这些基类意味着你可以以极少的样板代码boilerplate实现自定义增强并自动获得批次概率门控、same_on_batch、keepdim、参数回放replay与AugmentationSequential容器集成能力区分“刚性rigid”与“非刚性non-rigid”增强的实现路径为坐标类数据框、关键点自动推导变换矩阵理解p/p_batch/_params/transform_matrix等内部状态的准确语义避免在随机种子、序列化和导出场景中踩坑。基类继承链每一层只负责一个关注点Kornia 的增强基类形成一条链每一层增加一个关注点并被若干具体子类共享。官方文档建议选择能完成你所需功能的最浅层基类shallowest base。nn.Module └─ _BasicAugmentationBase parameter sampling the forward skeleton ├─ _AugmentationBase dispatch to image / mask / box / keypoint / class data keys │ ├─ AugmentationBase2D 2D tensor validation (subclass for a fully custom 2D op) │ │ └─ RigidAffineAugmentationBase2D transform-matrix machinery │ │ ├─ IntensityAugmentationBase2D intensity ops — override apply_transform │ │ └─ GeometricAugmentationBase2D warp ops — also override compute_transformation │ └─ AugmentationBase3D … (the 3D mirror of the 2D chain) └─ MixAugmentationBaseV2 mix ops (MixUp / CutMix) — bypass the per-key dispatch每一层都是一个独立、可复用的轴axis参数采样、数据键分发、2D 与 3D、刚性矩阵与自由形式、强度intensity与几何geometric。其中四个公开的*Base2D类是外部代码可直接继承的公开 API。对于大多数自定义增强场景直接继承IntensityAugmentationBase2D或GeometricAugmentationBase2D即可。各层的源码落位如下基类源码位置核心职责_BasicAugmentationBasekornia/augmentation/base.py概率门控、_param_generator、flags、forward 骨架_AugmentationBasekornia/augmentation/base.py面向 image/mask/box/keypoint/class 的数据键分发AugmentationBase2Dkornia/augmentation/_2d/base.py2D 张量校验(B, C, H, W)布局RigidAffineAugmentationBase2Dkornia/augmentation/_2d/base.py变换矩阵机制与transform_matrix属性IntensityAugmentationBase2Dkornia/augmentation/_2d/intensity/base.py强度类仅需覆写apply_transformGeometricAugmentationBase2Dkornia/augmentation/_2d/geometric/base.py几何类还需覆写compute_transformation并自带逆变换MixAugmentationBaseV2kornia/augmentation/_2d/mix/base.pyMixUp / CutMix 等混合增强绕过逐键分发AugmentationBase3D/RigidAffineAugmentationBase3Dkornia/augmentation/_3d/base.py3D 链(B, C, D, H, W)布局、(B, 4, 4)矩阵从源码结构看_BasicAugmentationBase直接继承torch.nn.Module见 kornia/augmentation/base.py并通过__init_subclass__为每个子类复制一份独立的forward实现以避免 Dynamo 编译缓存冲突——这是 Kornia 对torch.compile兼容性所做的一处重要工程细节。何时选择哪个基类像素级、不改变像素位置的增强亮度、噪声、颜色扰动→IntensityAugmentationBase2D仅覆写apply_transform会移动像素位置的刚性几何增强旋转、仿射、缩放→GeometricAugmentationBase2D覆写compute_transformation与apply_transform完全自定义的 2D 算子含非刚性操作→AugmentationBase2D混合类增强MixUp / CutMix / Mosaic / Jigsaw / Transplantation→MixAugmentationBaseV23D 数据体素、点云之外的立体图像→ 3D 链对应基类。预定义增强流程sample-apply 例程Kornia 的增强算子普遍遵循sample-apply先采样、后应用的例程sample采样Kornia 的目标是在张量级别做灵活的增强即批次中每张图像可以使用不同的参数和不同的概率。采样步骤首先抽取一组随机参数并将采样得到的增强状态存入增强算子的_params属性中供回放replay使用。应用时刻的随机数如部分 VAE 潜在变量可能需要额外的 RNG 控制详见下文“随机可复现性”小节。apply应用利用生成的或用户提供的参数执行增强。除了变换图像张量外Kornia 还支持反向操作inverse还原变换以及掩码mask、关键点keypoint、边界框bounding box等其他数据模态即 Kornia 中的data keys的变换。这些能力取决于具体算子的数据键处理器实现。在AugmentationSequential容器中几何坐标类变换的派发依赖几何基类类型仅在自定义刚性基类上实现一个矩阵并不会自动启用该派发参见 issue #4481。非刚性坐标变换也不会被自动提供参见 issue #4420。这两点意味着若自定义刚性增强需要容器自动派发到框/关键点应继承GeometricAugmentationBase2D而非更浅的RigidAffineAugmentationBase2D。forward 骨架中的关键内部状态_BasicAugmentationBase.__init__kornia/augmentation/base.py初始化了以下核心状态self.p逐样本element-wise应用概率self.p_batch整批batch-wise应用概率self.same_on_batch是否整批共享同一组采样参数self.keepdim输出是否保持与输入相同的形状True还是广播为批次形式Falseself._params最近一次采样的参数字典Dict[str, Tensor]用于回放self._param_generatorRandomGeneratorBase子类的实例负责生成随机参数self.flags静态非随机配置的普通字典。forward的完整链路为__unpack_input__→transform_tensor统一为(B, C, H, W)→forward_parameters生成batch_prob、采样参数、记录forward_input_shape→_process_kwargs_to_params_and_flags合并 kwargs 覆写→apply_func→ 按keepdim还原输出形状。其中forward_parameters还会写入forward_input_shape这是几何变换正确求逆所需的关键信息kornia/augmentation/base.py。自定义强度增强继承 IntensityAugmentationBase2DIntensityAugmentationBase2D面向不移动像素位置的变换它自带compute_transformation并返回单位矩阵为大多数标注数据提供默认透传passthrough处理器。子类可以覆写这些默认行为例如RandomErasing会同时擦除掩码区域强度类直接调用transform_boxes时框会原样返回容器也会跳过强度变换对框的处理。用户只需要覆写apply_transform。最常见的场景是带随机逐样本参数的像素级增强在__init__中声明一个参数生成器然后在apply_transform中读取采样值。文档给出的完整示例RandomAddValue如下from typing import Any, Dict, Optional import torch from torch import Tensor from kornia.augmentation import IntensityAugmentationBase2D from kornia.augmentation import random_generator as rg class RandomAddValue(IntensityAugmentationBase2D): Add a per-sample value drawn uniformly from add_range. def __init__(self, add_range(0.0, 0.2), same_on_batchFalse, p1.0, keepdimFalse): super().__init__(pp, same_on_batchsame_on_batch, keepdimkeepdim) # A PlainUniformGenerator sampler is a 4-tuple (range, name, center, bound): # sample a value inside range and expose it as params[name]. center and # bound (None here) are optional constraints for centred/bounded ranges. self._param_generator rg.PlainUniformGenerator((add_range, add, None, None)) def apply_transform( self, input: Tensor, params: Dict[str, Tensor], flags: Dict[str, Any], transform: Optional[Tensor] None, ) - Tensor: # params[add] has shape (B,) — reshape to broadcast over C, H, W. add params[add].to(input).view(-1, 1, 1, 1) return input add aug RandomAddValue((0.0, 0.2), p1.0) out aug(torch.rand(4, 3, 32, 32)) # a different value per sample again aug(torch.rand(4, 3, 32, 32), paramsaug._params) # reuse the recorded drawPlainUniformGenerator 的采样协议PlainUniformGenerator位于 kornia/augmentation/random_generator/_2d/plain_uniform.py。它的每个采样器是一个四元组(factor, name, center, bound)factor可以是二元组、(2,)形状的torch.Tensor或nn.Parameter。若传入nn.Parameter该范围会被注册为模块参数并支持梯度回传若传入普通 Tensor 则注册为 buffername返回字典中的键center与bound可选的居中/边界约束二者必须同时提供或同时为None提供时_range_bound会校验factor是否落在center ± bound范围内。forward以batch_shape[0]为批次大小通过UniformDistribution采样(B,)形状的数值并支持same_on_batch此时整个批次共享同一采样值。默认在 CPU 上以float32采样可通过set_rng_device_and_dtype(device..., dtype...)调整。注意_adapted_rsampling采样后会把结果.to(device_device, dtype_dtype)转换回构造参数所在设备与精度——这就是文档中“返回参数的放置位置与采样放置位置是两回事”的由来。静态配置请放入 self.flags非随机静态配置应放入self.flags一个普通 dict并在apply_transform的flags参数中读取。例如将value填入self.flags[value]。遵循这一约定自定义增强既能独立运行也能无需额外接线地嵌入AugmentationSequential。自定义几何增强继承 GeometricAugmentationBase2D对于刚性的几何增强需要实现compute_transformation返回(B, 3, 3)变换矩阵与apply_transform实现对应的图像操作。几何基类会将矩阵传播给框与关键点掩码处理使用其掩码处理器。若要支持图像求逆还需实现inverse_transform基类提供矩阵求逆compute_inverse_transformation默认调用_torch_inverse_cast见 kornia/augmentation/_2d/geometric/base.py。某些配置如 slice 模式裁剪会拒绝求逆。文档给出MyRandomTransform示例仅用于说明方法签名并非可逆的几何扭曲from typing import Any, Dict, Optional import torch from torch import Tensor import kornia as K from kornia.augmentation import GeometricAugmentationBase2D from kornia.augmentation import random_generator as rg class MyRandomTransform(GeometricAugmentationBase2D): def __init__( self, factor(0., 1.), same_on_batch: bool False, p: float 1.0, keepdim: bool False, ) - None: super().__init__(pp, same_on_batchsame_on_batch, keepdimkeepdim) self._param_generator rg.PlainUniformGenerator((factor, factor, None, None)) def compute_transformation(self, input, params, flags): # return the (B, 3, 3) transform matrix for this augmentation # identity shown only to illustrate the required matrix shape return K.eye_like(3, input) def apply_transform( self, input: Tensor, params: Dict[str, Tensor], flags: Dict[str, Any], transform: Optional[Tensor] None ) - Tensor: factor params[factor].to(input).view(-1, 1, 1, 1) return input * factor矩阵生成与惰性求值RigidAffineAugmentationBase2D增加了矩阵机制generate_transformation_matrix会先调用compute_transformation得到“应用矩阵”再结合batch_prob与单位矩阵做torch.where混合kornia/augmentation/_2d/base.py。当p 1.0 and p_batch 1.0时走快速路径直接返回应用矩阵跳过单位矩阵构建与混合——这是一处热路径优化。强度类设置_compute_matrix_lazily Truekornia/augmentation/_2d/intensity/base.py因为其apply_transform不读取矩阵构建矩阵属于纯开销矩阵推迟到首次读取transform_matrix时才计算。apply_func中通过_kornia_input_metadata_only标记检测四个关键方法transform_tensor、generate_transformation_matrix、compute_transformation、identity_matrix是否仅读取元数据从而决定惰性状态中保留的是紧凑的_InputMetadatashape/dtype/device还是完整输入张量kornia/augmentation/_2d/base.py。几何基类的坐标数据派发GeometricAugmentationBase2D为框与关键点提供矩阵驱动的处理器apply_transform_box若transform为None则取self.transform_matrix否则抛出RuntimeError随后调用input.transform_boxes_(transform)kornia/augmentation/_2d/geometric/base.pyapply_transform_keypoint同样逻辑调用input.transform_keypoints_(transform)掩码默认使用最近邻nearest重采样并默认关闭抗锯齿以保证离散标签不被破坏可通过transform_masks(..., resample..., antialiasTrue)显式覆写。inverse流程kornia/augmentation/_2d/geometric/base.py会先取当前变换矩阵get_transformation_matrix经compute_inverse_transformation求逆后逐数据键还原。文档同时提醒逆变换的重采样无法恢复裁剪、填充或插值造成的图像/掩码信息丢失张量形式的框经轴对齐包围后可能丢失旋转角。非刚性增强apply_transform* 与 apply_non_transform* API对于非刚性增强例如随机裁剪、挖洞用户可按需实现apply_transform*与apply_non_transform*系列 APIapply_transform*作用于批次中被选中执行增强的元素apply_non_transform*作用于被跳过batch_prob False的元素。例如裁剪操作会改变被选中元素的尺寸因此被跳过的元素也必须做尺寸调整如 resize以保证整个批次张量保持单一尺寸。_AugmentationBase中的_blend_by_probkornia/augmentation/base.py在两条分支形状一致时使用torch.where混合可 ONNX 导出、可 fullgraph 编译形状不一致时退化为 Python 分支不可 ONNX 导出。当p 1.0 and p_batch 1.0时直接返回变换结果跳过混合以支持 fullgraph 编译。AugmentationBase2D暴露的完整钩子方法包括apply_transform、apply_non_transform、apply_transform_mask、apply_non_transform_mask、apply_transform_box、apply_non_transform_box、apply_transform_keypoint、apply_non_transform_keypoint、apply_transform_class、apply_non_transform_class。其中未变换分支的掩码、框、关键点、类别处理器在基类中默认透传。3D 增强链AugmentationBase3D 体系3D 基类提供与 2D 链类似的形状、矩阵与数据键机制工作布局为(B, C, D, H, W)变换矩阵为(B, 4, 4)identity_matrix返回eye_like(4, input)见 kornia/augmentation/_3d/base.py。dtype 守卫仅接受float16、bfloat16、float32、float64。3D 基类不实现独立的几何求逆在AugmentationSequential中几何 3D 子节点在求逆时会抛出异常而强度 3D 子节点会被跳过其效果保持应用。3D 链构成AugmentationBase3D→RigidAffineAugmentationBase3D→GeometricAugmentationBase3D/IntensityAugmentationBase3D。其中RandomTransplantation3D例外地同时继承MixAugmentationBaseV2与AugmentationBase3D。混合增强MixAugmentationBaseV2混合类增强MixUp / CutMix 等派生自MixAugmentationBaseV2kornia/augmentation/_2d/mix/base.py拥有独立的 forward 与数据键契约不继承AugmentationBase2D的全部 forward 约定。实现者需要提供generate_parameters与apply_transform且apply_transform需在内部自行处理概率。几个关键约定混合类不是几何增强没有变换矩阵访问transform_matrix抛出RuntimeError也没有求逆inverse()只接受关键字参数并抛出RuntimeErrordata_keys决定哪些位置参数被派发支持input、image、mask、bbox、bbox_xyxy、bbox_xywh、keypoints、class、label未实现的键在采样前即抛出NotImplementedError移植类在两个移植类中于参数绘制后才抛出有些混合增强会组合不同样本RandomJigsaw在每张图像内部重排 patchRandomMixUpV2接受类别标签而RandomMosaic接受框但不实现类别标签变换RandomTransplantation与RandomTransplantation3D需要分割掩码仅掩码调用需传data_keys[mask]。对应实现位于 kornia/augmentation/_2d/mix/ 下的mixup.py、mosaic.py、jigsaw.py、transplantation.py等文件。进一步说明概率、随机生成器、可复现性与序列化概率p 与 p_batch 的双层门控_BasicAugmentationBase拥有逐样本的p与整批的p_batch两个门控。具体构造器可将公开的p映射到其中任意一个门控因此需要查阅对应类的契约例如RandomMixUpV2门控整批而RandomJigsaw门控单个样本。__batch_prob_generator__kornia/augmentation/base.py的实现细节当0 p_batch 1时基类在逐样本门控之前先抽取一次批次级 Bernoulli端点为 0 或 1 时跳过 Bernoulli 抽取p_batch 1时批次门为全 1p_batch 0时为全 0逐样本门控中p 1/p 0直接走常量分支可在编译期解析不产生 graph breaksame_on_batchTrue时只采样一次再 expand两层结果以无分支方式相乘合并避免数据相关的if分支造成 graph break。因此p1.0, p_batch0.0时没有任何样本被选中。注意只有部分具体构造器直接暴露p_batchissue #4425 跟踪此限制且门控在整批变换计算完成后才选择被跳过的样本仍可能抛错或携带 NaN 梯度issue #4576。随机生成器继承 RandomGeneratorBase要获得自动生成的、列出全部自定义参数的__repr__应通过继承RandomGeneratorBase实现_param_generator来生成随机参数并将所有静态参数放入self.flags。__repr__会自动拼接_param_generator的字符串表示与flags中的每一项枚举值会以name.lower()形式输出见 kornia/augmentation/base.py。PlainUniformGenerator可用于以更少样板代码生成简单均匀参数。随机可复现性参数采样通常在 CPU 上进行与图像所在设备无关。set_rng_device_and_dtype会请求新的采样器放置位置与精度但并非每个内部张量都会遵循该请求部分生成器/设备组合在前向过程中仍可能失败issue #4426。返回参数的放置位置与采样放置位置是独立的构造器中的范围与类型转换可能把采样后的张量放到其他设备或其他 dtype 上。完整细节全局种子、DataLoader worker 种子、消耗顺序、回放与采样器配置限制参见 docs/source/get-started/conventions.rst。值得注意应用时刻的随机性并非总是被记录RandomDissolving的 VAE 潜在变量需要控制其随机状态才能回放基类对关键字与不完整参数的处理并不统一适用于混合增强通过forward(x, params...)回放时参数字典是整体替换按引用存储而非合并torch.manual_seed在调用前设置即可复现同后端同 dtype 的采样。序列化几个与序列化相关的要点均以源码与文档为准若干构造器接受nn.Parameter范围并可向其传播梯度但load_state_dict之后数值范围 buffer 未必与缓存的采样器保持连接——要改变采样范围需重建这些配置issue #4428kornia.augmentation.auto策略只有在其中每个算子包装器都可 pickle 时才能被 pickle因此默认策略不能issue #4469而空的AutoAugment(policy[[]])与只含posterize条目的策略可以内置的 2D 强度增强、翻转、Resize含LongestMaxSize与SmallestMaxSize以及 slice 模式RandomResizedCrop只在被请求时才计算变换矩阵。它们的挂起矩阵状态仅保留输入的 shape、dtype 与 device 以及变换参数不保留图像张量本身。因此对这些模块执行pickle、copy.deepcopy、torch.save不会仅仅因为矩阵尚未被读取而携带最后一整批图像参数含其既有梯度连接仍可用于回放。自定义惰性子类若在transform_tensor、generate_transformation_matrix、compute_transformation或identity_matrix中使用像素值则保留原始的基于输入的惰性行为保留输入直到矩阵被读取或被下一次 forward 替换原样继承内置实现则保持紧凑的元数据状态——覆写并不隐式承诺可以脱离像素数据运行。结语Kornia 的增强基类体系将“采样—应用—派发—求逆—序列化”的关注点拆解到清晰的继承链中。日常使用中强度类仅需覆写apply_transform几何类补充compute_transformation以及可选的inverse_transform即可获得与内置算子一致的批次概率门控、坐标数据派发与容器集成能力。更深入的采样契约、回放语义与序列化边界可继续研读 docs/source/get-started/conventions.rst 以及 kornia/augmentation/ 下各基类与random_generator的实现。赞分享计算机视觉人工智能深度学习图像处理【免费下载链接】kornia Geometric Computer Vision Library for Spatial AI项目地址https://gitcode.com/gh_mirrors/ko/kornia点击查看免费下载相关推荐Electrobun 桌面应用框架1.5MB 的原生层为何能替代 ElectronElectrobun 桌面应用框架1.5MB 的原生层为何能替代 Electron Electrobun 桌面应用框架是一个用 TypeScript 编写主进计算机视觉深度学习人工智能图像处理Kornia 可微数据增强全解析2D/3D 几何、颜色空间与 Mix 增强操作指南Kornia 可微数据增强全解析2D/3D 几何、颜色空间与 Mix 增强操作指南 Kornia 的可微数据增强Differentiable Data Au计算机视觉人工智能深度学习图像处理Kornia 可微分数据增强Differentiable Data Augmentation完全指南从 2D/3D 算子到随机策略Kornia 可微分数据增强Differentiable Data Augmentation完全指南从 2D/3D 算子到随机策略 Kornia 的 ko计算机视觉深度学习人工智能图像处理上一篇从0到1集成ExpandableTextViewAndroid开发者的完整实践教程下一篇HsMod炉石传说BepInEx框架下的全能游戏增强插件创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进